From cde38a0a2db6fde8cdb014176681a485636b6a2b Mon Sep 17 00:00:00 2001 From: Martin Taillefer Date: Mon, 14 Sep 2026 22:00:38 +0000 Subject: [PATCH] Introduce the http_headers crate for fast and robust HTTP header parsing Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 11e4f879-cc17-4c84-8590-fca44b4a456b Copilot-Session: 421aaf0a-555e-4b07-998d-1ed192e45d7a Copilot-Session: 7ac05d58-6b4f-4276-86c1-af5168a8354c --- .anvil.lock | 10 + .cargo/mutants.toml | 2 + .github/workflows/repository-checks.yml | 47 +- .spelling | 51 + CHANGELOG.md | 2 + Cargo.lock | 155 +- Cargo.toml | 15 + README.md | 1 + crates/http_headers/CHANGELOG.md | 5 + crates/http_headers/Cargo.toml | 183 + crates/http_headers/README.md | 465 +++ .../benches/http_headers_auth_cors_shapes.rs | 488 +++ .../http_headers_authority_semantics.rs | 341 ++ .../benches/http_headers_common_values.rs | 7 + .../http_headers_conditional_range_shapes.rs | 139 + .../benches/http_headers_fixtures.rs | 197 ++ .../http_headers_location_semantics.rs | 279 ++ .../benches/http_headers_micro.rs | 3078 +++++++++++++++++ .../benches/http_headers_name_recognition.rs | 151 + .../http_headers_negotiation_semantics.rs | 646 ++++ .../http_headers_negotiation_shapes.rs | 250 ++ .../benches/http_headers_operations.rs | 171 + .../benches/http_headers_per_header.rs | 685 ++++ .../benches/http_headers_policy.rs | 257 ++ .../benches/http_headers_policy_shapes.rs | 127 + .../benches/http_headers_protocol.rs | 260 ++ .../benches/http_headers_shapes_common.rs | 268 ++ .../benches/http_headers_storage.rs | 680 ++++ .../benches/http_headers_typed_tokens.rs | 158 + .../benches/http_headers_websocket_shapes.rs | 158 + crates/http_headers/docs/COMPATIBILITY.md | 201 ++ crates/http_headers/docs/DESIGN.md | 388 +++ crates/http_headers/docs/PERF.md | 60 + crates/http_headers/docs/TODO.md | 1417 ++++++++ crates/http_headers/examples/axum.rs | 31 + crates/http_headers/examples/axum/app.rs | 46 + crates/http_headers/favicon.ico | 3 + crates/http_headers/logo.png | 3 + crates/http_headers/scripts/perf_report.rs | 505 +++ crates/http_headers/src/decode_error.rs | 206 ++ crates/http_headers/src/field.rs | 581 ++++ crates/http_headers/src/field_name.rs | 830 +++++ crates/http_headers/src/field_value.rs | 1399 ++++++++ .../http_headers/src/headers/authorization.rs | 1362 ++++++++ .../http_headers/src/headers/cache_control.rs | 1445 ++++++++ .../src/headers/conditional/if_match.rs | 24 + .../headers/conditional/if_modified_since.rs | 20 + .../src/headers/conditional/if_none_match.rs | 24 + .../src/headers/conditional/if_range.rs | 497 +++ .../conditional/if_unmodified_since.rs | 20 + .../src/headers/conditional/last_modified.rs | 20 + .../src/headers/conditional/mod.rs | 27 + .../src/headers/conditional/shared.rs | 1893 ++++++++++ .../src/headers/content_length.rs | 232 ++ .../http_headers/src/headers/content_type.rs | 1247 +++++++ .../cors/access_control_allow_credentials.rs | 356 ++ .../cors/access_control_allow_headers.rs | 40 + .../cors/access_control_allow_methods.rs | 703 ++++ .../cors/access_control_allow_origin.rs | 1364 ++++++++ .../access_control_allow_origin/components.rs | 177 + .../cors/access_control_expose_headers.rs | 40 + .../headers/cors/access_control_max_age.rs | 334 ++ .../cors/access_control_request_headers.rs | 21 + .../cors/access_control_request_method.rs | 537 +++ .../src/headers/cors/cors_header_names.rs | 54 + .../src/headers/cors/cors_methods.rs | 54 + .../src/headers/cors/cors_tokens.rs | 37 + crates/http_headers/src/headers/cors/mod.rs | 46 + .../http_headers/src/headers/cors/shared.rs | 1576 +++++++++ .../http_headers/src/headers/cors/test_map.rs | 47 + crates/http_headers/src/headers/etag.rs | 606 ++++ .../src/headers/extension_value.rs | 11 + .../src/headers/field_name_view.rs | 140 + .../src/headers/invalid_method.rs | 17 + crates/http_headers/src/headers/location.rs | 622 ++++ .../src/headers/location/component.rs | 57 + .../src/headers/location/construction.rs | 158 + .../src/headers/location/metadata.rs | 83 + .../src/headers/location/uri_authority.rs | 107 + .../src/headers/location/uri_reference.rs | 134 + .../http_headers/src/headers/method_view.rs | 121 + crates/http_headers/src/headers/mod.rs | 183 + .../src/headers/negotiation/accept.rs | 497 +++ .../headers/negotiation/accept_encoding.rs | 250 ++ .../negotiation/accept_encoding_entry.rs | 69 + .../src/headers/negotiation/accept_entry.rs | 294 ++ .../headers/negotiation/accept_language.rs | 270 ++ .../negotiation/accept_language_entry.rs | 69 + .../src/headers/negotiation/accept_scan.rs | 411 +++ .../src/headers/negotiation/allow.rs | 132 + .../src/headers/negotiation/content_coding.rs | 113 + .../src/headers/negotiation/host.rs | 1175 +++++++ .../headers/negotiation/host/components.rs | 165 + .../src/headers/negotiation/language_range.rs | 92 + .../src/headers/negotiation/media_range.rs | 100 + .../src/headers/negotiation/mod.rs | 66 + .../negotiation/negotiation_members.rs | 89 + .../negotiation/negotiation_parameter.rs | 185 + .../headers/negotiation/negotiation_token.rs | 88 + .../src/headers/negotiation/quality.rs | 296 ++ .../negotiation/recognition_test_support.rs | 88 + .../src/headers/negotiation/server.rs | 380 ++ .../src/headers/negotiation/shared.rs | 2193 ++++++++++++ .../src/headers/negotiation/vary.rs | 173 + .../headers/negotiation/vary_entry_view.rs | 70 + .../negotiation/weighted_token_scan.rs | 394 +++ .../src/headers/range/accept_ranges.rs | 891 +++++ .../src/headers/range/content_range.rs | 1077 ++++++ crates/http_headers/src/headers/range/mod.rs | 20 + .../http_headers/src/headers/range/range.rs | 1346 +++++++ .../http_headers/src/headers/range/shared.rs | 126 + .../security/content_security_policy.rs | 480 +++ .../http_headers/src/headers/security/mod.rs | 20 + .../src/headers/security/referrer_policy.rs | 909 +++++ .../security/strict_transport_security.rs | 1301 +++++++ .../security/x_content_type_options.rs | 270 ++ crates/http_headers/src/headers/set_cookie.rs | 576 +++ crates/http_headers/src/headers/shared.rs | 479 +++ crates/http_headers/src/headers/tokens.rs | 54 + crates/http_headers/src/headers/user_agent.rs | 359 ++ .../http_headers/src/headers/websocket/mod.rs | 25 + .../websocket/sec_web_socket_accept.rs | 368 ++ .../websocket/sec_web_socket_extensions.rs | 1419 ++++++++ .../headers/websocket/sec_web_socket_key.rs | 332 ++ .../websocket/sec_web_socket_protocol.rs | 604 ++++ .../websocket/sec_web_socket_version.rs | 762 ++++ .../src/headers/websocket/shared.rs | 672 ++++ crates/http_headers/src/http_adapter.rs | 627 ++++ crates/http_headers/src/lib.rs | 547 +++ crates/http_headers/src/miri_http_map.rs | 30 + crates/http_headers/src/serde_impls.rs | 957 +++++ .../http_headers/src/sink/encoded_values.rs | 479 +++ crates/http_headers/src/sink/field_encoder.rs | 319 ++ crates/http_headers/src/sink/field_sink.rs | 613 ++++ .../http_headers/src/sink/field_sink_ext.rs | 975 ++++++ crates/http_headers/src/sink/insert_error.rs | 115 + crates/http_headers/src/sink/mod.rs | 73 + .../src/source/delimited_items.rs | 380 ++ crates/http_headers/src/source/field_lines.rs | 1047 ++++++ .../http_headers/src/source/field_source.rs | 86 + .../src/source/list_item_count.rs | 45 + crates/http_headers/src/source/mod.rs | 34 + crates/http_headers/src/test_sink.rs | 67 + crates/http_headers/src/test_support.rs | 49 + crates/http_headers/src/validate.rs | 213 ++ .../http_headers/tests/__fuzz__/campaign.toml | 39 + .../tests/access_control_allow_credentials.rs | 64 + crates/http_headers/tests/axum_example.rs | 84 + .../http_headers/tests/benchmark_ownership.rs | 208 ++ crates/http_headers/tests/bolero_fuzz.rs | 1004 ++++++ .../http_headers/tests/collection_traits.rs | 179 + .../tests/common/http_headers_name_corpus.rs | 106 + .../common/http_headers_storage_operations.rs | 94 + crates/http_headers/tests/common/mod.rs | 6 + crates/http_headers/tests/common/test_map.rs | 43 + crates/http_headers/tests/conformance_api.rs | 424 +++ .../tests/content_type_identity.rs | 128 + crates/http_headers/tests/core_api.rs | 621 ++++ crates/http_headers/tests/cors_iteration.rs | 101 + crates/http_headers/tests/cors_tokens.rs | 144 + .../http_headers/tests/decode_mode_parity.rs | 387 +++ .../tests/documentation_examples.rs | 59 + .../http_headers/tests/duration_precision.rs | 100 + .../tests/error_and_encoded_values.rs | 160 + .../http_headers/tests/feature_boundaries.rs | 34 + .../tests/field_lines_iteration.rs | 81 + crates/http_headers/tests/field_value_text.rs | 79 + crates/http_headers/tests/header_families.rs | 95 + crates/http_headers/tests/host_debug.rs | 124 + .../tests/host_origin_components.rs | 876 +++++ crates/http_headers/tests/hsts_syntax.rs | 54 + crates/http_headers/tests/http_map.rs | 283 ++ crates/http_headers/tests/http_messages.rs | 88 + crates/http_headers/tests/location.rs | 935 +++++ crates/http_headers/tests/moved_public_api.rs | 1674 +++++++++ .../tests/negotiation_semantics.rs | 692 ++++ .../tests/negotiation_source_stability.rs | 181 + .../http_headers/tests/owned_token_lists.rs | 68 + crates/http_headers/tests/range_invariants.rs | 331 ++ .../tests/raw_source_validation.rs | 198 ++ .../http_headers/tests/response_failures.rs | 117 + .../tests/sec_web_socket_version.rs | 93 + crates/http_headers/tests/serde.rs | 574 +++ crates/http_headers/tests/source_limits.rs | 521 +++ crates/http_headers/tests/support_api.rs | 80 + crates/http_headers/tests/typed_tokens.rs | 222 ++ .../tests/websocket_source_limits.rs | 157 + crates/http_headers_simd/CHANGELOG.md | 5 + crates/http_headers_simd/Cargo.toml | 43 + crates/http_headers_simd/README.md | 37 + .../http_headers_simd_no_std_dispatch.rs | 238 ++ crates/http_headers_simd/favicon.ico | 3 + crates/http_headers_simd/logo.png | 3 + .../src/__fuzz__/campaign.toml | 24 + .../differential_properties/corpus/boundaries | 1 + .../differential_properties/crashes/.gitkeep | 0 .../neon_matches_scalar/corpus/boundaries | 1 + .../neon_matches_scalar/crashes/.gitkeep | 0 .../sse2_matches_scalar/corpus/boundaries | 1 + .../sse2_matches_scalar/crashes/.gitkeep | 0 crates/http_headers_simd/src/api.rs | 947 +++++ crates/http_headers_simd/src/arm.rs | 541 +++ crates/http_headers_simd/src/base64.rs | 35 + crates/http_headers_simd/src/benchmarking.rs | 322 ++ crates/http_headers_simd/src/dispatch.rs | 1163 +++++++ crates/http_headers_simd/src/lib.rs | 61 + crates/http_headers_simd/src/list.rs | 455 +++ crates/http_headers_simd/src/range.rs | 463 +++ crates/http_headers_simd/src/scalar.rs | 87 + crates/http_headers_simd/src/tracking.rs | 455 +++ crates/http_headers_simd/src/uri.rs | 144 + crates/http_headers_simd/src/x86.rs | 1247 +++++++ .../tests/__fuzz__/campaign.toml | 14 + .../crashes/.gitkeep | 0 crates/http_headers_simd/tests/bolero_fuzz.rs | 63 + docs/design/README.md | 70 + justfile | 65 + 217 files changed, 73781 insertions(+), 2 deletions(-) create mode 100644 crates/http_headers/CHANGELOG.md create mode 100644 crates/http_headers/Cargo.toml create mode 100644 crates/http_headers/README.md create mode 100644 crates/http_headers/benches/http_headers_auth_cors_shapes.rs create mode 100644 crates/http_headers/benches/http_headers_authority_semantics.rs create mode 100644 crates/http_headers/benches/http_headers_common_values.rs create mode 100644 crates/http_headers/benches/http_headers_conditional_range_shapes.rs create mode 100644 crates/http_headers/benches/http_headers_fixtures.rs create mode 100644 crates/http_headers/benches/http_headers_location_semantics.rs create mode 100644 crates/http_headers/benches/http_headers_micro.rs create mode 100644 crates/http_headers/benches/http_headers_name_recognition.rs create mode 100644 crates/http_headers/benches/http_headers_negotiation_semantics.rs create mode 100644 crates/http_headers/benches/http_headers_negotiation_shapes.rs create mode 100644 crates/http_headers/benches/http_headers_operations.rs create mode 100644 crates/http_headers/benches/http_headers_per_header.rs create mode 100644 crates/http_headers/benches/http_headers_policy.rs create mode 100644 crates/http_headers/benches/http_headers_policy_shapes.rs create mode 100644 crates/http_headers/benches/http_headers_protocol.rs create mode 100644 crates/http_headers/benches/http_headers_shapes_common.rs create mode 100644 crates/http_headers/benches/http_headers_storage.rs create mode 100644 crates/http_headers/benches/http_headers_typed_tokens.rs create mode 100644 crates/http_headers/benches/http_headers_websocket_shapes.rs create mode 100644 crates/http_headers/docs/COMPATIBILITY.md create mode 100644 crates/http_headers/docs/DESIGN.md create mode 100644 crates/http_headers/docs/PERF.md create mode 100644 crates/http_headers/docs/TODO.md create mode 100644 crates/http_headers/examples/axum.rs create mode 100644 crates/http_headers/examples/axum/app.rs create mode 100644 crates/http_headers/favicon.ico create mode 100644 crates/http_headers/logo.png create mode 100755 crates/http_headers/scripts/perf_report.rs create mode 100644 crates/http_headers/src/decode_error.rs create mode 100644 crates/http_headers/src/field.rs create mode 100644 crates/http_headers/src/field_name.rs create mode 100644 crates/http_headers/src/field_value.rs create mode 100644 crates/http_headers/src/headers/authorization.rs create mode 100644 crates/http_headers/src/headers/cache_control.rs create mode 100644 crates/http_headers/src/headers/conditional/if_match.rs create mode 100644 crates/http_headers/src/headers/conditional/if_modified_since.rs create mode 100644 crates/http_headers/src/headers/conditional/if_none_match.rs create mode 100644 crates/http_headers/src/headers/conditional/if_range.rs create mode 100644 crates/http_headers/src/headers/conditional/if_unmodified_since.rs create mode 100644 crates/http_headers/src/headers/conditional/last_modified.rs create mode 100644 crates/http_headers/src/headers/conditional/mod.rs create mode 100644 crates/http_headers/src/headers/conditional/shared.rs create mode 100644 crates/http_headers/src/headers/content_length.rs create mode 100644 crates/http_headers/src/headers/content_type.rs create mode 100644 crates/http_headers/src/headers/cors/access_control_allow_credentials.rs create mode 100644 crates/http_headers/src/headers/cors/access_control_allow_headers.rs create mode 100644 crates/http_headers/src/headers/cors/access_control_allow_methods.rs create mode 100644 crates/http_headers/src/headers/cors/access_control_allow_origin.rs create mode 100644 crates/http_headers/src/headers/cors/access_control_allow_origin/components.rs create mode 100644 crates/http_headers/src/headers/cors/access_control_expose_headers.rs create mode 100644 crates/http_headers/src/headers/cors/access_control_max_age.rs create mode 100644 crates/http_headers/src/headers/cors/access_control_request_headers.rs create mode 100644 crates/http_headers/src/headers/cors/access_control_request_method.rs create mode 100644 crates/http_headers/src/headers/cors/cors_header_names.rs create mode 100644 crates/http_headers/src/headers/cors/cors_methods.rs create mode 100644 crates/http_headers/src/headers/cors/cors_tokens.rs create mode 100644 crates/http_headers/src/headers/cors/mod.rs create mode 100644 crates/http_headers/src/headers/cors/shared.rs create mode 100644 crates/http_headers/src/headers/cors/test_map.rs create mode 100644 crates/http_headers/src/headers/etag.rs create mode 100644 crates/http_headers/src/headers/extension_value.rs create mode 100644 crates/http_headers/src/headers/field_name_view.rs create mode 100644 crates/http_headers/src/headers/invalid_method.rs create mode 100644 crates/http_headers/src/headers/location.rs create mode 100644 crates/http_headers/src/headers/location/component.rs create mode 100644 crates/http_headers/src/headers/location/construction.rs create mode 100644 crates/http_headers/src/headers/location/metadata.rs create mode 100644 crates/http_headers/src/headers/location/uri_authority.rs create mode 100644 crates/http_headers/src/headers/location/uri_reference.rs create mode 100644 crates/http_headers/src/headers/method_view.rs create mode 100644 crates/http_headers/src/headers/mod.rs create mode 100644 crates/http_headers/src/headers/negotiation/accept.rs create mode 100644 crates/http_headers/src/headers/negotiation/accept_encoding.rs create mode 100644 crates/http_headers/src/headers/negotiation/accept_encoding_entry.rs create mode 100644 crates/http_headers/src/headers/negotiation/accept_entry.rs create mode 100644 crates/http_headers/src/headers/negotiation/accept_language.rs create mode 100644 crates/http_headers/src/headers/negotiation/accept_language_entry.rs create mode 100644 crates/http_headers/src/headers/negotiation/accept_scan.rs create mode 100644 crates/http_headers/src/headers/negotiation/allow.rs create mode 100644 crates/http_headers/src/headers/negotiation/content_coding.rs create mode 100644 crates/http_headers/src/headers/negotiation/host.rs create mode 100644 crates/http_headers/src/headers/negotiation/host/components.rs create mode 100644 crates/http_headers/src/headers/negotiation/language_range.rs create mode 100644 crates/http_headers/src/headers/negotiation/media_range.rs create mode 100644 crates/http_headers/src/headers/negotiation/mod.rs create mode 100644 crates/http_headers/src/headers/negotiation/negotiation_members.rs create mode 100644 crates/http_headers/src/headers/negotiation/negotiation_parameter.rs create mode 100644 crates/http_headers/src/headers/negotiation/negotiation_token.rs create mode 100644 crates/http_headers/src/headers/negotiation/quality.rs create mode 100644 crates/http_headers/src/headers/negotiation/recognition_test_support.rs create mode 100644 crates/http_headers/src/headers/negotiation/server.rs create mode 100644 crates/http_headers/src/headers/negotiation/shared.rs create mode 100644 crates/http_headers/src/headers/negotiation/vary.rs create mode 100644 crates/http_headers/src/headers/negotiation/vary_entry_view.rs create mode 100644 crates/http_headers/src/headers/negotiation/weighted_token_scan.rs create mode 100644 crates/http_headers/src/headers/range/accept_ranges.rs create mode 100644 crates/http_headers/src/headers/range/content_range.rs create mode 100644 crates/http_headers/src/headers/range/mod.rs create mode 100644 crates/http_headers/src/headers/range/range.rs create mode 100644 crates/http_headers/src/headers/range/shared.rs create mode 100644 crates/http_headers/src/headers/security/content_security_policy.rs create mode 100644 crates/http_headers/src/headers/security/mod.rs create mode 100644 crates/http_headers/src/headers/security/referrer_policy.rs create mode 100644 crates/http_headers/src/headers/security/strict_transport_security.rs create mode 100644 crates/http_headers/src/headers/security/x_content_type_options.rs create mode 100644 crates/http_headers/src/headers/set_cookie.rs create mode 100644 crates/http_headers/src/headers/shared.rs create mode 100644 crates/http_headers/src/headers/tokens.rs create mode 100644 crates/http_headers/src/headers/user_agent.rs create mode 100644 crates/http_headers/src/headers/websocket/mod.rs create mode 100644 crates/http_headers/src/headers/websocket/sec_web_socket_accept.rs create mode 100644 crates/http_headers/src/headers/websocket/sec_web_socket_extensions.rs create mode 100644 crates/http_headers/src/headers/websocket/sec_web_socket_key.rs create mode 100644 crates/http_headers/src/headers/websocket/sec_web_socket_protocol.rs create mode 100644 crates/http_headers/src/headers/websocket/sec_web_socket_version.rs create mode 100644 crates/http_headers/src/headers/websocket/shared.rs create mode 100644 crates/http_headers/src/http_adapter.rs create mode 100644 crates/http_headers/src/lib.rs create mode 100644 crates/http_headers/src/miri_http_map.rs create mode 100644 crates/http_headers/src/serde_impls.rs create mode 100644 crates/http_headers/src/sink/encoded_values.rs create mode 100644 crates/http_headers/src/sink/field_encoder.rs create mode 100644 crates/http_headers/src/sink/field_sink.rs create mode 100644 crates/http_headers/src/sink/field_sink_ext.rs create mode 100644 crates/http_headers/src/sink/insert_error.rs create mode 100644 crates/http_headers/src/sink/mod.rs create mode 100644 crates/http_headers/src/source/delimited_items.rs create mode 100644 crates/http_headers/src/source/field_lines.rs create mode 100644 crates/http_headers/src/source/field_source.rs create mode 100644 crates/http_headers/src/source/list_item_count.rs create mode 100644 crates/http_headers/src/source/mod.rs create mode 100644 crates/http_headers/src/test_sink.rs create mode 100644 crates/http_headers/src/test_support.rs create mode 100644 crates/http_headers/src/validate.rs create mode 100644 crates/http_headers/tests/__fuzz__/campaign.toml create mode 100644 crates/http_headers/tests/access_control_allow_credentials.rs create mode 100644 crates/http_headers/tests/axum_example.rs create mode 100644 crates/http_headers/tests/benchmark_ownership.rs create mode 100644 crates/http_headers/tests/bolero_fuzz.rs create mode 100644 crates/http_headers/tests/collection_traits.rs create mode 100644 crates/http_headers/tests/common/http_headers_name_corpus.rs create mode 100644 crates/http_headers/tests/common/http_headers_storage_operations.rs create mode 100644 crates/http_headers/tests/common/mod.rs create mode 100644 crates/http_headers/tests/common/test_map.rs create mode 100644 crates/http_headers/tests/conformance_api.rs create mode 100644 crates/http_headers/tests/content_type_identity.rs create mode 100644 crates/http_headers/tests/core_api.rs create mode 100644 crates/http_headers/tests/cors_iteration.rs create mode 100644 crates/http_headers/tests/cors_tokens.rs create mode 100644 crates/http_headers/tests/decode_mode_parity.rs create mode 100644 crates/http_headers/tests/documentation_examples.rs create mode 100644 crates/http_headers/tests/duration_precision.rs create mode 100644 crates/http_headers/tests/error_and_encoded_values.rs create mode 100644 crates/http_headers/tests/feature_boundaries.rs create mode 100644 crates/http_headers/tests/field_lines_iteration.rs create mode 100644 crates/http_headers/tests/field_value_text.rs create mode 100644 crates/http_headers/tests/header_families.rs create mode 100644 crates/http_headers/tests/host_debug.rs create mode 100644 crates/http_headers/tests/host_origin_components.rs create mode 100644 crates/http_headers/tests/hsts_syntax.rs create mode 100644 crates/http_headers/tests/http_map.rs create mode 100644 crates/http_headers/tests/http_messages.rs create mode 100644 crates/http_headers/tests/location.rs create mode 100644 crates/http_headers/tests/moved_public_api.rs create mode 100644 crates/http_headers/tests/negotiation_semantics.rs create mode 100644 crates/http_headers/tests/negotiation_source_stability.rs create mode 100644 crates/http_headers/tests/owned_token_lists.rs create mode 100644 crates/http_headers/tests/range_invariants.rs create mode 100644 crates/http_headers/tests/raw_source_validation.rs create mode 100644 crates/http_headers/tests/response_failures.rs create mode 100644 crates/http_headers/tests/sec_web_socket_version.rs create mode 100644 crates/http_headers/tests/serde.rs create mode 100644 crates/http_headers/tests/source_limits.rs create mode 100644 crates/http_headers/tests/support_api.rs create mode 100644 crates/http_headers/tests/typed_tokens.rs create mode 100644 crates/http_headers/tests/websocket_source_limits.rs create mode 100644 crates/http_headers_simd/CHANGELOG.md create mode 100644 crates/http_headers_simd/Cargo.toml create mode 100644 crates/http_headers_simd/README.md create mode 100644 crates/http_headers_simd/benches/http_headers_simd_no_std_dispatch.rs create mode 100644 crates/http_headers_simd/favicon.ico create mode 100644 crates/http_headers_simd/logo.png create mode 100644 crates/http_headers_simd/src/__fuzz__/campaign.toml create mode 100644 crates/http_headers_simd/src/__fuzz__/differential_properties/corpus/boundaries create mode 100644 crates/http_headers_simd/src/__fuzz__/differential_properties/crashes/.gitkeep create mode 100644 crates/http_headers_simd/src/__fuzz__/neon_matches_scalar/corpus/boundaries create mode 100644 crates/http_headers_simd/src/__fuzz__/neon_matches_scalar/crashes/.gitkeep create mode 100644 crates/http_headers_simd/src/__fuzz__/sse2_matches_scalar/corpus/boundaries create mode 100644 crates/http_headers_simd/src/__fuzz__/sse2_matches_scalar/crashes/.gitkeep create mode 100644 crates/http_headers_simd/src/api.rs create mode 100644 crates/http_headers_simd/src/arm.rs create mode 100644 crates/http_headers_simd/src/base64.rs create mode 100644 crates/http_headers_simd/src/benchmarking.rs create mode 100644 crates/http_headers_simd/src/dispatch.rs create mode 100644 crates/http_headers_simd/src/lib.rs create mode 100644 crates/http_headers_simd/src/list.rs create mode 100644 crates/http_headers_simd/src/range.rs create mode 100644 crates/http_headers_simd/src/scalar.rs create mode 100644 crates/http_headers_simd/src/tracking.rs create mode 100644 crates/http_headers_simd/src/uri.rs create mode 100644 crates/http_headers_simd/src/x86.rs create mode 100644 crates/http_headers_simd/tests/__fuzz__/campaign.toml create mode 100644 crates/http_headers_simd/tests/__fuzz__/public_scanners_match_scalar_oracles/crashes/.gitkeep create mode 100644 crates/http_headers_simd/tests/bolero_fuzz.rs diff --git a/.anvil.lock b/.anvil.lock index 644a0e3f2..d8fbcf66e 100644 --- a/.anvil.lock +++ b/.anvil.lock @@ -446,6 +446,16 @@ host = "crates/http_extensions/Cargo.toml" id = "anvil-lints" checksum = "sha256:2dd7c0f21339fd17092b8dedfe924aa86732c3520baab84f914c2d8f4103ac40" +[[region]] +host = "crates/http_headers/Cargo.toml" +id = "anvil-lints" +checksum = "sha256:2dd7c0f21339fd17092b8dedfe924aa86732c3520baab84f914c2d8f4103ac40" + +[[region]] +host = "crates/http_headers_simd/Cargo.toml" +id = "anvil-lints" +checksum = "sha256:2dd7c0f21339fd17092b8dedfe924aa86732c3520baab84f914c2d8f4103ac40" + [[region]] host = "crates/http_path_template/Cargo.toml" id = "anvil-lints" diff --git a/.cargo/mutants.toml b/.cargo/mutants.toml index 9d0bf3faa..17389807c 100644 --- a/.cargo/mutants.toml +++ b/.cargo/mutants.toml @@ -5,6 +5,8 @@ examine_globs = ["crates/**"] exclude_globs = [ # Fixture scaffolding for the `fetch_winhttp` tests, benchmarks and examples. "crates/fetch_winhttp_impl/src/testing/**", + "crates/http_headers/**", + "crates/http_headers_simd/**", "crates/observed_testing/**", "crates/rest_over_grpc_examples/**", "crates/rest_over_grpc_tests/**", diff --git a/.github/workflows/repository-checks.yml b/.github/workflows/repository-checks.yml index 771505518..15fde2f22 100644 --- a/.github/workflows/repository-checks.yml +++ b/.github/workflows/repository-checks.yml @@ -40,17 +40,62 @@ jobs: shell: pwsh run: just test-scripts + simd-no-std-arm: + name: SIMD no_std contract (native AArch64) + runs-on: ubuntu-24.04-arm + steps: + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: ./.github/actions/anvil-setup + with: + group: none + - name: Install Rust + run: just anvil-toolchain-stable-install + - name: Test isolated AArch64 no_std configuration + run: just test-http-headers-simd-no-std-arm + + simd-no-std-x86: + name: "SIMD no_std contract (baseline ${{ matrix.target }})" + strategy: + fail-fast: false + matrix: + target: [x86_64-unknown-linux-gnu, i686-unknown-linux-gnu, i586-unknown-linux-gnu] + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@3d3c42e5aac5ba805825da76410c181273ba90b1 # v7.0.1 + - uses: ./.github/actions/anvil-setup + with: + group: none + - name: Install 32-bit linker and C runtime + if: matrix.target != 'x86_64-unknown-linux-gnu' + run: | + sudo apt-get update --quiet + sudo apt-get install --yes --quiet gcc-multilib + - name: Install Rust target + run: just setup-http-headers-simd-no-std-x86 --target "${{ matrix.target }}" + - name: Test isolated baseline no_std configuration + run: just test-http-headers-simd-no-std-x86 --target "${{ matrix.target }}" + required-repository-checks: name: Required repository checks if: always() - needs: [release-script-tests] + needs: [release-script-tests, simd-no-std-arm, simd-no-std-x86] runs-on: ubuntu-latest steps: - name: Verify repository checks env: RELEASE_SCRIPT_TESTS: ${{ needs.release-script-tests.result }} + SIMD_NO_STD_ARM: ${{ needs.simd-no-std-arm.result }} + SIMD_NO_STD_X86: ${{ needs.simd-no-std-x86.result }} run: | if [[ "$RELEASE_SCRIPT_TESTS" != "success" ]]; then echo "::error::release-script-tests concluded $RELEASE_SCRIPT_TESTS" exit 1 fi + if [[ "$SIMD_NO_STD_ARM" != "success" ]]; then + echo "::error::simd-no-std-arm concluded $SIMD_NO_STD_ARM" + exit 1 + fi + if [[ "$SIMD_NO_STD_X86" != "success" ]]; then + echo "::error::simd-no-std-x86 concluded $SIMD_NO_STD_X86" + exit 1 + fi diff --git a/.spelling b/.spelling index 90ea592d8..09e1b4e1c 100644 --- a/.spelling +++ b/.spelling @@ -72,6 +72,7 @@ DevOps DotNet Dyn Enum +EPYC Extended FFI FFI-compatible @@ -204,6 +205,54 @@ Win32 Xamarin ZST ZSTs +AArch64 +CORS +DQUOTE +Fetch +Fetch's +GUID +HSTS +HTAB +HTTP's +IDNA +LF +OWS +SSE4 +SSSE3 +codings +cryptographically +formatters +indexable +intrinsics +lowercased +lowercases +lowercasing +movemasks +nonces +preload +reparse +revalidate +revalidated +revalidating +splitter +subdomains +subprotocol +subprotocols +subsecond +subtag +subtags +subtype +token68 +unaccelerated +unescaping +unmarks +userinfo +username +validator +validators +vectorizes +zeroized +zeroizes _arc _rc accessor @@ -922,3 +971,5 @@ JIT CSV subcommands POSIX +preflight +predecoded diff --git a/CHANGELOG.md b/CHANGELOG.md index 5ca20cd9b..4288bb6af 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -27,6 +27,8 @@ Please see each crate's change log below: - [`fundle_macros_impl`](./crates/fundle_macros_impl/CHANGELOG.md) - [`http_compression`](./crates/http_compression/CHANGELOG.md) - [`http_extensions`](./crates/http_extensions/CHANGELOG.md) +- [`http_headers`](./crates/http_headers/CHANGELOG.md) +- [`http_headers_simd`](./crates/http_headers_simd/CHANGELOG.md) - [`http_path_template`](./crates/http_path_template/CHANGELOG.md) - [`internity`](./crates/internity/CHANGELOG.md) - [`layered`](./crates/layered/CHANGELOG.md) diff --git a/Cargo.lock b/Cargo.lock index 0ca36ed3f..1e105f48f 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -546,6 +546,8 @@ dependencies = [ "percent-encoding", "pin-project-lite", "serde_core", + "serde_json", + "serde_path_to_error", "sync_wrapper", "tokio", "tower", @@ -725,6 +727,15 @@ version = "2.13.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b588b76d00fde79687d7646a9b5bdf3cc0f655e0bbd080335a95d7e96f3587da" +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + [[package]] name = "blocking" version = "1.7.0" @@ -1210,6 +1221,19 @@ dependencies = [ "static_assertions", ] +[[package]] +name = "compact_str" +version = "0.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "79fcda08c33bb58b97008b2cdada6622500e949e060f5913361763121abd2416" +dependencies = [ + "castaway", + "cfg-if", + "itoa", + "static_assertions", + "zmij", +] + [[package]] name = "compressors" version = "0.1.1" @@ -1418,6 +1442,16 @@ version = "0.2.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "460fbee9c2c2f33933d720630a6a0bac33ba7053db5344fac858d4b8952d77d5" +[[package]] +name = "crypto-common" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +dependencies = [ + "generic-array", + "typenum", +] + [[package]] name = "ctor" version = "1.0.13" @@ -1638,6 +1672,16 @@ dependencies = [ "unicode-xid", ] +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", +] + [[package]] name = "displaydoc" version = "0.2.7" @@ -2010,6 +2054,15 @@ dependencies = [ "zlib-rs", ] +[[package]] +name = "fluent-uri" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "17c704e9dbe1ddd863da1e6ff3567795087b1eb201ce80d8fa81162e1516500d" +dependencies = [ + "bitflags 1.3.2", +] + [[package]] name = "fnv" version = "1.0.7" @@ -2227,6 +2280,16 @@ dependencies = [ "windows-result", ] +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + [[package]] name = "getrandom" version = "0.2.17" @@ -2441,6 +2504,30 @@ dependencies = [ "foldhash 0.2.0", ] +[[package]] +name = "headers" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b3314d5adb5d94bcdf56771f2e50dbbc80bb4bdf88967526706205ac9eff24eb" +dependencies = [ + "base64 0.22.1", + "bytes", + "headers-core", + "http", + "httpdate", + "mime", + "sha1", +] + +[[package]] +name = "headers-core" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "54b4a22553d4242c49fddb9ba998a99962b5cc6f22cb5a3482bec22522403ce4" +dependencies = [ + "http", +] + [[package]] name = "heapless" version = "0.9.3" @@ -2549,6 +2636,44 @@ dependencies = [ "uuid", ] +[[package]] +name = "http_headers" +version = "0.1.0" +dependencies = [ + "axum", + "base64 0.23.1", + "bolero", + "bytes", + "compact_str 0.10.0", + "criterion", + "fluent-uri", + "headers", + "http", + "http_headers_simd", + "httpdate", + "idna", + "itoa", + "metabench", + "mutants", + "pastey", + "serde", + "serde_json", + "sha1", + "smallvec", + "tokio", + "tower", + "zeroize", +] + +[[package]] +name = "http_headers_simd" +version = "0.1.0" +dependencies = [ + "bolero", + "criterion", + "metabench", +] + [[package]] name = "http_path_template" version = "0.2.1" @@ -4667,7 +4792,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "cbb175c433c8e28a809d1f5773a2ae96e68c0ce40db865cbab1020bf33ae479c" dependencies = [ "bitflags 2.13.1", - "compact_str", + "compact_str 0.9.1", "hashbrown 0.17.1", "itertools 0.14.0", "kasuari", @@ -5316,6 +5441,17 @@ dependencies = [ "zmij", ] +[[package]] +name = "serde_path_to_error" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10a9ff822e371bb5403e391ecd83e182e0e77ba7f6fe0160b795797109d1b457" +dependencies = [ + "itoa", + "serde", + "serde_core", +] + [[package]] name = "serde_spanned" version = "1.1.1" @@ -5359,6 +5495,17 @@ dependencies = [ "syn 3.0.5", ] +[[package]] +name = "sha1" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a978451301f4db1d02937a4ab3ccce137717b81826e79b7d49ffe3244a13c3b8" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest", +] + [[package]] name = "sharded-slab" version = "0.1.7" @@ -6224,6 +6371,12 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bc7d623258602320d5c55d1bc22793b57daff0ec7efc270ea7d55ce1d5f5471c" +[[package]] +name = "typenum" +version = "1.20.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" + [[package]] name = "typespec" version = "1.1.0" diff --git a/Cargo.toml b/Cargo.toml index 14d88238a..9fa29fa1e 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -33,8 +33,11 @@ homepage = "https://github.com/microsoft/oxidizer" # # The `docs/**/*.md` glob matches only Markdown, so compile-time `include_str!` # doc fragments are packaged while binary diagram assets beside them are not. +# Bolero corpus and crash files are development artifacts regardless of +# whether their targets are integration tests or private source-level tests. include = [ "/src/**", + "!/src/__fuzz__/**", "/build.rs", "/tests/**", "!/tests/__fuzz__/**", @@ -100,6 +103,7 @@ chumsky = { version = "0.13.0", default-features = false } clap = { version = "4.6.4", default-features = false } # The latest command-group release still uses nix 0.27 on Unix; accept the duplicate until upstream updates. command-group = { version = "5.0.1", default-features = false } +compact_str = { version = "0.10.0", default-features = false } compressors = { path = "crates/compressors", default-features = false, version = "0.1.1" } const-hex = { version = "1.15.0", default-features = false } criterion = { version = "0.8.2", default-features = false } @@ -125,6 +129,7 @@ fetch_tls = { path = "crates/fetch_tls", default-features = false, version = "0. fetch_winhttp = { path = "crates/fetch_winhttp", default-features = false, version = "0.2.1" } fetch_winhttp_impl = { path = "crates/fetch_winhttp_impl", default-features = false, version = "0.2.1" } flate2 = { version = "1.1.10", default-features = false } +fluent-uri = { version = "0.1.4", default-features = false } foldhash = { version = "0.2.0", default-features = false } fundle = { path = "crates/fundle", default-features = false, version = "0.4.0" } fundle_macros = { path = "crates/fundle_macros", default-features = false, version = "=0.4.0" } @@ -140,17 +145,23 @@ gungraun-summary = { version = "=6.0.0", default-features = false } h3 = { version = "0.0.8", default-features = false } h3-quinn = { version = "0.0.10", default-features = false } hashbrown = { version = "0.17.0", default-features = false } +# Pinned exactly because the differential benchmarks use this release as their baseline. +headers = { version = "=0.4.1", default-features = false } heck = { version = "0.5.0", default-features = false } http = { version = "1.4.1", default-features = false, features = ["std"] } http-body = { version = "1.0.1", default-features = false } http-body-util = { version = "0.1.3", default-features = false } http_compression = { path = "crates/http_compression", default-features = false, version = "0.1.1" } http_extensions = { path = "crates/http_extensions", default-features = false, version = "0.11.1" } +http_headers = { path = "crates/http_headers", default-features = false, version = "0.1.0" } +http_headers_simd = { path = "crates/http_headers_simd", default-features = false, version = "=0.1.0" } http_path_template = { path = "crates/http_path_template", default-features = false, version = "0.2.1" } +httpdate = { version = "1.0.3", default-features = false } hyper = { version = "1.10.1", default-features = false } hyper-rustls = { version = "0.27.9", default-features = false } hyper-tls = { version = "0.6.0", default-features = false } hyper-util = { version = "0.1.20", default-features = false } +idna = { version = "1.1.0", default-features = false } infinity_pool = { version = "0.8.1", default-features = false } insta = { version = "1.44.1", default-features = false } internity = { path = "crates/internity", default-features = false, version = "0.2.1" } @@ -196,6 +207,7 @@ opentelemetry-stdout = { version = "0.32.0", default-features = false } opentelemetry_sdk = { version = "0.32.0", default-features = false } opool = { version = "0.2.0", default-features = false } parking_lot = { version = "0.12.5", default-features = false } +paste = { package = "pastey", version = "0.2.3", default-features = false } path-tree = { version = "0.8.3", default-features = false } pbjson = { version = "0.9.0", default-features = false } pbjson-build = { version = "0.9.0", default-features = false } @@ -250,6 +262,7 @@ serde_html_form = { version = "0.4.1", default-features = false } serde_json = { version = "1.0.145", default-features = false } serde_urlencoded = { version = "0.7.1", default-features = false } serial_test = { version = "4.0.1", default-features = false } +sha1 = { version = "0.10.7", default-features = false } sharded-slab = { version = "0.1.7", default-features = false } slab = { version = "0.4.12", default-features = false } slotmap = { version = "1.1.1", default-features = false } @@ -301,6 +314,7 @@ windows-sys = { version = "0.61.2", default-features = false } wiremock = { version = "0.6.5", default-features = false } xxhash-rust = { version = "0.8.15", default-features = false } zerocopy = { version = "0.8.26", default-features = false } +zeroize = { version = "1.9.0", default-features = false } zstd-safe = { version = "8.0.0", default-features = false } # >>> anvil-managed: anvil-workspace-lints @@ -380,6 +394,7 @@ clippy.multiple_unsafe_ops_per_block = "warn" clippy.redundant_type_annotations = "warn" clippy.renamed_function_params = "warn" clippy.semicolon_outside_block = "warn" +clippy.too_long_first_doc_paragraph = "warn" clippy.undocumented_unsafe_blocks = "warn" clippy.unnecessary_safety_comment = "warn" clippy.unnecessary_safety_doc = "warn" diff --git a/README.md b/README.md index 46037e98b..084161d33 100644 --- a/README.md +++ b/README.md @@ -46,6 +46,7 @@ These are the primary crates built out of this repo: - [`fundle`](./crates/fundle/README.md) - Compile-time safe dependency injection for Rust. - [`http_compression`](./crates/http_compression/README.md) - HTTP request and response body compression and decompression. - [`http_extensions`](./crates/http_extensions/README.md) - Shared HTTP types and extension traits for clients and servers. +- [`http_headers`](./crates/http_headers/README.md) - Fast, ergonomic typed HTTP headers with borrowed views. - [`http_path_template`](./crates/http_path_template/README.md) - Parser for the google.api.http path-template grammar. - [`internity`](./crates/internity/README.md) - Blazingly fast string interning with compact handles, compact storage, and concurrent fill support. - [`layered`](./crates/layered/README.md) - A foundational service abstraction for building composable, middleware-driven systems. diff --git a/crates/http_headers/CHANGELOG.md b/crates/http_headers/CHANGELOG.md new file mode 100644 index 000000000..6a8522106 --- /dev/null +++ b/crates/http_headers/CHANGELOG.md @@ -0,0 +1,5 @@ +# Changelog + +## [0.1.0] + +- Initial integration into the Oxidizer workspace. diff --git a/crates/http_headers/Cargo.toml b/crates/http_headers/Cargo.toml new file mode 100644 index 000000000..6ae7ba327 --- /dev/null +++ b/crates/http_headers/Cargo.toml @@ -0,0 +1,183 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +[package] +name = "http_headers" +description = "Fast, ergonomic typed HTTP headers with borrowed views." +version = "0.1.0" +readme = "README.md" +edition.workspace = true +rust-version.workspace = true +authors.workspace = true +license.workspace = true +homepage.workspace = true +include.workspace = true +repository = "https://github.com/microsoft/oxidizer/tree/main/crates/http_headers" +documentation = "https://docs.rs/http_headers" +keywords = ["http", "headers", "zero-copy", "simd"] +categories = ["web-programming"] +autobenches = false + +[package.metadata.cargo_check_external_types] +allowed_external_types = ["bytes::bytes::Bytes", "http::*", "serde_core::de::*", "serde_core::ser::*"] + +[package.metadata.docs.rs] +all-features = true + +[features] +default = ["headers-all"] +benchmarking = ["http_headers_simd/benchmarking"] +headers-all = [ + "headers-authorization", + "headers-cache-control", + "headers-conditional", + "headers-content-length", + "headers-content-type", + "headers-cors", + "headers-etag", + "headers-location", + "headers-negotiation", + "headers-range", + "headers-security", + "headers-set-cookie", + "headers-user-agent", + "headers-websocket", +] +headers-authorization = ["dep:base64", "dep:zeroize"] +headers-cache-control = ["dep:compact_str"] +headers-conditional = ["dep:httpdate", "headers-etag"] +headers-content-length = [] +headers-content-type = [] +headers-cors = [] +headers-etag = [] +headers-location = ["dep:fluent-uri"] +headers-negotiation = ["dep:idna"] +headers-range = [] +headers-security = ["dep:compact_str"] +headers-set-cookie = [] +headers-user-agent = [] +headers-websocket = ["dep:base64", "dep:sha1"] +# Optional adapter for the external `http` crate. The adapter may retain +# `HeaderValue` storage internally without changing typed-header semantics. +http = ["dep:http"] +serde = ["dep:serde"] + +[dependencies] +base64 = { workspace = true, optional = true, features = ["std"] } +bytes = { workspace = true, features = ["std"] } +compact_str = { workspace = true, optional = true, features = ["std"] } +fluent-uri = { workspace = true, optional = true, features = ["ipv_future"] } +http = { workspace = true, optional = true } +http_headers_simd = { workspace = true, features = ["std"] } +httpdate = { workspace = true, optional = true } +idna = { workspace = true, optional = true, features = ["compiled_data", "std"] } +itoa.workspace = true +serde = { workspace = true, optional = true, features = ["derive", "std"] } +sha1 = { workspace = true, optional = true, features = ["std"] } +smallvec = { workspace = true, features = ["const_new"] } +zeroize = { workspace = true, optional = true, features = ["alloc"] } + +[dev-dependencies] +axum = { workspace = true, features = ["http1", "json", "tokio"] } +base64 = { workspace = true, features = ["std"] } +bolero = { workspace = true, features = ["std"] } +compact_str = { workspace = true, features = ["std"] } +criterion = { workspace = true } +fluent-uri = { workspace = true, features = ["ipv_future"] } +headers = { workspace = true } +http = { workspace = true } +httpdate = { workspace = true } +idna = { workspace = true, features = ["compiled_data", "std"] } +metabench = { workspace = true } +mutants = { workspace = true } +paste = { workspace = true } +serde = { workspace = true, features = ["derive", "std"] } +serde_json = { workspace = true, features = ["std"] } +sha1 = { workspace = true, features = ["std"] } +tokio = { workspace = true, features = ["macros", "net", "rt-multi-thread"] } +tower = { workspace = true, features = ["util"] } +zeroize = { workspace = true, features = ["alloc"] } + +[[example]] +name = "axum" +required-features = ["headers-cache-control", "headers-user-agent", "http"] + +[[bench]] +name = "http_headers_micro" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_storage" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_per_header" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_name_recognition" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_policy" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_typed_tokens" +harness = false +required-features = ["benchmarking", "http", "headers-negotiation"] + +[[bench]] +name = "http_headers_location_semantics" +harness = false +required-features = ["benchmarking", "http", "headers-location"] + +[[bench]] +name = "http_headers_authority_semantics" +harness = false +required-features = ["benchmarking", "http", "headers-negotiation", "headers-cors"] + +[[bench]] +name = "http_headers_negotiation_semantics" +harness = false +required-features = ["benchmarking", "http", "headers-negotiation"] + +[[bench]] +name = "http_headers_protocol" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_websocket_shapes" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_negotiation_shapes" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_auth_cors_shapes" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_conditional_range_shapes" +harness = false +required-features = ["benchmarking", "http"] + +[[bench]] +name = "http_headers_policy_shapes" +harness = false +required-features = ["benchmarking", "http"] + +# >>> anvil-managed: anvil-lints +[lints] +workspace = true +# <<< anvil-managed: anvil-lints diff --git a/crates/http_headers/README.md b/crates/http_headers/README.md new file mode 100644 index 000000000..116ac1cfe --- /dev/null +++ b/crates/http_headers/README.md @@ -0,0 +1,465 @@ +
+ Http Headers Logo + +# Http Headers + +[![crate.io](https://img.shields.io/crates/v/http_headers.svg)](https://crates.io/crates/http_headers) +[![docs.rs](https://docs.rs/http_headers/badge.svg)](https://docs.rs/http_headers) +[![MSRV](https://img.shields.io/crates/msrv/http_headers)](https://crates.io/crates/http_headers) +[![CI](https://github.com/microsoft/oxidizer/actions/workflows/anvil-pr.yml/badge.svg)](https://github.com/microsoft/oxidizer/actions/workflows/anvil-pr.yml) +[![Coverage](https://codecov.io/gh/microsoft/oxidizer/graph/badge.svg?token=FCUG0EL5TI)](https://codecov.io/gh/microsoft/oxidizer) +[![License](https://img.shields.io/badge/license-MIT-blue.svg)](https://github.com/microsoft/oxidizer/blob/main/LICENSE) +This crate was developed as part of the Oxidizer project + +
+ +Efficient and robust HTTP header parsing and creation. + +This crate provides: + +* Highly optimized parsing of incoming HTTP headers which produce owned or borrowed + strongly-typed Rust structs. These parsers insulate your code from badly formed + headers. + +* Highly optimized production of headers, ensuring the headers are well-formed. + +Header parsing and production are abstracted over their source and destination. +The optional `http` feature integrates with +[`HeaderMap`][__link0] +plus generic [`Request`][__link1] +and [`Response`][__link2] +values from the [`http`][__link3] crate. + +## Parsing headers + +Headers are parsed from an implementation of the [`source::FieldSource`][__link4] trait. The `http` crate feature +implements this trait for [`HeaderMap`][__link5], +[`Request`][__link6], and +[`Response`][__link7]. +Once you have a source, you can choose to parse into borrowed views or owned structs. +Prefer borrowed views when the decoded value does not need to outlive the +source as they are generally faster. Use owned structs when the parsed header +data needs to be retained (such as in a cache). + +Source and sink operations use static field-name descriptors, including +custom names stored in a `static LazyLock`. Locally constructed +runtime names are supported by [`FieldName`][__link8] for validation and conversion, +but dynamic lookup and mutation must use the container’s native API. + +[`Field::view`][__link9] returns a header’s +borrowed `*View` type, whose lifetime is tied to the source. + +```rust +use http::HeaderMap; +use http_headers::Field; +use http_headers::headers::{ContentType, UserAgent}; + +// create a HeaderMap to show how to read from it +let mut headers = HeaderMap::new(); +headers.insert( + http::header::USER_AGENT, + http::HeaderValue::from_static("example-client/1.0"), +); +headers.insert( + http::header::CONTENT_TYPE, + http::HeaderValue::from_static("application/json; charset=utf-8"), +); + +if let Some(agent) = UserAgent::view(&headers)? { + assert_eq!(agent.as_str()?, "example-client/1.0"); +} + +if let Some(content_type) = ContentType::view(&headers)? { + assert_eq!(content_type.type_()?, "application"); + assert_eq!(content_type.subtype()?, "json"); + assert_eq!( + content_type.parameter("charset")?, + Some(b"utf-8".as_slice()) + ); +} +``` + +Prefer `view` unless the decoded value must outlive the source. Use +[`Field::owned`][__link10] when you need to retain, move, or independently +store the result: + +```rust +use http::HeaderMap; +use http_headers::Field; +use http_headers::headers::UserAgent; + +let mut headers = HeaderMap::new(); +headers.insert( + http::header::USER_AGENT, + http::HeaderValue::from_static("example-client/1.0"), +); + +let owned = UserAgent::owned(&headers)?.expect("User-Agent is present"); +drop(headers); +assert_eq!(owned.as_bytes(), b"example-client/1.0"); +``` + +Both methods return `Ok(None)` when the header is absent and `Err` when a +present value is malformed. + +### Reading validated members + +Structured headers expose semantic values as well as their original wire +representation. `AllowOwned::methods()` yields case-sensitive method tokens; +`VaryOwned::entries()` distinguishes wildcard members from case-insensitive +field names. These borrowed member reads do not allocate. + +```rust +use http_headers::headers::{AllowOwned, MethodView, VaryOwned}; + +let allow = AllowOwned::try_from("GET, HEAD, CUSTOM")?; +assert!(allow.methods().any(|method| method == MethodView::GET)); +assert!(!allow.methods().any(|method| method == MethodView::POST)); + +let vary = VaryOwned::try_from("Accept-Encoding, X-Tenant")?; +assert!(!vary.contains_wildcard()); +assert!(vary.entries().any(|entry| { + entry + .field_name() + .is_some_and(|name| name.eq_ignore_ascii_case("x-tenant")) +})); +``` + +The Accept family exposes typed ranges, parameters, and exact quality +weights through `entries()`. Location exposes URI-reference components +through `uri_reference()`. Host and Allow-Origin retain parsed authority +components. Semantic access does not sort lists or replace the original +field lines used for forwarding. + +## Producing headers + +You produce headers by populating an implementation of the [`sink::FieldSink`][__link11] trait. The +`http` cargo feature implements this trait for +[`HeaderMap`][__link12], +[`Request`][__link13], and +[`Response`][__link14]. + +Enabling a header-family feature exposes `sink::FieldSinkExt`, whose fluent +methods work for any sink. Core-only builds use [`sink::FieldSink`][__link15] directly. + +```rust +use std::time::Duration; + +use http::HeaderMap; +use http_headers::headers::{CacheControl, ContentType}; +use http_headers::sink::FieldSinkExt; + +let mut headers = HeaderMap::new(); +headers + .set_content_type(ContentType::json())? + .set_content_length(1_024)? + .set_cache_control(CacheControl::public().max_age(Duration::from_secs(60)))?; +``` + +An owned value can insert itself when it has already been constructed: + +```rust +use http::HeaderMap; +use http_headers::headers::LocationOwned; + +let mut headers = HeaderMap::new(); +LocationOwned::try_from("/next")?.insert_into(&mut headers)?; +``` + +A borrowed view can also be forwarded directly to another sink: + +```rust +use http::HeaderMap; +use http_headers::Field; +use http_headers::headers::UserAgent; + +let mut incoming = HeaderMap::new(); +incoming.insert( + http::header::USER_AGENT, + http::HeaderValue::from_static("example-client/1.0"), +); + +let mut outgoing = HeaderMap::new(); +if let Some(agent) = UserAgent::view(&incoming)? { + agent.insert_into(&mut outgoing)?; +} +``` + +[`Field::insert`][__link16] is the generic alternative when the descriptor type is +already known. It replaces all existing field lines for that header; +[`Field::remove`][__link17] removes them instead. + +```rust +use http::HeaderMap; +use http_headers::Field; +use http_headers::headers::{UserAgent, UserAgentOwned}; + +let mut headers = HeaderMap::new(); +UserAgent::insert( + &mut headers, + UserAgentOwned::try_from_static("example-client/1.0")?, +)?; +UserAgent::remove(&mut headers); +``` + +Repeated field lines remain separate. In particular, `Set-Cookie` values are +never comma-joined: + +```rust +use http::HeaderMap; +use http_headers::Field; +use http_headers::headers::{SetCookie, SetCookieOwned}; + +let mut cookies = SetCookieOwned::new(); +cookies.push_str("session=abc; Path=/; HttpOnly")?; +cookies.push_str("theme=dark; Path=/")?; + +let mut headers = HeaderMap::new(); +SetCookie::insert(&mut headers, cookies)?; +assert_eq!(headers.get_all(http::header::SET_COOKIE).iter().count(), 2); +``` + +## Serialization + +The `serde` cargo feature implements Serde serialization and deserialization for every owned header +struct, [`FieldName`][__link18], [`FieldValue`][__link19], and [`sink::EncodedValues`][__link20]. +Headers serialize as an ordered sequence of physical field values and +deserialize through relaxed validation, which includes strict syntax and +the documented interoperability deviations. This preserves round trips for +every owned value produced by the public API: + +```rust +use http_headers::headers::UserAgentOwned; + +let header = UserAgentOwned::try_from("example-client/1.0")?; +let json = serde_json::to_string(&header)?; +let decoded: UserAgentOwned = serde_json::from_str(&json)?; +assert_eq!(decoded.as_bytes(), header.as_bytes()); +``` + +Repeated lines retain their boundaries: + +```rust +use http_headers::headers::SetCookieOwned; + +let mut cookies = SetCookieOwned::new(); +cookies.push_str("session=abc")?; +cookies.push_str("theme=dark")?; +let json = serde_json::to_string(&cookies)?; +let decoded: SetCookieOwned = serde_json::from_str(&json)?; +assert_eq!(decoded.len(), 2); +``` + +Serialization is not redaction. Sensitive values include their original +bytes and an explicit sensitivity marker, so serialized data must be +protected like the header value itself: + +```rust +use http_headers::{FieldSensitivity, FieldValue}; + +let secret = + FieldValue::from_static("credential").with_sensitivity(FieldSensitivity::Sensitive); +let json = serde_json::to_string(&secret)?; +let decoded: FieldValue = serde_json::from_str(&json)?; +assert_eq!(decoded.as_bytes(), b"credential"); +assert!(decoded.is_sensitive()); +``` + +## Strict and relaxed reads + +[`Field::view`][__link21] and [`Field::owned`][__link22] use strict syntax. Applications that +must accept specific common deviations can request [`DecodeMode::Relaxed`][__link23] +through [`Field::view_with`][__link24] or [`Field::owned_with`][__link25]: + +```rust +use http_headers::headers::AcceptEncoding; +use http_headers::{DecodeMode, Field}; + +let mut headers = http::HeaderMap::new(); +headers.insert( + http::header::ACCEPT_ENCODING, + http::HeaderValue::from_static("gzip; q = .5"), +); + +assert!(AcceptEncoding::view(&headers).is_err()); +assert!(AcceptEncoding::view_with(&headers, DecodeMode::Relaxed)?.is_some()); +``` + +Relaxed mode is not a general validation bypass. Each header documents the +additional forms it accepts, and the original field bytes are preserved. + +## Sensitive values + +Authorization, `Location`, and cookie values are marked sensitive so their +`Debug` representations and compatible sinks do not reveal their contents. +Basic authentication can be read through a borrowed view while reusing +caller-owned decode storage: + +```rust +use http::HeaderMap; +use http_headers::Field; +use http_headers::headers::{Authorization, Basic, BasicCredentials}; + +let mut headers = HeaderMap::new(); +headers.insert( + http::header::AUTHORIZATION, + http::HeaderValue::from_static("Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ=="), +); + +let authorization = Authorization::::view(&headers)?.expect("Authorization is present"); +let mut credentials = BasicCredentials::new(); +let decoded = authorization.extract(&mut credentials)?; +assert_eq!(decoded.username(), b"Aladdin"); +credentials.clear(); +``` + +`BasicCredentials` zeroizes decoded bytes when cleared, reused, or dropped. + +## Defining a custom single-value header + +Implement [`SingleValueField`][__link26] when a custom header is represented by exactly +one field line. The crate then supplies its [`Field`][__link27] implementation, +including borrowed and owned reads, singleton cardinality checks, insertion, +and removal. + +```rust +use std::sync::LazyLock; + +use http_headers::{ + DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, SingleValueField, +}; + +static REQUEST_ID: LazyLock = + LazyLock::new(|| FieldName::from_static("x-request-id")); + +struct RequestId; + +#[derive(Clone, Debug, Eq, PartialEq)] +struct RequestIdOwned(FieldValue); + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +struct RequestIdView<'a>(FieldValueRef<'a>); + +fn is_token(bytes: &[u8]) -> bool { + !bytes.is_empty() + && bytes + .iter() + .all(|byte| byte.is_ascii_alphanumeric() || b"!#$%&'*+-.^_`|~".contains(byte)) +} + +impl SingleValueField for RequestId { + type View<'a> = RequestIdView<'a>; + type Owned = RequestIdOwned; + + fn name() -> &'static FieldName { + &REQUEST_ID + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + if is_token(value.as_bytes()) { + Ok(RequestIdView(value)) + } else { + Err(DecodeError::new(&REQUEST_ID, DecodeErrorKind::InvalidToken)) + } + } + + fn decode_owned(value: FieldValue) -> Result { + if is_token(value.as_bytes()) { + Ok(RequestIdOwned(value)) + } else { + Err(DecodeError::new(&REQUEST_ID, DecodeErrorKind::InvalidToken)) + } + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.0 + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.0 + } +} +``` + +## Performance + +[`docs/PERF.md`][__link28] +records comparative typed-decode-and-read measurements against `headers 0.4.1`. +Results vary by header, ownership mode, and hardware, and include both faster +and slower cases. The table does not measure comparative header production or +end-to-end request processing. Borrowed reads generally avoid allocations. + +## Cargo features + +* `headers-all` (enabled by default): all built-in typed header families. +* `headers-authorization`, `headers-cache-control`, `headers-conditional`, + `headers-content-length`, `headers-content-type`, `headers-cors`, + `headers-etag`, `headers-location`, `headers-negotiation`, `headers-range`, + `headers-security`, `headers-set-cookie`, `headers-user-agent`, and + `headers-websocket`: individual built-in header families. +* `http`: optional adapter for `http::HeaderMap` and the `http` crate’s name, + value, and method types. +* `serde`: serialization and deserialization for owned headers, + [`FieldName`][__link29], [`FieldValue`][__link30], and [`sink::EncodedValues`][__link31]. + +Disable default features to use only the core source, sink, name, and value +APIs, then enable only the header families an application needs. + +## What about trailers? + +Although this crate is named `http_headers`, it fully supports trailers as well. +The crate doesn’t currently expose any trailer-specific structs however, so you +would need to define those structs and implement the parsers yourself as implementations +of the traits in this crate. + +## Alternate crates + +This crate is an alternative to the popular [`headers`][__link32] crate. +`http_headers` has the following benefits: + +* Faster decoding for some headers in the measured configurations +* Supports more headers +* Performs more robust validation to avoid downstream surprises +* Supports explicit relaxed parsing options to support common malformed headers +* Supports serde + + +
+ +This crate was developed as part of The Oxidizer Project. Browse this crate's source code. + + + [__cargo_doc2readme_dependencies_info]: ggGmYW0CYXZlMC43LjNhdIQborR2_k_xJd4bTcf2krrNPIcbP72Pw1UdRjkbim_eMDe2BBthYvRhcoQbfVFs3NqFhWgbBn84idxhrs4bC_HSVxhRRoYbYww44PeKrh5hZIGCbGh0dHBfaGVhZGVyc2UwLjEuMA + [__link0]: https://docs.rs/http/latest/http/header/struct.HeaderMap.html + [__link1]: https://docs.rs/http/latest/http/request/struct.Request.html + [__link10]: https://docs.rs/http_headers/0.1.0/http_headers/?search=Field::owned + [__link11]: https://docs.rs/http_headers/0.1.0/http_headers/?search=sink::FieldSink + [__link12]: https://docs.rs/http/latest/http/header/struct.HeaderMap.html + [__link13]: https://docs.rs/http/latest/http/request/struct.Request.html + [__link14]: https://docs.rs/http/latest/http/response/struct.Response.html + [__link15]: https://docs.rs/http_headers/0.1.0/http_headers/?search=sink::FieldSink + [__link16]: https://docs.rs/http_headers/0.1.0/http_headers/?search=Field::insert + [__link17]: https://docs.rs/http_headers/0.1.0/http_headers/?search=Field::remove + [__link18]: https://docs.rs/http_headers/0.1.0/http_headers/?search=FieldName + [__link19]: https://docs.rs/http_headers/0.1.0/http_headers/?search=FieldValue + [__link2]: https://docs.rs/http/latest/http/response/struct.Response.html + [__link20]: https://docs.rs/http_headers/0.1.0/http_headers/?search=sink::EncodedValues + [__link21]: https://docs.rs/http_headers/0.1.0/http_headers/?search=Field::view + [__link22]: https://docs.rs/http_headers/0.1.0/http_headers/?search=Field::owned + [__link23]: https://docs.rs/http_headers/0.1.0/http_headers/?search=DecodeMode::Relaxed + [__link24]: https://docs.rs/http_headers/0.1.0/http_headers/?search=Field::view_with + [__link25]: https://docs.rs/http_headers/0.1.0/http_headers/?search=Field::owned_with + [__link26]: https://docs.rs/http_headers/0.1.0/http_headers/?search=SingleValueField + [__link27]: https://docs.rs/http_headers/0.1.0/http_headers/?search=Field + [__link28]: https://github.com/microsoft/oxidizer/blob/main/crates/http_headers/docs/PERF.md + [__link29]: https://docs.rs/http_headers/0.1.0/http_headers/?search=FieldName + [__link3]: https://crates.io/crates/http + [__link30]: https://docs.rs/http_headers/0.1.0/http_headers/?search=FieldValue + [__link31]: https://docs.rs/http_headers/0.1.0/http_headers/?search=sink::EncodedValues + [__link32]: https://crates.io/crates/headers + [__link4]: https://docs.rs/http_headers/0.1.0/http_headers/?search=source::FieldSource + [__link5]: https://docs.rs/http/latest/http/header/struct.HeaderMap.html + [__link6]: https://docs.rs/http/latest/http/request/struct.Request.html + [__link7]: https://docs.rs/http/latest/http/response/struct.Response.html + [__link8]: https://docs.rs/http_headers/0.1.0/http_headers/?search=FieldName + [__link9]: https://docs.rs/http_headers/0.1.0/http_headers/?search=Field::view diff --git a/crates/http_headers/benches/http_headers_auth_cors_shapes.rs b/crates/http_headers/benches/http_headers_auth_cors_shapes.rs new file mode 100644 index 000000000..337bd9c2d --- /dev/null +++ b/crates/http_headers/benches/http_headers_auth_cors_shapes.rs @@ -0,0 +1,488 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! HTTP-backed and raw-source decode shapes for authorization, CORS and content length. + +use http_headers::DecodeErrorKind; +use http_headers::headers::{ + AccessControlAllowCredentials, AccessControlAllowHeaders, AccessControlAllowMethods, AccessControlAllowOrigin, + AccessControlExposeHeaders, AccessControlMaxAge, AccessControlRequestHeaders, AccessControlRequestMethod, Authorization, Basic, Bearer, + ContentLength, +}; + +#[path = "http_headers_shapes_common.rs"] +mod shapes; + +use shapes::Expected; + +shapes::define_shapes!( + "http_headers_auth_cors_shapes/parse"; + (authorization_basic_short, Authorization, &["Basic dTpw"], Strict, Expected::Valid), + ( + authorization_basic_existing_fixture, + Authorization, + &["Basic YWxhZGRpbjpvcGVuc2VzYW1l"], + Strict, + Expected::Valid + ), + (authorization_basic_padded, Authorization, &["Basic dXNlcjpwYXNzd29yZA=="], Strict, Expected::Valid), + ( + authorization_basic_service, + Authorization, + &["Basic Z2F0ZXdheS1zZXJ2aWNlLWFjY291bnQ6c3ludGhldGljLXBhc3N3b3JkLXdpdGgtNjQtY2hhcmFjdGVycy0wMTIzNDU2Nzg5LWFiY2RlZmdoaWprbG1ub3A="], + Strict, + Expected::Valid + ), + ( + authorization_basic_long_username, + Authorization, + &["Basic bG9uZy1zZXJ2aWNlLWFjY291bnQtbmFtZS1mb3ItYXV0aGVudGljYXRpb24tcHJveHktdGVzdGluZzpzaG9ydA=="], + Strict, + Expected::Valid + ), + (authorization_basic_binary_password, Authorization, &["Basic dXNlcjoA/w=="], Strict, Expected::Valid), + (authorization_basic_mixed_case_spaces, Authorization, &["bAsIc dTpw"], Strict, Expected::Valid), + (authorization_basic_relaxed, Authorization, &["bAsIc dTpw"], Relaxed, Expected::Valid), + (authorization_basic_absent, Authorization, &[], Strict, Expected::Absent), + ( + authorization_basic_noncanonical_pad_bits, + Authorization, + &["Basic Oh=="], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + authorization_basic_missing_colon, + Authorization, + &["Basic bm8tY29sb24="], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + authorization_basic_repeated, + Authorization, + &["Basic dTpw", "Basic dTpw"], + Strict, + Expected::Error(DecodeErrorKind::UnexpectedMultipleValues) + ), + (authorization_bearer_short, Authorization, &["Bearer abc.def"], Strict, Expected::Valid), + (authorization_bearer_token_eight, Authorization, &["Bearer test1234"], Strict, Expected::Valid), + (authorization_bearer_token_nine, Authorization, &["Bearer test12345"], Strict, Expected::Valid), + ( + authorization_bearer_token_31, + Authorization, + &["Bearer abcdefghijklmnopqrstuvwxyz01234"], + Strict, + Expected::Valid + ), + ( + authorization_bearer_token_32, + Authorization, + &["Bearer abcdefghijklmnopqrstuvwxyz012345"], + Strict, + Expected::Valid + ), + ( + authorization_bearer_token_33, + Authorization, + &["Bearer abcdefghijklmnopqrstuvwxyz0123456"], + Strict, + Expected::Valid + ), + ( + authorization_bearer_synthetic_jwt_shape, + Authorization, + &["Bearer syntheticHeader0123456789.syntheticPayloadForGatewayServiceAccountWithScopesReadWriteAndExpiry0123456789.syntheticSignatureAbCdEfGhIjKlMnOpQrStUvWxYz0123456789_-"], + Strict, + Expected::Valid + ), + (authorization_bearer_padding, Authorization, &["Bearer YWJjZA=="], Strict, Expected::Valid), + (authorization_bearer_mixed_case_spaces, Authorization, &["bEaReR abc.def"], Strict, Expected::Valid), + (authorization_bearer_absent, Authorization, &[], Strict, Expected::Absent), + ( + authorization_bearer_interior_padding, + Authorization, + &["Bearer ab=c"], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + authorization_bearer_invalid_relaxed, + Authorization, + &["Bearer ab=c"], + Relaxed, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + authorization_bearer_repeated, + Authorization, + &["Bearer abc.def", "Bearer abc.def"], + Strict, + Expected::Error(DecodeErrorKind::UnexpectedMultipleValues) + ), + (access_control_allow_credentials_true, AccessControlAllowCredentials, &["true"], Strict, Expected::Valid), + (access_control_allow_credentials_leading_ows, AccessControlAllowCredentials, &[" true"], Strict, Expected::Valid), + (access_control_allow_credentials_trailing_ows, AccessControlAllowCredentials, &["true\t"], Strict, Expected::Valid), + (access_control_allow_credentials_framed, AccessControlAllowCredentials, &[" \ttrue\t "], Strict, Expected::Valid), + ( + access_control_allow_credentials_wide_ows, + AccessControlAllowCredentials, + &[" true "], + Strict, + Expected::Valid + ), + (access_control_allow_credentials_absent, AccessControlAllowCredentials, &[], Strict, Expected::Absent), + ( + access_control_allow_credentials_empty, + AccessControlAllowCredentials, + &[""], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + access_control_allow_credentials_case_relaxed, + AccessControlAllowCredentials, + &["TRUE"], + Relaxed, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + access_control_allow_credentials_suffix, + AccessControlAllowCredentials, + &[" truee "], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + access_control_allow_credentials_repeated, + AccessControlAllowCredentials, + &["true", "true"], + Strict, + Expected::Error(DecodeErrorKind::UnexpectedMultipleValues) + ), + ( + access_control_allow_headers_common_pair, + AccessControlAllowHeaders, + &["content-type, x-request-id"], + Strict, + Expected::Valid + ), + (access_control_allow_headers_custom, AccessControlAllowHeaders, &["x-custom-header"], Strict, Expected::Valid), + ( + access_control_allow_headers_mixed_case, + AccessControlAllowHeaders, + &["Content-Type, X-Correlation-Id"], + Strict, + Expected::Valid + ), + ( + access_control_allow_headers_large, + AccessControlAllowHeaders, + &["authorization, content-type, x-request-id, x-correlation-id, x-client-version, x-tenant-id, x-trace-id, x-idempotency-key, x-api-version, x-requested-with, traceparent, tracestate"], + Strict, + Expected::Valid + ), + ( + access_control_allow_headers_repeated, + AccessControlAllowHeaders, + &["content-type, x-request-id", "authorization", "x-custom-header, content-type"], + Strict, + Expected::Valid + ), + ( + access_control_allow_headers_sixteen_lines, + AccessControlAllowHeaders, + &["x-00", "x-01", "x-02", "x-03", "x-04", "x-05", "x-06", "x-07", "x-08", "x-09", "x-10", "x-11", "x-12", "x-13", "x-14", "x-15"], + Strict, + Expected::Valid + ), + (access_control_allow_headers_empty, AccessControlAllowHeaders, &[""], Strict, Expected::Valid), + ( + access_control_allow_headers_empty_slots, + AccessControlAllowHeaders, + &[" ,\t, content-type,, X-Trace-Id, "], + Strict, + Expected::Valid + ), + (access_control_allow_headers_wildcard, AccessControlAllowHeaders, &["*"], Strict, Expected::Valid), + ( + access_control_allow_headers_late_error, + AccessControlAllowHeaders, + &["content-type", "x-request-id", "bad:name"], + Strict, + Expected::Error(DecodeErrorKind::InvalidToken) + ), + (access_control_allow_headers_absent, AccessControlAllowHeaders, &[], Strict, Expected::Absent), + (access_control_allow_methods_pair, AccessControlAllowMethods, &["GET, POST"], Strict, Expected::Valid), + ( + access_control_allow_methods_crud, + AccessControlAllowMethods, + &["GET, HEAD, POST, PUT, PATCH, DELETE, OPTIONS"], + Strict, + Expected::Valid + ), + ( + access_control_allow_methods_extensions, + AccessControlAllowMethods, + &["PROPFIND, PROPPATCH, MKCOL, COPY, MOVE, LOCK, UNLOCK"], + Strict, + Expected::Valid + ), + ( + access_control_allow_methods_repeated, + AccessControlAllowMethods, + &["GET, X-PURGE,,", "PATCH, GET", "POST"], + Strict, + Expected::Valid + ), + (access_control_allow_methods_empty, AccessControlAllowMethods, &[""], Strict, Expected::Valid), + ( + access_control_allow_methods_whitespace, + AccessControlAllowMethods, + &["\t GET ,\tPOST, , PATCH \t"], + Strict, + Expected::Valid + ), + (access_control_allow_methods_wildcard, AccessControlAllowMethods, &["*"], Strict, Expected::Valid), + (access_control_allow_methods_absent, AccessControlAllowMethods, &[], Strict, Expected::Absent), + ( + access_control_allow_methods_late_error, + AccessControlAllowMethods, + &["GET, POST", "PATCH", "bad method"], + Strict, + Expected::Error(DecodeErrorKind::InvalidToken) + ), + (access_control_allow_methods_relaxed, AccessControlAllowMethods, &["get, X-PURGE"], Relaxed, Expected::Valid), + (access_control_allow_origin_https, AccessControlAllowOrigin, &["https://example.com"], Strict, Expected::Valid), + (access_control_allow_origin_http, AccessControlAllowOrigin, &["http://example.com"], Strict, Expected::Valid), + (access_control_allow_origin_wildcard, AccessControlAllowOrigin, &["*"], Strict, Expected::Valid), + (access_control_allow_origin_null, AccessControlAllowOrigin, &["null"], Strict, Expected::Valid), + (access_control_allow_origin_port, AccessControlAllowOrigin, &["https://api.example.com:8443"], Strict, Expected::Valid), + ( + access_control_allow_origin_long_domain, + AccessControlAllowOrigin, + &["https://service-authentication.production.westus2.customer-tenant-0123456789.internal.example.com"], + Strict, + Expected::Valid + ), + (access_control_allow_origin_ipv4, AccessControlAllowOrigin, &["http://127.0.0.1"], Strict, Expected::Valid), + (access_control_allow_origin_ipv6_port, AccessControlAllowOrigin, &["https://[2001:db8::1]:8443"], Strict, Expected::Valid), + ( + access_control_allow_origin_ipv6_uncompressed, + AccessControlAllowOrigin, + &["https://[2001:db8:1:2:3:4:5:6]"], + Strict, + Expected::Valid + ), + (access_control_allow_origin_wss, AccessControlAllowOrigin, &["wss://example.com"], Strict, Expected::Valid), + (access_control_allow_origin_whitespace, AccessControlAllowOrigin, &[" \thttps://example.com\t "], Strict, Expected::Valid), + (access_control_allow_origin_absent, AccessControlAllowOrigin, &[], Strict, Expected::Absent), + ( + access_control_allow_origin_uppercase_host, + AccessControlAllowOrigin, + &["https://Example.com"], + Relaxed, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + access_control_allow_origin_default_port, + AccessControlAllowOrigin, + &["https://example.com:443"], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + access_control_allow_origin_ipv6_leading_zero, + AccessControlAllowOrigin, + &["https://[2001:0db8::1]"], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + access_control_allow_origin_repeated, + AccessControlAllowOrigin, + &["https://example.com", "https://example.com"], + Strict, + Expected::Error(DecodeErrorKind::UnexpectedMultipleValues) + ), + (access_control_expose_headers_pair, AccessControlExposeHeaders, &["etag, x-request-id"], Strict, Expected::Valid), + ( + access_control_expose_headers_response, + AccessControlExposeHeaders, + &["content-length, content-range, etag"], + Strict, + Expected::Valid + ), + ( + access_control_expose_headers_mixed_case, + AccessControlExposeHeaders, + &["X-Request-Id, Content-Length"], + Strict, + Expected::Valid + ), + ( + access_control_expose_headers_large, + AccessControlExposeHeaders, + &["etag, content-length, content-range, x-request-id, x-correlation-id, x-ratelimit-limit, x-ratelimit-remaining, x-ratelimit-reset, retry-after, server-timing, x-api-version, x-total-count"], + Strict, + Expected::Valid + ), + ( + access_control_expose_headers_repeated, + AccessControlExposeHeaders, + &["etag, x-request-id", "X-Trace-Id", "etag, content-length"], + Strict, + Expected::Valid + ), + (access_control_expose_headers_empty, AccessControlExposeHeaders, &[""], Strict, Expected::Valid), + (access_control_expose_headers_wildcard, AccessControlExposeHeaders, &["*"], Strict, Expected::Valid), + ( + access_control_expose_headers_empty_slots, + AccessControlExposeHeaders, + &["\t, etag,, X-Request-Id,\t"], + Strict, + Expected::Valid + ), + ( + access_control_expose_headers_late_error, + AccessControlExposeHeaders, + &["etag", "x-request-id", "x-bad:name"], + Strict, + Expected::Error(DecodeErrorKind::InvalidToken) + ), + (access_control_expose_headers_absent, AccessControlExposeHeaders, &[], Strict, Expected::Absent), + (access_control_max_age_600, AccessControlMaxAge, &["600"], Strict, Expected::Valid), + (access_control_max_age_zero, AccessControlMaxAge, &["0"], Strict, Expected::Valid), + (access_control_max_age_whitespace, AccessControlMaxAge, &[" \t00600\t "], Strict, Expected::Valid), + (access_control_max_age_nineteen_digits, AccessControlMaxAge, &["9999999999999999999"], Strict, Expected::Valid), + (access_control_max_age_maximum, AccessControlMaxAge, &["18446744073709551615"], Strict, Expected::Valid), + (access_control_max_age_many_zeroes, AccessControlMaxAge, &["00000000000000000000000000000600"], Strict, Expected::Valid), + (access_control_max_age_absent, AccessControlMaxAge, &[], Strict, Expected::Absent), + ( + access_control_max_age_overflow, + AccessControlMaxAge, + &["18446744073709551616"], + Strict, + Expected::Error(DecodeErrorKind::InvalidNumber) + ), + (access_control_max_age_invalid_relaxed, AccessControlMaxAge, &["-1"], Relaxed, Expected::Error(DecodeErrorKind::InvalidNumber)), + ( + access_control_max_age_repeated, + AccessControlMaxAge, + &["600", "600"], + Strict, + Expected::Error(DecodeErrorKind::UnexpectedMultipleValues) + ), + (access_control_request_headers_pair, AccessControlRequestHeaders, &["content-type, x-request-id"], Strict, Expected::Valid), + (access_control_request_headers_single, AccessControlRequestHeaders, &["content-type"], Strict, Expected::Valid), + ( + access_control_request_headers_mixed_case, + AccessControlRequestHeaders, + &["X-Trace-Id, Content-Type"], + Strict, + Expected::Valid + ), + (access_control_request_headers_custom, AccessControlRequestHeaders, &["x-idempotency-key"], Strict, Expected::Valid), + ( + access_control_request_headers_large, + AccessControlRequestHeaders, + &["authorization, content-type, traceparent, tracestate, x-api-version, x-client-version, x-correlation-id, x-idempotency-key, x-request-id, x-tenant-id, x-trace-id"], + Strict, + Expected::Valid + ), + ( + access_control_request_headers_repeated, + AccessControlRequestHeaders, + &["X-Trace-Id, content-type", "x-trace-id", "authorization"], + Strict, + Expected::Valid + ), + ( + access_control_request_headers_empty_then_member, + AccessControlRequestHeaders, + &[" , ", "content-type", ""], + Strict, + Expected::Valid + ), + (access_control_request_headers_absent, AccessControlRequestHeaders, &[], Strict, Expected::Absent), + ( + access_control_request_headers_empty_repeated, + AccessControlRequestHeaders, + &[" , ", "\t,,", ""], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + access_control_request_headers_late_error, + AccessControlRequestHeaders, + &["content-type", "x-trace-id", "bad:name"], + Strict, + Expected::Error(DecodeErrorKind::InvalidToken) + ), + (access_control_request_method_post, AccessControlRequestMethod, &["POST"], Strict, Expected::Valid), + (access_control_request_method_get, AccessControlRequestMethod, &["GET"], Strict, Expected::Valid), + (access_control_request_method_patch, AccessControlRequestMethod, &["PATCH"], Strict, Expected::Valid), + (access_control_request_method_options, AccessControlRequestMethod, &["OPTIONS"], Strict, Expected::Valid), + (access_control_request_method_extension, AccessControlRequestMethod, &["PROPFIND"], Strict, Expected::Valid), + ( + access_control_request_method_long_extension, + AccessControlRequestMethod, + &["X-REBUILD-SEARCH-INDEX-FOR-TENANT"], + Strict, + Expected::Valid + ), + (access_control_request_method_lowercase, AccessControlRequestMethod, &["post"], Strict, Expected::Valid), + (access_control_request_method_framed_registered, AccessControlRequestMethod, &[" \tPOST\t "], Strict, Expected::Valid), + (access_control_request_method_absent, AccessControlRequestMethod, &[], Strict, Expected::Absent), + ( + access_control_request_method_invalid_relaxed, + AccessControlRequestMethod, + &["GET, POST"], + Relaxed, + Expected::Error(DecodeErrorKind::InvalidToken) + ), + ( + access_control_request_method_repeated, + AccessControlRequestMethod, + &["POST", "POST"], + Strict, + Expected::Error(DecodeErrorKind::UnexpectedMultipleValues) + ), + (content_length_348, ContentLength, &["348"], Strict, Expected::Valid), + (content_length_zero, ContentLength, &["0"], Strict, Expected::Valid), + (content_length_megabyte, ContentLength, &["1048576"], Strict, Expected::Valid), + (content_length_nineteen_digits, ContentLength, &["9999999999999999999"], Strict, Expected::Valid), + (content_length_maximum, ContentLength, &["18446744073709551615"], Strict, Expected::Valid), + (content_length_many_zeroes, ContentLength, &["00000000000000000000000000000348"], Strict, Expected::Valid), + (content_length_whitespace, ContentLength, &[" \t348\t "], Strict, Expected::Valid), + (content_length_joined_equal, ContentLength, &["348, 348"], Strict, Expected::Valid), + ( + content_length_repeated_numeric_equivalence, + ContentLength, + &["000348, 348", " 348 ", "348,348", "00348"], + Strict, + Expected::Valid + ), + (content_length_absent, ContentLength, &[], Strict, Expected::Absent), + ( + content_length_overflow, + ContentLength, + &["18446744073709551616"], + Strict, + Expected::Error(DecodeErrorKind::InvalidNumber) + ), + ( + content_length_repeated_conflict, + ContentLength, + &["348", "348", "349"], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + content_length_late_malformed, + ContentLength, + &["348", "348", "348x"], + Strict, + Expected::Error(DecodeErrorKind::InvalidNumber) + ), +); diff --git a/crates/http_headers/benches/http_headers_authority_semantics.rs b/crates/http_headers/benches/http_headers_authority_semantics.rs new file mode 100644 index 000000000..67fe21906 --- /dev/null +++ b/crates/http_headers/benches/http_headers_authority_semantics.rs @@ -0,0 +1,341 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Authority inspection for routing, allow-list comparison, and forwarding. +//! +//! Each pair compares retained components with explicit consumer-side parsing +//! of the existing textual accessors. Both use the current header decoder, so +//! these are workload comparisons, not historical decoder baselines. The +//! `decode` groups include validation; `reads` groups reuse predecoded values. + +use std::borrow::Cow; +use std::hint::black_box; +use std::net::{Ipv4Addr, Ipv6Addr}; +use std::sync::OnceLock; + +use criterion::{BatchSize, BenchmarkId, Criterion}; +use http_headers::headers::{ + AccessControlAllowOrigin, AccessControlAllowOriginKind, AccessControlAllowOriginOwned, AccessControlAllowOriginView, Host, HostKind, + HostOwned, HostView, OriginHost, OriginScheme, PortConversionError, PortConversionErrorKind, +}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeMode, Field, FieldName, FieldValue, SingleValueField}; + +#[derive(Debug, Eq, PartialEq)] +enum HostSemantic<'a> { + Name(Cow<'a, str>), + Ipv4(Ipv4Addr), + Ipv6(Ipv6Addr), + Future(&'a str, &'a str), +} + +fn retained_host<'a>(view: &'a HostView<'_>) -> (HostSemantic<'a>, Result, PortConversionErrorKind>) { + let host = match view.kind() { + HostKind::RegisteredName(name) => HostSemantic::Name(Cow::Borrowed(name.normalized())), + HostKind::Ipv4(address) => HostSemantic::Ipv4(address), + HostKind::Ipv6(address) => HostSemantic::Ipv6(address), + HostKind::IpvFuture(future) => HostSemantic::Future(future.version(), future.address()), + }; + (host, view.network_port().map_err(PortConversionError::kind)) +} + +fn consumer_host<'a>(view: &HostView<'a>) -> (HostSemantic<'a>, Result, PortConversionErrorKind>) { + let host = view.host(); + let semantic = if let Some(literal) = host.strip_prefix('[').and_then(|host| host.strip_suffix(']')) { + if let Some(future) = literal.strip_prefix('v').or_else(|| literal.strip_prefix('V')) { + let (version, address) = future.split_once('.').expect("validated IPvFuture contains a version delimiter"); + HostSemantic::Future(version, address) + } else { + HostSemantic::Ipv6(literal.parse().expect("validated IPv6 parses")) + } + } else if !host.is_ascii() { + HostSemantic::Name(Cow::Owned( + idna::domain_to_ascii(host).expect("relaxed host already passed IDNA validation"), + )) + } else if let Ok(address) = host.parse::() { + HostSemantic::Ipv4(address) + } else { + HostSemantic::Name(Cow::Borrowed(host)) + }; + let port = view + .port() + .map(|port| { + if port.is_empty() { + Err(PortConversionErrorKind::Empty) + } else { + port.parse::().map_err(|_invalid| PortConversionErrorKind::Overflow) + } + }) + .transpose(); + (semantic, port) +} + +struct HostFixture { + wire: FieldValue, + mode: DecodeMode, + owned: HostOwned, +} + +impl HostFixture { + fn new(wire: &'static str, mode: DecodeMode) -> Self { + let wire = FieldValue::try_from(wire).expect("fixture is a field value"); + let owned = ::decode_owned_with(wire.clone(), mode).expect("fixture is a valid host"); + let view = owned.as_view(); + assert_eq!(retained_host(&view), consumer_host(&view)); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + drop(view); + Self { wire, mode, owned } + } +} + +#[derive(Debug, Eq, PartialEq)] +enum OriginSemantic<'a> { + Wildcard, + Null, + Tuple(OriginScheme, HostSemantic<'a>, Option, u16), +} + +fn retained_origin(view: AccessControlAllowOriginView<'_>) -> OriginSemantic<'_> { + match view.kind() { + AccessControlAllowOriginKind::Wildcard => OriginSemantic::Wildcard, + AccessControlAllowOriginKind::Null => OriginSemantic::Null, + AccessControlAllowOriginKind::Origin(origin) => { + let host = match origin.host() { + OriginHost::Domain(domain) => HostSemantic::Name(Cow::Borrowed(domain.as_str())), + OriginHost::Ipv4(address) => HostSemantic::Ipv4(address), + OriginHost::Ipv6(address) => HostSemantic::Ipv6(address), + }; + OriginSemantic::Tuple(origin.scheme(), host, origin.port(), origin.effective_port()) + } + } +} + +#[expect(clippy::panic, reason = "the benchmark consumes only previously validated tuple-origin schemes")] +fn consumer_origin(view: AccessControlAllowOriginView<'_>) -> OriginSemantic<'_> { + let serialized = view.as_str(); + if serialized == "*" { + return OriginSemantic::Wildcard; + } + if serialized == "null" { + return OriginSemantic::Null; + } + let (scheme, authority) = serialized.split_once("://").expect("validated tuple origin has a scheme"); + let scheme = match scheme { + "ftp" => OriginScheme::Ftp, + "http" => OriginScheme::Http, + "https" => OriginScheme::Https, + "ws" => OriginScheme::Ws, + "wss" => OriginScheme::Wss, + _ => panic!("validated tuple-origin scheme"), + }; + let (host, port) = if let Some(literal) = authority.strip_prefix('[') { + let (host, suffix) = literal.split_once(']').expect("validated IPv6 has a closing bracket"); + let address = host.parse().expect("validated origin IPv6 parses"); + (HostSemantic::Ipv6(address), suffix.strip_prefix(':')) + } else { + let (host, port) = authority + .split_once(':') + .map_or((authority, None), |(host, port)| (host, Some(port))); + let host = host + .strip_suffix('.') + .unwrap_or(host) + .parse::() + .map_or_else(|_| HostSemantic::Name(Cow::Borrowed(host)), HostSemantic::Ipv4); + (host, port) + }; + let port = port.map(|port| port.parse::().expect("validated origin port fits u16")); + OriginSemantic::Tuple(scheme, host, port, port.unwrap_or_else(|| scheme.default_port())) +} + +struct OriginFixture { + wire: FieldValue, + owned: AccessControlAllowOriginOwned, +} + +impl OriginFixture { + fn new(wire: &'static str) -> Self { + let wire = FieldValue::try_from(wire).expect("fixture is a field value"); + let owned = AccessControlAllowOriginOwned::try_from(wire.clone()).expect("fixture is a valid origin"); + let view = owned.as_view(); + assert_eq!(retained_origin(view), consumer_origin(view)); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + Self { wire, owned } + } +} + +impl FieldSource for OriginFixture { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::AccessControlAllowOrigin).then(|| FieldLines::single(name, self.wire.as_bytes())) + } +} + +macro_rules! host_cases { + ($(($id:ident, $wire:literal, $mode:ident)),+ $(,)?) => { + $( + fn $id() -> &'static HostFixture { + static FIXTURE: OnceLock = OnceLock::new(); + FIXTURE.get_or_init(|| HostFixture::new($wire, DecodeMode::$mode)) + } + )+ + + #[metabench::benchmark(HOST_DECODE_RETAINED, "http_headers_authority_semantics/host_decode", "retained")] + $(#[bench::$id(setup = $id)])+ + fn host_decode_retained(fixture: &HostFixture) { + let view = ::decode_view_with(black_box(fixture.wire.as_field_value_ref()), fixture.mode) + .expect("validated fixture"); + drop(black_box(retained_host(black_box(&view)))); + black_box(view.as_field_value().as_bytes()); + } + + #[metabench::benchmark(HOST_DECODE_CONSUMER, "http_headers_authority_semantics/host_decode", "consumer_parse")] + $(#[bench::$id(setup = $id)])+ + fn host_decode_consumer(fixture: &HostFixture) { + let view = ::decode_view_with(black_box(fixture.wire.as_field_value_ref()), fixture.mode) + .expect("validated fixture"); + drop(black_box(consumer_host(black_box(&view)))); + black_box(view.as_field_value().as_bytes()); + } + + #[metabench::benchmark(HOST_READS_RETAINED, "http_headers_authority_semantics/host_reads", "retained")] + $(#[bench::$id(setup = $id)])+ + fn host_reads_retained(fixture: &HostFixture) { + let view = black_box(&fixture.owned).as_view(); + drop(black_box(retained_host(black_box(&view)))); + black_box(view.as_field_value().as_bytes()); + } + + #[metabench::benchmark(HOST_READS_CONSUMER, "http_headers_authority_semantics/host_reads", "consumer_parse")] + $(#[bench::$id(setup = $id)])+ + fn host_reads_consumer(fixture: &HostFixture) { + let view = black_box(&fixture.owned).as_view(); + drop(black_box(consumer_host(black_box(&view)))); + black_box(view.as_field_value().as_bytes()); + } + + fn host_benchmarks(criterion: &mut Criterion) { + let mut group = criterion.benchmark_group("http_headers_authority_semantics/host_decode"); + $( + group.bench_function(BenchmarkId::new("retained", stringify!($id)), + |b| b.iter_batched($id, host_decode_retained, BatchSize::SmallInput)); + group.bench_function(BenchmarkId::new("consumer_parse", stringify!($id)), + |b| b.iter_batched($id, host_decode_consumer, BatchSize::SmallInput)); + )+ + group.finish(); + let mut group = criterion.benchmark_group("http_headers_authority_semantics/host_reads"); + $( + group.bench_function(BenchmarkId::new("retained", stringify!($id)), + |b| b.iter_batched($id, host_reads_retained, BatchSize::SmallInput)); + group.bench_function(BenchmarkId::new("consumer_parse", stringify!($id)), + |b| b.iter_batched($id, host_reads_consumer, BatchSize::SmallInput)); + )+ + group.finish(); + } + }; +} + +macro_rules! origin_cases { + ($(($id:ident, $wire:literal)),+ $(,)?) => { + $( + fn $id() -> &'static OriginFixture { + static FIXTURE: OnceLock = OnceLock::new(); + FIXTURE.get_or_init(|| OriginFixture::new($wire)) + } + )+ + + #[metabench::benchmark(ORIGIN_DECODE_RETAINED, "http_headers_authority_semantics/origin_decode", "retained")] + $(#[bench::$id(setup = $id)])+ + fn origin_decode_retained(fixture: &OriginFixture) { + let view = ::view(black_box(fixture)).expect("validated fixture").expect("present"); + black_box(retained_origin(black_box(view))); + black_box(view.as_field_value().as_bytes()); + } + + #[metabench::benchmark(ORIGIN_DECODE_CONSUMER, "http_headers_authority_semantics/origin_decode", "consumer_parse")] + $(#[bench::$id(setup = $id)])+ + fn origin_decode_consumer(fixture: &OriginFixture) { + let view = ::view(black_box(fixture)).expect("validated fixture").expect("present"); + black_box(consumer_origin(black_box(view))); + black_box(view.as_field_value().as_bytes()); + } + + #[metabench::benchmark(ORIGIN_READS_RETAINED, "http_headers_authority_semantics/origin_reads", "retained")] + $(#[bench::$id(setup = $id)])+ + fn origin_reads_retained(fixture: &OriginFixture) { + let view = black_box(&fixture.owned).as_view(); + black_box(retained_origin(black_box(view))); + black_box(view.as_field_value().as_bytes()); + } + + #[metabench::benchmark(ORIGIN_READS_CONSUMER, "http_headers_authority_semantics/origin_reads", "consumer_parse")] + $(#[bench::$id(setup = $id)])+ + fn origin_reads_consumer(fixture: &OriginFixture) { + let view = black_box(&fixture.owned).as_view(); + black_box(consumer_origin(black_box(view))); + black_box(view.as_field_value().as_bytes()); + } + + fn origin_benchmarks(criterion: &mut Criterion) { + let mut group = criterion.benchmark_group("http_headers_authority_semantics/origin_decode"); + $( + group.bench_function(BenchmarkId::new("retained", stringify!($id)), + |b| b.iter_batched($id, origin_decode_retained, BatchSize::SmallInput)); + group.bench_function(BenchmarkId::new("consumer_parse", stringify!($id)), + |b| b.iter_batched($id, origin_decode_consumer, BatchSize::SmallInput)); + )+ + group.finish(); + let mut group = criterion.benchmark_group("http_headers_authority_semantics/origin_reads"); + $( + group.bench_function(BenchmarkId::new("retained", stringify!($id)), + |b| b.iter_batched($id, origin_reads_retained, BatchSize::SmallInput)); + group.bench_function(BenchmarkId::new("consumer_parse", stringify!($id)), + |b| b.iter_batched($id, origin_reads_consumer, BatchSize::SmallInput)); + )+ + group.finish(); + } + }; +} + +host_cases!( + (domain, "api.example.com:8443", Strict), + (ipv4, "192.0.2.1:443", Strict), + (ipv6, "[2001:db8::1]:443", Strict), + (ipvfuture, "[vFFFFFFFFFFFFFFFFFFFF.alpha:beta]:443", Strict), + (numeric_name, "192.168.001.1:80", Strict), + (percent_name, "api%2Eexample.com:443", Strict), + (empty_port, "example.com:", Strict), + (overflow_port, "example.com:65536", Strict), + (leading_zero_port, "example.com:0000000000000000000000443", Strict), + (idna_name, "münich.example:443", Relaxed), +); + +origin_cases!( + (origin_wildcard, "*"), + (origin_null, "null"), + (origin_domain, "https://api.example.com"), + (origin_explicit_port, "https://api.example.com:8443"), + (origin_ipv4, "https://192.0.2.1:8443"), + (origin_ipv6, "https://[2001:db8::1]:8443"), +); + +fn criterion_benchmarks(criterion: &mut Criterion) { + host_benchmarks(criterion); + origin_benchmarks(criterion); +} + +metabench::main!( + criterion = { + factory = Criterion::default, + benchmarks = criterion_benchmarks, + unit = "ns", + }, + benchmarks = [ + HOST_DECODE_RETAINED, + HOST_DECODE_CONSUMER, + HOST_READS_RETAINED, + HOST_READS_CONSUMER, + ORIGIN_DECODE_RETAINED, + ORIGIN_DECODE_CONSUMER, + ORIGIN_READS_RETAINED, + ORIGIN_READS_CONSUMER, + ], +); diff --git a/crates/http_headers/benches/http_headers_common_values.rs b/crates/http_headers/benches/http_headers_common_values.rs new file mode 100644 index 000000000..3a9b39fa3 --- /dev/null +++ b/crates/http_headers/benches/http_headers_common_values.rs @@ -0,0 +1,7 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Wire values shared by `http_headers` benchmark targets. + +/// `Basic` credentials for `aladdin:opensesame`. +pub(crate) const BASIC_AUTHORIZATION: &str = "Basic YWxhZGRpbjpvcGVuc2VzYW1l"; diff --git a/crates/http_headers/benches/http_headers_conditional_range_shapes.rs b/crates/http_headers/benches/http_headers_conditional_range_shapes.rs new file mode 100644 index 000000000..cdce3f207 --- /dev/null +++ b/crates/http_headers/benches/http_headers_conditional_range_shapes.rs @@ -0,0 +1,139 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Decode-only shapes for conditional and range headers, excluding semantic reader costs. + +use http_headers::DecodeErrorKind; +use http_headers::headers::{ + AcceptRanges, ContentRange, ETag, IfMatch, IfModifiedSince, IfNoneMatch, IfRange, IfUnmodifiedSince, LastModified, Range, +}; + +#[path = "http_headers_shapes_common.rs"] +mod shapes; + +use shapes::Expected; + +shapes::define_shapes!( + "http_headers_conditional_range_shapes/parse"; + // Extra lengths isolate the scalar/word scanner and inline/shared storage boundaries. + (etag_absent, ETag, &[], Strict, Expected::Absent), + (etag_short, ETag, &["\"x\""], Strict, Expected::Valid), + (etag_revision, ETag, &["\"revision-42\""], Strict, Expected::Valid), + (etag_weak_revision, ETag, &["W/\"revision-42\""], Strict, Expected::Valid), + (etag_len_seven, ETag, &["\"1234567\""], Strict, Expected::Valid), + (etag_len_eight, ETag, &["\"12345678\""], Strict, Expected::Valid), + (etag_len_nine, ETag, &["\"123456789\""], Strict, Expected::Valid), + (etag_wire_sixty_four, ETag, &["\"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcd\""], Strict, Expected::Valid), + (etag_wire_sixty_five, ETag, &["\"0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcde\""], Strict, Expected::Valid), + (etag_lowercase_weak_relaxed, ETag, &["w/\"revision-42\""], Relaxed, Expected::Valid), + (etag_space_late, ETag, &["\"0123456789abcdef0123456789abcdef bad\""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (etag_repeated, ETag, &["\"one\"", "\"two\""], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + + (if_match_one, IfMatch, &["\"revision-42\""], Strict, Expected::Valid), + (if_match_wildcard, IfMatch, &["*"], Strict, Expected::Valid), + (if_match_two, IfMatch, &["\"a\", W/\"b\""], Strict, Expected::Valid), + (if_match_long, IfMatch, &["\"sha256-0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef\""], Strict, Expected::Valid), + (if_match_many, IfMatch, &["\"revision-0\", \"revision-1\", \"revision-2\", \"revision-3\", \"revision-4\", \"revision-5\", \"revision-6\", \"revision-7\", \"revision-8\", \"revision-9\", \"revision-10\", \"revision-11\", \"revision-12\", \"revision-13\", \"revision-14\", \"revision-15\""], Strict, Expected::Valid), + (if_match_two_lines, IfMatch, &["\"one\"", "W/\"two\""], Strict, Expected::Valid), + (if_match_comma_backslash, IfMatch, &["\"one,two\", \"comma,slash\\\""], Strict, Expected::Valid), + (if_match_lowercase_relaxed, IfMatch, &["w/\"revision\", \"next\""], Relaxed, Expected::Valid), + (if_match_late_invalid, IfMatch, &["\"one\"", "W/\"two\"", "not-a-tag"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (if_match_empty_field, IfMatch, &[""], Strict, Expected::Error(DecodeErrorKind::MissingValue)), + + (if_none_match_one, IfNoneMatch, &["\"revision-42\""], Strict, Expected::Valid), + (if_none_match_wildcard, IfNoneMatch, &["*"], Strict, Expected::Valid), + (if_none_match_two, IfNoneMatch, &["W/\"a\", \"b\""], Strict, Expected::Valid), + (if_none_match_long, IfNoneMatch, &["\"sha256-0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef\""], Strict, Expected::Valid), + (if_none_match_many, IfNoneMatch, &["\"revision-0\", \"revision-1\", \"revision-2\", \"revision-3\", \"revision-4\", \"revision-5\", \"revision-6\", \"revision-7\", \"revision-8\", \"revision-9\", \"revision-10\", \"revision-11\", \"revision-12\", \"revision-13\", \"revision-14\", \"revision-15\""], Strict, Expected::Valid), + (if_none_match_two_lines, IfNoneMatch, &["\"one\"", "W/\"two\""], Strict, Expected::Valid), + (if_none_match_comma_backslash, IfNoneMatch, &["\"one,two\", \"comma,slash\\\""], Strict, Expected::Valid), + (if_none_match_lowercase_relaxed, IfNoneMatch, &["w/\"revision\", \"next\""], Relaxed, Expected::Valid), + (if_none_match_late_invalid, IfNoneMatch, &["\"one\"", "W/\"two\"", "not-a-tag"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (if_none_match_empty_field, IfNoneMatch, &[""], Strict, Expected::Error(DecodeErrorKind::MissingValue)), + + (if_range_revision, IfRange, &["\"revision-42\""], Strict, Expected::Valid), + (if_range_long_tag, IfRange, &["\"sha256-0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef\""], Strict, Expected::Valid), + (if_range_imf_date, IfRange, &["Sun, 06 Nov 1994 08:49:37 GMT"], Strict, Expected::Valid), + (if_range_rfc850, IfRange, &["Sunday, 06-Nov-94 08:49:37 GMT"], Strict, Expected::Valid), + (if_range_asctime, IfRange, &["Sun Nov 6 08:49:37 1994"], Strict, Expected::Valid), + (if_range_relaxed_date, IfRange, &[" Tue, 8 Nov 1994 8:49:37 UTC "], Relaxed, Expected::Valid), + (if_range_relaxed_strict_date, IfRange, &["Sun, 06 Nov 1994 08:49:37 GMT"], Relaxed, Expected::Valid), + (if_range_weak, IfRange, &["W/\"revision-42\""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (if_range_unterminated, IfRange, &["\"unterminated"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (if_range_repeated, IfRange, &["\"revision-42\"", "Sun, 06 Nov 1994 08:49:37 GMT"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + + (if_modified_since_imf, IfModifiedSince, &["Sun, 06 Nov 1994 08:49:37 GMT"], Strict, Expected::Valid), + (if_modified_since_leap, IfModifiedSince, &["Thu, 29 Feb 2024 00:00:00 GMT"], Strict, Expected::Valid), + (if_modified_since_rfc850, IfModifiedSince, &["Sunday, 06-Nov-94 08:49:37 GMT"], Strict, Expected::Valid), + (if_modified_since_asctime, IfModifiedSince, &["Sun Nov 6 08:49:37 1994"], Strict, Expected::Valid), + (if_modified_since_relaxed_canonical, IfModifiedSince, &["Sun, 06 Nov 1994 08:49:37 GMT"], Relaxed, Expected::Valid), + (if_modified_since_ows_relaxed, IfModifiedSince, &[" Sun, 06 Nov 1994 08:49:37 GMT "], Relaxed, Expected::Valid), + (if_modified_since_normalized, IfModifiedSince, &[" Tue, 8 Nov 1994 8:49:37 UTC "], Relaxed, Expected::Valid), + (if_modified_since_wrong_weekday, IfModifiedSince, &["Mon, 06 Nov 1994 08:49:37 GMT"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (if_modified_since_empty, IfModifiedSince, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (if_modified_since_repeated, IfModifiedSince, &["Sun, 06 Nov 1994 08:49:37 GMT", "Mon, 07 Nov 1994 08:49:37 GMT"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + + (if_unmodified_since_imf, IfUnmodifiedSince, &["Sun, 06 Nov 1994 08:49:37 GMT"], Strict, Expected::Valid), + (if_unmodified_since_leap, IfUnmodifiedSince, &["Thu, 29 Feb 2024 00:00:00 GMT"], Strict, Expected::Valid), + (if_unmodified_since_rfc850, IfUnmodifiedSince, &["Sunday, 06-Nov-94 08:49:37 GMT"], Strict, Expected::Valid), + (if_unmodified_since_asctime, IfUnmodifiedSince, &["Sun Nov 6 08:49:37 1994"], Strict, Expected::Valid), + (if_unmodified_since_relaxed_canonical, IfUnmodifiedSince, &["Sun, 06 Nov 1994 08:49:37 GMT"], Relaxed, Expected::Valid), + (if_unmodified_since_ows_relaxed, IfUnmodifiedSince, &[" Sun, 06 Nov 1994 08:49:37 GMT "], Relaxed, Expected::Valid), + (if_unmodified_since_normalized, IfUnmodifiedSince, &[" Tue, 8 Nov 1994 8:49:37 UTC "], Relaxed, Expected::Valid), + (if_unmodified_since_wrong_weekday, IfUnmodifiedSince, &["Mon, 06 Nov 1994 08:49:37 GMT"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (if_unmodified_since_empty, IfUnmodifiedSince, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (if_unmodified_since_repeated, IfUnmodifiedSince, &["Sun, 06 Nov 1994 08:49:37 GMT", "Mon, 07 Nov 1994 08:49:37 GMT"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + + (last_modified_imf, LastModified, &["Sun, 06 Nov 1994 08:49:37 GMT"], Strict, Expected::Valid), + (last_modified_leap, LastModified, &["Thu, 29 Feb 2024 00:00:00 GMT"], Strict, Expected::Valid), + (last_modified_rfc850, LastModified, &["Sunday, 06-Nov-94 08:49:37 GMT"], Strict, Expected::Valid), + (last_modified_asctime, LastModified, &["Sun Nov 6 08:49:37 1994"], Strict, Expected::Valid), + (last_modified_relaxed_canonical, LastModified, &["Sun, 06 Nov 1994 08:49:37 GMT"], Relaxed, Expected::Valid), + (last_modified_ows_relaxed, LastModified, &[" Sun, 06 Nov 1994 08:49:37 GMT "], Relaxed, Expected::Valid), + (last_modified_normalized, LastModified, &[" Tue, 8 Nov 1994 8:49:37 UTC "], Relaxed, Expected::Valid), + (last_modified_wrong_weekday, LastModified, &["Mon, 06 Nov 1994 08:49:37 GMT"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (last_modified_empty, LastModified, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (last_modified_repeated, LastModified, &["Sun, 06 Nov 1994 08:49:37 GMT", "Mon, 07 Nov 1994 08:49:37 GMT"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + + // Extra range cases distinguish scanner fallbacks from numeric and cardinality errors. + (range_closed, Range, &["bytes=0-499"], Strict, Expected::Valid), + (range_existing, Range, &["bytes=0-499, 1000-"], Strict, Expected::Valid), + (range_suffix, Range, &["bytes=-500"], Strict, Expected::Valid), + (range_many_members, Range, &["bytes=0-9, 20-29, 40-49, 60-69, 80-89, 100-109, 120-129, 140-149, 160-169, 180-189, 200-209, 220-229, 240-249, 260-269, 280-289, 300-309"], Strict, Expected::Valid), + (range_large_offsets, Range, &["bytes=1048576-2097151, 3145728-4194303"], Strict, Expected::Valid), + (range_maximum_open, Range, &["bytes=18446744073709551615-"], Strict, Expected::Valid), + (range_leading_zeroes, Range, &["bytes=00000000000000000000005-6"], Strict, Expected::Valid), + (range_case_fallback, Range, &["Bytes=0-499, 1000-"], Strict, Expected::Valid), + (range_extension, Range, &["items=1-5"], Strict, Expected::Valid), + (range_relaxed_whitespace, Range, &[" Bytes = 0 - 9 , - 5 "], Relaxed, Expected::Valid), + (range_inverted, Range, &["bytes=10-9"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (range_overflow, Range, &["bytes=18446744073709551616-"], Strict, Expected::Error(DecodeErrorKind::InvalidNumber)), + (range_repeated, Range, &["bytes=0-1", "bytes=2-3"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + + // Quote, token and cardinality errors take distinct fallback paths. + (accept_ranges_bytes, AcceptRanges, &["bytes"], Strict, Expected::Valid), + (accept_ranges_none, AcceptRanges, &["none"], Strict, Expected::Valid), + (accept_ranges_case_bytes, AcceptRanges, &["Bytes"], Strict, Expected::Valid), + (accept_ranges_ows, AcceptRanges, &[" \tbytes \t"], Strict, Expected::Valid), + (accept_ranges_two, AcceptRanges, &["bytes, items"], Strict, Expected::Valid), + (accept_ranges_repeated, AcceptRanges, &["bytes", "items", "records"], Strict, Expected::Valid), + (accept_ranges_many_lines, AcceptRanges, &["bytes", "items", "records", "frames", "pages", "lines", "chunks", "segments", "blocks", "entries", "samples", "packets", "rows", "columns", "cells", "objects"], Strict, Expected::Valid), + (accept_ranges_none_conflict, AcceptRanges, &["none, bytes"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (accept_ranges_late_invalid, AcceptRanges, &["bytes", "items", "bad/unit"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (accept_ranges_unterminated, AcceptRanges, &["\"bytes"], Strict, Expected::Error(DecodeErrorKind::UnterminatedQuote)), + (accept_ranges_empty_field, AcceptRanges, &[""], Strict, Expected::Error(DecodeErrorKind::MissingValue)), + + // Empty extensions, unknown lengths and unsatisfied responses have different grammars. + (content_range_existing, ContentRange, &["bytes 0-499/1234"], Strict, Expected::Valid), + (content_range_unknown, ContentRange, &["bytes 500-999/*"], Strict, Expected::Valid), + (content_range_unsatisfied, ContentRange, &["bytes */1234"], Strict, Expected::Valid), + (content_range_large, ContentRange, &["bytes 1048576-2097151/4294967296"], Strict, Expected::Valid), + (content_range_leading_zeroes, ContentRange, &["bytes 00000000000000000000005-0006/000010"], Strict, Expected::Valid), + (content_range_case, ContentRange, &["Bytes 0-499/1234"], Strict, Expected::Valid), + (content_range_extension_empty, ContentRange, &["items "], Strict, Expected::Valid), + (content_range_extension_long, ContentRange, &["example-unit opaque response segment-0000000000000001 to segment-0000000000000002"], Strict, Expected::Valid), + (content_range_relaxed_spaces, ContentRange, &["Bytes 0 - 499 / 1234"], Relaxed, Expected::Valid), + (content_range_overflow, ContentRange, &["bytes 0-1/18446744073709551616"], Strict, Expected::Error(DecodeErrorKind::InvalidNumber)), + (content_range_length_too_small, ContentRange, &["bytes 0-499/499"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (content_range_repeated, ContentRange, &["bytes 0-9/100", "bytes 20-29/100"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), +); diff --git a/crates/http_headers/benches/http_headers_fixtures.rs b/crates/http_headers/benches/http_headers_fixtures.rs new file mode 100644 index 000000000..4d687fc61 --- /dev/null +++ b/crates/http_headers/benches/http_headers_fixtures.rs @@ -0,0 +1,197 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! The corpus every `http_headers` benchmark shares with its `headers` twin. +//! +//! Both competitors read byte-identical [`HeaderMap`]s built here, so no +//! comparison degenerates into a comparison of fixtures. Included with +//! `#[path]` rather than reached through the library, because each benchmark +//! file is its own crate and the corpus must not enter the public API. +//! +//! Values are built with [`HeaderValue::from_bytes`] rather than +//! `from_static`, so their storage is heap-backed exactly as it is after a +//! real request is parsed off a socket. That matters: cloning a static value +//! is free, and a fixture of static values would hide every clone the +//! competitor performs. + +use http::header::{ + ACCEPT, ACCEPT_ENCODING, ACCEPT_LANGUAGE, AUTHORIZATION, CACHE_CONTROL, CONNECTION, CONTENT_LENGTH, CONTENT_TYPE, COOKIE, HOST, + REFERER, USER_AGENT, +}; +use http::{HeaderMap, HeaderName, HeaderValue}; + +#[path = "http_headers_common_values.rs"] +mod common_values; + +/// A browser `User-Agent`, long enough to cross the SIMD dispatch threshold. +pub(crate) const USER_AGENT_VALUE: &[u8] = + b"Mozilla/5.0 (X11; Linux x86_64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/124.0.0.0 Safari/537.36"; + +/// The `Content-Type` of an ordinary JSON API request. +pub(crate) const CONTENT_TYPE_VALUE: &[u8] = b"application/json; charset=utf-8"; + +/// A `Content-Type` neither crate can parse. +pub(crate) const CONTENT_TYPE_MALFORMED: &[u8] = b"application"; + +/// A single-line `Cache-Control` with three directives. +pub(crate) const CACHE_CONTROL_VALUE: &[u8] = b"max-age=3600, public, must-revalidate"; + +/// A directive set delivered as three separate field lines. +pub(crate) const CACHE_CONTROL_LINES: [&[u8]; 3] = [b"max-age=3600", b"public, must-revalidate", b"no-transform, s-maxage=120"]; + +/// A signed JWT credential of realistic length. +pub(crate) const BEARER_VALUE: &[u8] = b"Bearer eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwibmFtZSI6IkpvaG4gRG9lIiwiaWF0IjoxNTE2MjM5MDIyfQ.SflKxwRJSMeKKF2QT4fwpMeJf36POk6yJV_adQssw5c"; + +/// The credential portion of [`BEARER_VALUE`], after the scheme. +pub(crate) const BEARER_TOKEN: &[u8] = b"eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwibmFtZSI6IkpvaG4gRG9lIiwiaWF0IjoxNTE2MjM5MDIyfQ.SflKxwRJSMeKKF2QT4fwpMeJf36POk6yJV_adQssw5c"; + +/// `Basic` credentials for `aladdin:opensesame`. +pub(crate) const BASIC_VALUE: &[u8] = common_values::BASIC_AUTHORIZATION.as_bytes(); + +/// The username encoded in [`BASIC_VALUE`]. +pub(crate) const BASIC_USERNAME: &[u8] = b"aladdin"; + +/// The password encoded in [`BASIC_VALUE`]. +pub(crate) const BASIC_PASSWORD: &[u8] = b"opensesame"; + +/// Four `Set-Cookie` field lines of the kind a login response emits. +pub(crate) const SET_COOKIE_VALUES: [&[u8]; 4] = [ + b"session=6f1c2a9b4e7d8f3a5c0b1d2e3f4a5b6c; Path=/; HttpOnly; Secure; SameSite=Lax", + b"csrf=9a8b7c6d5e4f3a2b1c0d9e8f7a6b5c4d; Path=/; Secure; SameSite=Strict", + b"theme=dark; Path=/; Max-Age=31536000", + b"locale=en-US; Path=/; Max-Age=31536000", +]; + +/// A downstream extension header carried by every request in the corpus. +pub(crate) static REQUEST_ID_NAME: HeaderName = HeaderName::from_static("x-request-id"); + +/// The value carried by [`REQUEST_ID_NAME`]. +pub(crate) const REQUEST_ID_VALUE: &[u8] = b"0f6c9a3e-84d1-4f2b-9c77-2b5e1a0d6f38"; + +/// The `charset` parameter value both crates look up. +pub(crate) const CHARSET: &[u8] = b"utf-8"; + +/// Builds a header value, failing loudly rather than measuring a bad fixture. +pub(crate) fn value(bytes: &[u8]) -> HeaderValue { + HeaderValue::from_bytes(bytes).expect("benchmark fixture is not a legal header value") +} + +/// Builds the [`http_headers::FieldValue`] counterpart of [`value`]. +pub(crate) fn field_value(bytes: &[u8]) -> http_headers::FieldValue { + http_headers::FieldValue::from_bytes(bytes).expect("benchmark fixture is not a legal field value") +} + +/// Asserts a measured operation produced exactly `expected` and returns its +/// length, so the optimizer cannot delete the work that produced it. +/// +/// Both competitors call this identical function on identical bytes, so the +/// comparison it guards costs the same on either side. +pub(crate) fn expect_bytes(actual: &[u8], expected: &[u8]) -> usize { + assert!(actual == expected, "benchmark produced unexpected bytes"); + actual.len() +} + +/// Asserts a measured operation produced exactly `expected`. +pub(crate) fn expect_usize(actual: usize, expected: usize) -> usize { + assert!(actual == expected, "benchmark produced unexpected count"); + actual +} + +/// Consumes a fallible insertion, deliberately without requiring the error to +/// be `Debug`. +/// +/// `Field::insert` is fallible because `HeaderMap` has a maximum capacity, and +/// the error types that describe that condition are not all `Debug`: +/// `http`'s own `TryEntryError` explicitly is not. Going through `is_ok` +/// rather than `expect` keeps every call site here compiling whichever error +/// the API settles on, and it is also what a real caller writes: one branch, +/// which is exactly what the measured insertion cases should be paying for. +pub(crate) fn expect_inserted(result: Result<(), E>) -> usize { + let inserted = match result { + Ok(()) => 1, + Err(_full) => 0, + }; + expect_usize(inserted, 1) +} + +/// The headers of an ordinary authenticated JSON API request. +pub(crate) fn json_request() -> HeaderMap { + let mut map = HeaderMap::with_capacity(16); + let _ = map.insert(HOST, value(b"api.example.com")); + let _ = map.insert(USER_AGENT, value(USER_AGENT_VALUE)); + let _ = map.insert(ACCEPT, value(b"application/json, text/plain;q=0.9, */*;q=0.8")); + let _ = map.insert(ACCEPT_ENCODING, value(b"gzip, deflate, br")); + let _ = map.insert(ACCEPT_LANGUAGE, value(b"en-US,en;q=0.9")); + let _ = map.insert(CONTENT_TYPE, value(CONTENT_TYPE_VALUE)); + let _ = map.insert(CONTENT_LENGTH, value(b"348")); + let _ = map.insert(AUTHORIZATION, value(BEARER_VALUE)); + let _ = map.insert(CACHE_CONTROL, value(CACHE_CONTROL_VALUE)); + let _ = map.insert(COOKIE, value(b"session=6f1c2a9b4e7d8f3a5c0b1d2e3f4a5b6c; theme=dark")); + let _ = map.insert(&REQUEST_ID_NAME, value(REQUEST_ID_VALUE)); + let _ = map.insert(REFERER, value(b"https://app.example.com/dashboard")); + let _ = map.insert(CONNECTION, value(b"keep-alive")); + map +} + +/// The same request, with `Authorization` carrying `Basic` credentials. +pub(crate) fn basic_auth_request() -> HeaderMap { + let mut map = json_request(); + let _ = map.insert(AUTHORIZATION, value(BASIC_VALUE)); + map +} + +/// A request whose optional headers are absent or malformed. +/// +/// Middleware meets this shape constantly: no `Authorization`, no +/// `Cache-Control`, and a `Content-Type` that fails the grammar. +pub(crate) fn degraded_request() -> HeaderMap { + let mut map = HeaderMap::with_capacity(16); + let _ = map.insert(HOST, value(b"api.example.com")); + let _ = map.insert(USER_AGENT, value(USER_AGENT_VALUE)); + let _ = map.insert(CONTENT_TYPE, value(CONTENT_TYPE_MALFORMED)); + let _ = map.insert(CONTENT_LENGTH, value(b"0")); + let _ = map.insert(&REQUEST_ID_NAME, value(REQUEST_ID_VALUE)); + map +} + +/// A map whose `Cache-Control` arrives as [`CACHE_CONTROL_LINES`]. +pub(crate) fn cache_control_multi_line() -> HeaderMap { + let mut map = json_request(); + let _ = map.remove(CACHE_CONTROL); + for line in CACHE_CONTROL_LINES { + let _ = map.append(CACHE_CONTROL, value(line)); + } + map +} + +/// A legal but adversarial request: every value sits at the large end of what +/// a gateway accepts, and none of it is invalid. +pub(crate) fn adversarial_request() -> HeaderMap { + let mut map = HeaderMap::with_capacity(8); + let mut agent = Vec::with_capacity(4096); + while agent.len() < 4000 { + agent.extend_from_slice(b"Component/1.0 (build 20240101; feature-set-extended) "); + } + agent.truncate(4000); + let _ = map.insert(USER_AGENT, value(&agent)); + + let mut content_type = Vec::with_capacity(2048); + content_type.extend_from_slice(b"application/vnd.example.v3+json"); + for index in 0..48 { + content_type.extend_from_slice(format!("; param{index}=value{index}").as_bytes()); + } + content_type.extend_from_slice(b"; charset=utf-8"); + let _ = map.insert(CONTENT_TYPE, value(&content_type)); + + let mut cache_control = Vec::with_capacity(1024); + cache_control.extend_from_slice(b"max-age=3600"); + for index in 0..32 { + cache_control.extend_from_slice(format!(", ext{index}=value{index}").as_bytes()); + } + cache_control.extend_from_slice(b", public"); + let _ = map.insert(CACHE_CONTROL, value(&cache_control)); + + let _ = map.insert(AUTHORIZATION, value(BEARER_VALUE)); + let _ = map.insert(&REQUEST_ID_NAME, value(REQUEST_ID_VALUE)); + map +} diff --git a/crates/http_headers/benches/http_headers_location_semantics.rs b/crates/http_headers/benches/http_headers_location_semantics.rs new file mode 100644 index 000000000..c3c673839 --- /dev/null +++ b/crates/http_headers/benches/http_headers_location_semantics.rs @@ -0,0 +1,279 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Location decode, retained reads, and redirect forwarding workloads. +//! +//! The caller baselines reproduce pre-metadata strict validation followed by +//! one caller-side RFC 3986 parse, retained across 1/2/8 reads. They do not +//! charge a fresh parse for every accessor. + +use std::hint::black_box; +use std::sync::OnceLock; + +use criterion::{BatchSize, Criterion}; +use fluent_uri::Uri; +use http::{HeaderMap, HeaderValue}; +use http_headers::headers::{Location, LocationOwned, LocationView, UriReference}; +use http_headers::{DecodeError, DecodeErrorKind, DecodeMode, FieldName, FieldValue, SingleValueField}; + +const GROUP: &str = "http_headers_location_semantics/reads"; +const SIMPLE: &str = "https://example.com:8443/docs/next?tab=1#top"; +const GENERAL: &str = "https://user:secret@[2001:db8::1]:8443/a%2Fb?tab=%31#top"; +const RELAXED: &str = r"https:\\example.com:8443\docs\next?tab=1#top"; + +fn simple() -> FieldValue { + FieldValue::from_static(SIMPLE) +} + +fn general() -> FieldValue { + FieldValue::from_static(GENERAL) +} + +fn relaxed() -> FieldValue { + FieldValue::from_static(RELAXED) +} + +fn retained() -> &'static LocationView<'static> { + static FIELD: FieldValue = FieldValue::from_static(SIMPLE); + static VIEW: OnceLock> = OnceLock::new(); + VIEW.get_or_init(|| ::decode_view(FIELD.as_field_value_ref()).expect("valid fixture")) +} + +fn retained_relaxed() -> &'static LocationView<'static> { + static FIELD: FieldValue = FieldValue::from_static(RELAXED); + static VIEW: OnceLock> = OnceLock::new(); + VIEW.get_or_init(|| { + ::decode_view_with(FIELD.as_field_value_ref(), DecodeMode::Relaxed).expect("valid fixture") + }) +} + +fn owned() -> LocationOwned { + LocationOwned::try_from(SIMPLE).expect("valid fixture") +} + +fn redirect_headers() -> HeaderMap { + let mut headers = HeaderMap::new(); + headers.insert("location", HeaderValue::from_static(SIMPLE)); + headers +} + +fn old_decode(value: &FieldValue) -> Result<&str, DecodeError> { + let invalid = || DecodeError::new(&FieldName::Location, DecodeErrorKind::InvalidSyntax); + if let Some(text) = http_headers_simd::as_simple_uri_reference(value.as_bytes()) { + return Ok(text); + } + let text = std::str::from_utf8(value.as_bytes()).map_err(|_invalid| invalid())?; + i32::try_from(text.len()).map_err(|_invalid| invalid())?; + Uri::parse(text).map_err(|_invalid| invalid())?; + Ok(text) +} + +fn observe(uri: UriReference<'_>) -> usize { + uri.scheme().map_or(0, str::len) + + uri.authority().map_or(0, |authority| { + authority.userinfo().map_or(0, str::len) + authority.host().len() + authority.port().map_or(0, str::len) + }) + + uri.path().len() + + uri.query().map_or(0, str::len) + + uri.fragment().map_or(0, str::len) +} + +fn observe_caller(uri: &Uri<&str>) -> usize { + uri.scheme().map_or(0, |scheme| scheme.as_str().len()) + + uri.authority().map_or(0, |authority| { + authority.userinfo().map_or(0, |userinfo| userinfo.as_str().len()) + + authority.host().as_str().len() + + authority.port().map_or(0, str::len) + }) + + uri.path().as_str().len() + + uri.query().map_or(0, |query| query.as_str().len()) + + uri.fragment().map_or(0, |fragment| fragment.as_str().len()) +} + +fn structured_reads(value: FieldValue) -> (usize, FieldValue) { + let view = ::decode_view(black_box(value.as_field_value_ref())).expect("valid fixture"); + let uri = view.uri_reference(); + let sum = (0..READS).map(|_| black_box(observe(black_box(uri)))).sum(); + (sum, value) +} + +fn caller_reads(value: FieldValue) -> (usize, FieldValue) { + let text = old_decode(black_box(&value)).expect("valid fixture"); + let uri = Uri::parse(black_box(text)).expect("validated reference"); + let sum = (0..READS).map(|_| black_box(observe_caller(black_box(&uri)))).sum(); + (sum, value) +} + +#[metabench::benchmark(DECODE_SIMPLE, GROUP, "decode_simple", gungraun_setup = simple)] +fn decode_simple(value: FieldValue) -> (usize, FieldValue) { + let view = ::decode_view(black_box(value.as_field_value_ref())).expect("valid fixture"); + let length = black_box(view).as_bytes().len(); + (length, value) +} + +#[metabench::benchmark(DECODE_GENERAL, GROUP, "decode_general", gungraun_setup = general)] +fn decode_general(value: FieldValue) -> (usize, FieldValue) { + decode_simple(value) +} + +#[metabench::benchmark(DECODE_RELAXED, GROUP, "decode_relaxed", gungraun_setup = relaxed)] +fn decode_relaxed(value: FieldValue) -> (usize, FieldValue) { + let view = ::decode_view_with(black_box(value.as_field_value_ref()), DecodeMode::Relaxed) + .expect("valid fixture"); + let length = black_box(view).as_bytes().len(); + (length, value) +} + +#[metabench::benchmark(STRUCTURED_1, GROUP, "structured_1", gungraun_setup = simple)] +fn structured_1(value: FieldValue) -> (usize, FieldValue) { + structured_reads::<1>(value) +} + +#[metabench::benchmark(STRUCTURED_2, GROUP, "structured_2", gungraun_setup = simple)] +fn structured_2(value: FieldValue) -> (usize, FieldValue) { + structured_reads::<2>(value) +} + +#[metabench::benchmark(STRUCTURED_8, GROUP, "structured_8", gungraun_setup = simple)] +fn structured_8(value: FieldValue) -> (usize, FieldValue) { + structured_reads::<8>(value) +} + +#[metabench::benchmark(CALLER_1, GROUP, "caller_parse_1", gungraun_setup = simple)] +fn caller_1(value: FieldValue) -> (usize, FieldValue) { + caller_reads::<1>(value) +} + +#[metabench::benchmark(CALLER_2, GROUP, "caller_parse_2", gungraun_setup = simple)] +fn caller_2(value: FieldValue) -> (usize, FieldValue) { + caller_reads::<2>(value) +} + +#[metabench::benchmark(CALLER_8, GROUP, "caller_parse_8", gungraun_setup = simple)] +fn caller_8(value: FieldValue) -> (usize, FieldValue) { + caller_reads::<8>(value) +} + +#[metabench::benchmark(STRUCTURED_GENERAL_8, GROUP, "structured_general_8", gungraun_setup = general)] +fn structured_general_8(value: FieldValue) -> (usize, FieldValue) { + structured_reads::<8>(value) +} + +#[metabench::benchmark(CALLER_GENERAL_8, GROUP, "caller_general_8", gungraun_setup = general)] +fn caller_general_8(value: FieldValue) -> (usize, FieldValue) { + caller_reads::<8>(value) +} + +#[metabench::benchmark(STRUCTURED_RELAXED_8, GROUP, "structured_relaxed_8", gungraun_setup = relaxed)] +fn structured_relaxed_8(value: FieldValue) -> (usize, FieldValue) { + let view = ::decode_view_with(black_box(value.as_field_value_ref()), DecodeMode::Relaxed) + .expect("valid fixture"); + let uri = view.uri_reference(); + let sum = (0..8).map(|_| black_box(observe(black_box(uri)))).sum(); + (sum, value) +} + +#[metabench::benchmark(RETAINED_1, GROUP, "retained_1", gungraun_setup = retained)] +fn retained_1(value: &'static LocationView<'static>) -> usize { + observe(black_box(value.uri_reference())) +} + +#[metabench::benchmark(RETAINED_2, GROUP, "retained_2", gungraun_setup = retained)] +fn retained_2(value: &'static LocationView<'static>) -> usize { + let uri = value.uri_reference(); + (0..2).map(|_| black_box(observe(black_box(uri)))).sum() +} + +#[metabench::benchmark(RETAINED_8, GROUP, "retained_8", gungraun_setup = retained)] +fn retained_8(value: &'static LocationView<'static>) -> usize { + let uri = value.uri_reference(); + (0..8).map(|_| black_box(observe(black_box(uri)))).sum() +} + +#[metabench::benchmark(RETAINED_RELAXED_8, GROUP, "retained_relaxed_8", gungraun_setup = retained_relaxed)] +fn retained_relaxed_8(value: &'static LocationView<'static>) -> usize { + retained_8(value) +} + +#[metabench::benchmark(RETAINED_OWNED_8, GROUP, "retained_owned_8", gungraun_setup = owned)] +fn retained_owned_8(value: LocationOwned) -> (usize, LocationOwned) { + let uri = value.uri_reference(); + let sum = (0..8).map(|_| black_box(observe(black_box(uri)))).sum(); + (sum, value) +} + +#[metabench::benchmark(INSPECT_FORWARD, GROUP, "inspect_forward", gungraun_setup = redirect_headers)] +fn inspect_forward(headers: HeaderMap) -> (usize, HeaderMap, HeaderMap) { + let view = Location::view(black_box(&headers)) + .expect("valid fixture") + .expect("present fixture"); + let uri = view.uri_reference(); + let observed = observe(black_box(uri)); + let mut forwarded = HeaderMap::new(); + if black_box(uri.path().starts_with("/docs") && uri.query().is_some()) { + view.insert_into(&mut forwarded).expect("valid forwarding"); + } + (observed, headers, forwarded) +} + +fn criterion_benchmarks(criterion: &mut Criterion) { + // Only call uninstrumented helpers here: metabench measures the first + // instrumented invocation in an allocation worker. + assert_eq!(caller_reads::<1>(simple()).0, structured_reads::<1>(simple()).0); + assert_eq!(caller_reads::<2>(simple()).0, structured_reads::<2>(simple()).0); + assert_eq!(caller_reads::<8>(simple()).0, structured_reads::<8>(simple()).0); + assert_eq!(caller_reads::<8>(general()).0, structured_reads::<8>(general()).0); + assert_eq!(observe(retained_relaxed().uri_reference()), observe(retained().uri_reference())); + + let mut group = criterion.benchmark_group(DECODE_SIMPLE.group_name()); + macro_rules! case { + ($id:ident, $setup:ident, $function:ident) => { + group.bench_function($id.benchmark_name(), |bencher| { + bencher.iter_batched($setup, $function, BatchSize::SmallInput); + }); + }; + } + case!(DECODE_SIMPLE, simple, decode_simple); + case!(DECODE_GENERAL, general, decode_general); + case!(DECODE_RELAXED, relaxed, decode_relaxed); + case!(STRUCTURED_1, simple, structured_1); + case!(STRUCTURED_2, simple, structured_2); + case!(STRUCTURED_8, simple, structured_8); + case!(CALLER_1, simple, caller_1); + case!(CALLER_2, simple, caller_2); + case!(CALLER_8, simple, caller_8); + case!(STRUCTURED_GENERAL_8, general, structured_general_8); + case!(CALLER_GENERAL_8, general, caller_general_8); + case!(STRUCTURED_RELAXED_8, relaxed, structured_relaxed_8); + case!(RETAINED_1, retained, retained_1); + case!(RETAINED_2, retained, retained_2); + case!(RETAINED_8, retained, retained_8); + case!(RETAINED_RELAXED_8, retained_relaxed, retained_relaxed_8); + case!(RETAINED_OWNED_8, owned, retained_owned_8); + case!(INSPECT_FORWARD, redirect_headers, inspect_forward); + group.finish(); +} + +metabench::main!( + criterion = criterion_benchmarks, + benchmarks = [ + DECODE_SIMPLE, + DECODE_GENERAL, + DECODE_RELAXED, + STRUCTURED_1, + STRUCTURED_2, + STRUCTURED_8, + CALLER_1, + CALLER_2, + CALLER_8, + STRUCTURED_GENERAL_8, + CALLER_GENERAL_8, + STRUCTURED_RELAXED_8, + RETAINED_1, + RETAINED_2, + RETAINED_8, + RETAINED_RELAXED_8, + RETAINED_OWNED_8, + INSPECT_FORWARD, + ] +); diff --git a/crates/http_headers/benches/http_headers_micro.rs b/crates/http_headers/benches/http_headers_micro.rs new file mode 100644 index 000000000..c9af3f563 --- /dev/null +++ b/crates/http_headers/benches/http_headers_micro.rs @@ -0,0 +1,3078 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Unified microbenchmarks for `http_headers` beside `headers`. +//! +//! # How a pair is formed +//! +//! Every case is two functions sharing one `#[bench]` id, in a group carrying +//! `compare_by_id = true`, so each engine reports the pair under one case id. +//! The function name is the id followed by the arm: +//! `user_agent_borrowed_http_headers` and +//! `user_agent_borrowed_headers` are the two arms of `user_agent_borrowed`. +//! +//! Giving every case its own function, rather than one function switching on +//! an argument, keeps the dispatch out of the measured region: what Callgrind +//! attributes to a case is the operation and nothing else. +//! +//! # Fairness +//! +//! Both arms of a pair receive the same prebuilt [`HeaderMap`] from a setup +//! function, and setup and teardown run outside the measured region. Both end +//! by calling the same `expect_bytes`/`expect_usize` assertion against the same +//! expected answer, so a decoder that skipped work fails instead of posting a +//! better number. +//! +//! Where the crates cannot do the same thing the case says so rather than +//! inventing an equivalence. `headers` has no borrowed view, so +//! `user_agent_borrowed` races a borrowed lookup against an owned one and is +//! published as a capability difference; `user_agent_owned` is the +//! like-for-like row. `headers::SetCookie` exposes no accessor at all, so its +//! arm can only decode where ours decodes *and* reads every cookie. +//! +//! # Reading these numbers +//! +//! Metabench combines Criterion time, Gungraun instruction counts, and +//! allocation measurements. Results are comparable within this binary and not +//! across files, because the optimizer's inlining decisions inside a measured +//! region depend on the rest of the binary. + +use std::hint::black_box; +use std::str::{self, FromStr as _}; +use std::sync::LazyLock; +use std::time::Duration; + +use criterion::measurement::WallTime; +use criterion::{BatchSize, BenchmarkGroup, BenchmarkId, Criterion}; +use headers::HeaderMapExt as TheirMapExt; +use http::HeaderMap; +use http::header::{AUTHORIZATION, CACHE_CONTROL, CONTENT_TYPE, SET_COOKIE, USER_AGENT}; +use http_headers::headers::{ + Authorization, AuthorizationOwned, Basic, BasicCredentials, Bearer, CacheControl, CacheControlOwned, ContentType, ContentTypeOwned, + ContentTypeView, SetCookie, SetCookieOwned, UserAgent, UserAgentOwned, +}; +use http_headers::sink::{EncodedValues, FieldSink, InsertError}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldValue, FieldValueRef}; + +#[path = "http_headers_fixtures.rs"] +mod fixtures; + +use fixtures::{ + BASIC_PASSWORD, BASIC_USERNAME, BEARER_TOKEN, CHARSET, SET_COOKIE_VALUES, USER_AGENT_VALUE, expect_bytes, expect_inserted, + expect_usize, field_value, value, +}; + +// ── corpus, all of it built outside every measured region ──────────────────── + +/// A `Content-Type` carrying several parameters, one of them quoted. +const CONTENT_TYPE_PARAMETERS: &[u8] = b"multipart/form-data; boundary=------------------------1a2b3c; charset=utf-8; name=\"upload\""; + +/// The base64 payload of the corpus `Basic` credentials. +const BASIC_ENCODED: &[u8] = b"YWxhZGRpbjpvcGVuc2VzYW1l"; + +/// `boundary` + `charset` + `name`, the parameter names both crates report. +const PARAMETER_NAME_BYTES: usize = 19; + +const COOKIE_BYTES_1: usize = SET_COOKIE_VALUES[0].len(); +const COOKIE_BYTES_4: usize = + SET_COOKIE_VALUES[0].len() + SET_COOKIE_VALUES[1].len() + SET_COOKIE_VALUES[2].len() + SET_COOKIE_VALUES[3].len(); +const COOKIE_BYTES_12: usize = 3 * COOKIE_BYTES_4; + +const MAX_AGE: Duration = Duration::from_hours(1); +const CREDENTIAL_ROUNDS: usize = 8; + +static JSON_REQUEST: LazyLock = LazyLock::new(fixtures::json_request); +static BASIC_REQUEST: LazyLock = LazyLock::new(fixtures::basic_auth_request); +static DEGRADED_REQUEST: LazyLock = LazyLock::new(fixtures::degraded_request); +static CACHE_CONTROL_MULTI: LazyLock = LazyLock::new(fixtures::cache_control_multi_line); +static ADVERSARIAL_REQUEST: LazyLock = LazyLock::new(fixtures::adversarial_request); +static COOKIES_1: LazyLock = LazyLock::new(|| cookies(1)); +static COOKIES_4: LazyLock = LazyLock::new(|| cookies(4)); +static COOKIES_12: LazyLock = LazyLock::new(|| cookies(12)); +static CONTENT_TYPE_MANY: LazyLock = LazyLock::new(|| { + let mut map = HeaderMap::with_capacity(2); + let _ = map.insert(CONTENT_TYPE, value(CONTENT_TYPE_PARAMETERS)); + map +}); +static LARGE_BASIC: LazyLock = LazyLock::new(large_basic_map); +/// One directive, so the `list` rows span 1, 3, and 34 and the fixed part of a +/// decode can be separated from its per-directive part by measurement. +static CACHE_CONTROL_ONE: LazyLock = LazyLock::new(|| { + let mut map = HeaderMap::with_capacity(2); + let _ = map.insert(CACHE_CONTROL, value(b"max-age=3600")); + map +}); + +fn cookies(count: usize) -> HeaderMap { + let mut map = HeaderMap::with_capacity(count.next_power_of_two()); + for index in 0..count { + let _ = map.append(SET_COOKIE, value(SET_COOKIE_VALUES[index % 4])); + } + map +} + +fn large_basic_map() -> HeaderMap { + let mut password = Vec::with_capacity(4096); + while password.len() < 2976 { + password.extend_from_slice(b"0123456789abcdef"); + } + password.truncate(2976); + let credentials = AuthorizationOwned::::basic(b"gateway-service-account", &password).expect("legal credentials"); + let mut map = empty_map(); + expect_inserted(Authorization::::insert(&mut map, credentials)); + map +} + +// ── shared setup ───────────────────────────────────────────────────────────── + +fn json_map() -> &'static HeaderMap { + &JSON_REQUEST +} + +fn bearer_map() -> &'static HeaderMap { + black_box(http_headers_simd::is_token68(BEARER_TOKEN)); + &JSON_REQUEST +} + +fn basic_map() -> &'static HeaderMap { + &BASIC_REQUEST +} + +fn degraded_map() -> &'static HeaderMap { + &DEGRADED_REQUEST +} + +fn cache_one_map() -> &'static HeaderMap { + &CACHE_CONTROL_ONE +} + +fn cache_multi_map() -> &'static HeaderMap { + &CACHE_CONTROL_MULTI +} + +fn adversarial_map() -> &'static HeaderMap { + &ADVERSARIAL_REQUEST +} + +fn content_type_many_map() -> &'static HeaderMap { + &CONTENT_TYPE_MANY +} + +fn cookies_1() -> &'static HeaderMap { + &COOKIES_1 +} + +fn cookies_4() -> &'static HeaderMap { + &COOKIES_4 +} + +fn cookies_12() -> &'static HeaderMap { + &COOKIES_12 +} + +fn large_map() -> &'static HeaderMap { + &LARGE_BASIC +} + +fn empty_map() -> HeaderMap { + HeaderMap::with_capacity(8) +} + +fn drop_it(value: T) { + drop(value); +} + +/// Keeps an owned result alive past the measured region without letting the +/// optimizer delete the work that produced it. +fn consume(value: T) { + drop(black_box(value)); +} + +// ── group: lookup ──────────────────────────────────────────────────────────── + +#[metabench::benchmark( + USER_AGENT_BORROWED_HTTP_HEADERS, + "lookup", + "user_agent_borrowed_http_headers", + gungraun_setup = json_map, +)] +#[bench::user_agent_borrowed()] +fn user_agent_borrowed_http_headers(map: &'static HeaderMap) -> usize { + let view = UserAgent::view(map).expect("valid user agent").expect("present user agent"); + expect_bytes(view.as_bytes(), USER_AGENT_VALUE) +} + +#[metabench::benchmark( + USER_AGENT_BORROWED_HEADERS, + "lookup", + "user_agent_borrowed_headers", + gungraun_setup = json_map, +)] +#[bench::user_agent_borrowed()] +fn user_agent_borrowed_headers(map: &'static HeaderMap) -> usize { + let agent = TheirMapExt::typed_try_get::(map) + .expect("valid user agent") + .expect("present user agent"); + let length = expect_bytes(agent.as_str().as_bytes(), USER_AGENT_VALUE); + consume(agent); + length +} + +#[metabench::benchmark( + USER_AGENT_OWNED_HTTP_HEADERS, + "lookup", + "user_agent_owned_http_headers", + gungraun_setup = json_map, +)] +#[bench::user_agent_owned()] +fn user_agent_owned_http_headers(map: &'static HeaderMap) -> usize { + let agent = UserAgent::owned(map).expect("valid user agent").expect("present user agent"); + let length = expect_bytes(agent.as_bytes(), USER_AGENT_VALUE); + consume(agent); + length +} + +#[metabench::benchmark( + USER_AGENT_OWNED_HEADERS, + "lookup", + "user_agent_owned_headers", + gungraun_setup = json_map, +)] +#[bench::user_agent_owned()] +fn user_agent_owned_headers(map: &'static HeaderMap) -> usize { + let agent = TheirMapExt::typed_try_get::(map) + .expect("valid user agent") + .expect("present user agent"); + let length = expect_bytes(agent.as_str().as_bytes(), USER_AGENT_VALUE); + consume(agent); + length +} + +#[metabench::benchmark( + USER_AGENT_ABSENT_HTTP_HEADERS, + "lookup", + "user_agent_absent_http_headers", + gungraun_setup = cookies_4, +)] +#[bench::user_agent_absent()] +fn user_agent_absent_http_headers(map: &'static HeaderMap) -> usize { + let absent = UserAgent::view(map).expect("absent is not an error"); + expect_usize(usize::from(absent.is_none()), 1) +} + +#[metabench::benchmark( + USER_AGENT_ABSENT_HEADERS, + "lookup", + "user_agent_absent_headers", + gungraun_setup = cookies_4, +)] +#[bench::user_agent_absent()] +fn user_agent_absent_headers(map: &'static HeaderMap) -> usize { + let absent = TheirMapExt::typed_try_get::(map).expect("absent is not an error"); + expect_usize(usize::from(absent.is_none()), 1) +} + +// ── group: repeated ────────────────────────────────────────────────────────── + +fn read_cookies(map: &HeaderMap, expected: usize) -> usize { + let view = SetCookie::view(map).expect("valid cookies").expect("present cookies"); + let mut total = 0; + for cookie in view.iter() { + total += cookie.as_bytes().len(); + } + expect_usize(total, expected) +} + +/// `headers::SetCookie` has no accessor, so its only measurable operation is +/// the decode, which clones every field value into a `Vec`. +fn decode_cookies(map: &HeaderMap) -> usize { + let cookies = TheirMapExt::typed_try_get::(map) + .expect("valid cookies") + .expect("present cookies"); + consume(cookies); + 1 +} + +#[metabench::benchmark( + SET_COOKIE_BORROWED_1_HTTP_HEADERS, + "repeated", + "set_cookie_borrowed_1_http_headers", + gungraun_setup = cookies_1, +)] +#[bench::set_cookie_borrowed_1()] +fn set_cookie_borrowed_1_http_headers(map: &'static HeaderMap) -> usize { + read_cookies(map, COOKIE_BYTES_1) +} + +#[metabench::benchmark( + SET_COOKIE_BORROWED_1_HEADERS, + "repeated", + "set_cookie_borrowed_1_headers", + gungraun_setup = cookies_1, +)] +#[bench::set_cookie_borrowed_1()] +fn set_cookie_borrowed_1_headers(map: &'static HeaderMap) -> usize { + decode_cookies(map) +} + +#[metabench::benchmark( + SET_COOKIE_BORROWED_4_HTTP_HEADERS, + "repeated", + "set_cookie_borrowed_4_http_headers", + gungraun_setup = cookies_4, +)] +#[bench::set_cookie_borrowed_4()] +fn set_cookie_borrowed_4_http_headers(map: &'static HeaderMap) -> usize { + read_cookies(map, COOKIE_BYTES_4) +} + +#[metabench::benchmark( + SET_COOKIE_BORROWED_4_HEADERS, + "repeated", + "set_cookie_borrowed_4_headers", + gungraun_setup = cookies_4, +)] +#[bench::set_cookie_borrowed_4()] +fn set_cookie_borrowed_4_headers(map: &'static HeaderMap) -> usize { + decode_cookies(map) +} + +#[metabench::benchmark( + SET_COOKIE_BORROWED_12_HTTP_HEADERS, + "repeated", + "set_cookie_borrowed_12_http_headers", + gungraun_setup = cookies_12, +)] +#[bench::set_cookie_borrowed_12()] +fn set_cookie_borrowed_12_http_headers(map: &'static HeaderMap) -> usize { + read_cookies(map, COOKIE_BYTES_12) +} + +#[metabench::benchmark( + SET_COOKIE_BORROWED_12_HEADERS, + "repeated", + "set_cookie_borrowed_12_headers", + gungraun_setup = cookies_12, +)] +#[bench::set_cookie_borrowed_12()] +fn set_cookie_borrowed_12_headers(map: &'static HeaderMap) -> usize { + decode_cookies(map) +} + +#[metabench::benchmark( + SET_COOKIE_OWNED_4_HTTP_HEADERS, + "repeated", + "set_cookie_owned_4_http_headers", + gungraun_setup = cookies_4, +)] +#[bench::set_cookie_owned_4()] +fn set_cookie_owned_4_http_headers(map: &'static HeaderMap) -> usize { + let cookies = SetCookie::owned(map).expect("valid cookies").expect("present cookies"); + let mut total = 0; + for cookie in &cookies { + total += cookie.as_bytes().len(); + } + let total = expect_usize(total, COOKIE_BYTES_4); + consume(cookies); + total +} + +#[metabench::benchmark( + SET_COOKIE_OWNED_4_HEADERS, + "repeated", + "set_cookie_owned_4_headers", + gungraun_setup = cookies_4, +)] +#[bench::set_cookie_owned_4()] +fn set_cookie_owned_4_headers(map: &'static HeaderMap) -> usize { + decode_cookies(map) +} + +// ── group: list ────────────────────────────────────────────────────────────── + +#[metabench::benchmark( + CACHE_CONTROL_SINGLE_LINE_HTTP_HEADERS, + "list", + "cache_control_single_line_http_headers", + gungraun_setup = json_map, +)] +#[bench::cache_control_single_line()] +fn cache_control_single_line_http_headers(map: &'static HeaderMap) -> usize { + let view = CacheControl::view(map) + .expect("valid cache control") + .expect("present cache control"); + let answer = usize::from(view.max_age() == Some(MAX_AGE)) + usize::from(!view.no_cache()); + expect_usize(answer, 2) +} + +#[metabench::benchmark( + CACHE_CONTROL_SINGLE_LINE_HEADERS, + "list", + "cache_control_single_line_headers", + gungraun_setup = json_map, +)] +#[bench::cache_control_single_line()] +fn cache_control_single_line_headers(map: &'static HeaderMap) -> usize { + let control = TheirMapExt::typed_try_get::(map) + .expect("valid cache control") + .expect("present cache control"); + let answer = usize::from(control.max_age() == Some(MAX_AGE)) + usize::from(!control.no_cache()); + let answer = expect_usize(answer, 2); + consume(control); + answer +} + +#[metabench::benchmark( + CACHE_CONTROL_MULTI_LINE_HTTP_HEADERS, + "list", + "cache_control_multi_line_http_headers", + gungraun_setup = cache_multi_map, +)] +#[bench::cache_control_multi_line()] +fn cache_control_multi_line_http_headers(map: &'static HeaderMap) -> usize { + let view = CacheControl::view(map) + .expect("valid cache control") + .expect("present cache control"); + let answer = usize::from(view.max_age() == Some(MAX_AGE)) + view.directives().count(); + expect_usize(answer, 6) +} + +#[metabench::benchmark( + CACHE_CONTROL_MULTI_LINE_HEADERS, + "list", + "cache_control_multi_line_headers", + gungraun_setup = cache_multi_map, +)] +#[bench::cache_control_multi_line()] +fn cache_control_multi_line_headers(map: &'static HeaderMap) -> usize { + let control = TheirMapExt::typed_try_get::(map) + .expect("valid cache control") + .expect("present cache control"); + let recognized = usize::from(control.public()) + + usize::from(control.must_revalidate()) + + usize::from(control.no_transform()) + + usize::from(control.s_max_age().is_some()) + + usize::from(control.max_age().is_some()); + let answer = usize::from(control.max_age() == Some(MAX_AGE)) + recognized; + let answer = expect_usize(answer, 6); + consume(control); + answer +} + +#[metabench::benchmark( + CACHE_CONTROL_OWNED_HTTP_HEADERS, + "list", + "cache_control_owned_http_headers", + gungraun_setup = json_map, +)] +#[bench::cache_control_owned()] +fn cache_control_owned_http_headers(map: &'static HeaderMap) -> usize { + let control = CacheControl::owned(map) + .expect("valid cache control") + .expect("present cache control"); + let answer = usize::from(control.max_age() == Some(MAX_AGE)) + usize::from(!control.no_cache()); + let answer = expect_usize(answer, 2); + consume(control); + answer +} + +#[metabench::benchmark( + CACHE_CONTROL_OWNED_HEADERS, + "list", + "cache_control_owned_headers", + gungraun_setup = json_map, +)] +#[bench::cache_control_owned()] +fn cache_control_owned_headers(map: &'static HeaderMap) -> usize { + let control = TheirMapExt::typed_try_get::(map) + .expect("valid cache control") + .expect("present cache control"); + let answer = usize::from(control.max_age() == Some(MAX_AGE)) + usize::from(!control.no_cache()); + let answer = expect_usize(answer, 2); + consume(control); + answer +} + +#[metabench::benchmark( + CACHE_CONTROL_ADVERSARIAL_HTTP_HEADERS, + "list", + "cache_control_adversarial_http_headers", + gungraun_setup = adversarial_map, +)] +#[bench::cache_control_adversarial()] +fn cache_control_adversarial_http_headers(map: &'static HeaderMap) -> usize { + let view = CacheControl::view(map) + .expect("valid cache control") + .expect("present cache control"); + expect_usize(usize::from(view.max_age() == Some(MAX_AGE)), 1) +} + +#[metabench::benchmark( + CACHE_CONTROL_ADVERSARIAL_HEADERS, + "list", + "cache_control_adversarial_headers", + gungraun_setup = adversarial_map, +)] +#[bench::cache_control_adversarial()] +fn cache_control_adversarial_headers(map: &'static HeaderMap) -> usize { + let control = TheirMapExt::typed_try_get::(map) + .expect("valid cache control") + .expect("present cache control"); + let answer = expect_usize(usize::from(control.max_age() == Some(MAX_AGE)), 1); + consume(control); + answer +} + +#[metabench::benchmark( + CACHE_CONTROL_ONE_DIRECTIVE_HTTP_HEADERS, + "list", + "cache_control_one_directive_http_headers", + gungraun_setup = cache_one_map, +)] +#[bench::cache_control_one_directive()] +fn cache_control_one_directive_http_headers(map: &'static HeaderMap) -> usize { + let view = CacheControl::view(map) + .expect("valid cache control") + .expect("present cache control"); + expect_usize(usize::from(view.max_age() == Some(MAX_AGE)), 1) +} + +#[metabench::benchmark( + CACHE_CONTROL_ONE_DIRECTIVE_HEADERS, + "list", + "cache_control_one_directive_headers", + gungraun_setup = cache_one_map, +)] +#[bench::cache_control_one_directive()] +fn cache_control_one_directive_headers(map: &'static HeaderMap) -> usize { + let control = TheirMapExt::typed_try_get::(map) + .expect("valid cache control") + .expect("present cache control"); + let answer = expect_usize(usize::from(control.max_age() == Some(MAX_AGE)), 1); + consume(control); + answer +} + +#[metabench::benchmark( + CACHE_CONTROL_MAX_AGE_ONLY_HTTP_HEADERS, + "list", + "cache_control_max_age_only_http_headers", + gungraun_setup = json_map, +)] +#[bench::cache_control_max_age_only()] +fn cache_control_max_age_only_http_headers(map: &'static HeaderMap) -> usize { + let view = CacheControl::view(map) + .expect("valid cache control") + .expect("present cache control"); + expect_usize(usize::from(view.max_age() == Some(MAX_AGE)), 1) +} + +#[metabench::benchmark( + CACHE_CONTROL_MAX_AGE_ONLY_HEADERS, + "list", + "cache_control_max_age_only_headers", + gungraun_setup = json_map, +)] +#[bench::cache_control_max_age_only()] +fn cache_control_max_age_only_headers(map: &'static HeaderMap) -> usize { + let control = TheirMapExt::typed_try_get::(map) + .expect("valid cache control") + .expect("present cache control"); + let answer = expect_usize(usize::from(control.max_age() == Some(MAX_AGE)), 1); + consume(control); + answer +} + +// ── group: directives ──────────────────────────────────────────────────────── + +/// `max-age=3600` is the only directive in the single-line corpus with a value. +const SINGLE_LINE_VALUE_BYTES: usize = 4; + +/// One `max-age` plus 32 `ext=value` directives in the adversarial corpus. +const ADVERSARIAL_VALUE_BYTES: usize = 218; + +fn sum_value_bytes(map: &HeaderMap, expected: usize) -> usize { + let view = CacheControl::view(map) + .expect("valid cache control") + .expect("present cache control"); + let mut total = 0; + for directive in view.directives() { + total += directive.value().map_or(0, <[u8]>::len); + } + expect_usize(total, expected) +} + +fn sum_value_str(map: &HeaderMap, expected: usize) -> usize { + let view = CacheControl::view(map) + .expect("valid cache control") + .expect("present cache control"); + let mut total = 0; + for directive in view.directives() { + total += directive.value_str().expect("utf-8 directive value").map_or(0, str::len); + } + expect_usize(total, expected) +} + +#[metabench::benchmark( + DIRECTIVE_VALUE_BYTES, + "directives", + "directive_value_bytes", + gungraun_setup = json_map, +)] +#[bench::directive_value()] +fn directive_value_bytes(map: &'static HeaderMap) -> usize { + sum_value_bytes(map, SINGLE_LINE_VALUE_BYTES) +} + +#[metabench::benchmark( + DIRECTIVE_VALUE_STR, + "directives", + "directive_value_str", + gungraun_setup = json_map, +)] +#[bench::directive_value()] +fn directive_value_str(map: &'static HeaderMap) -> usize { + sum_value_str(map, SINGLE_LINE_VALUE_BYTES) +} + +#[metabench::benchmark( + DIRECTIVE_VALUE_MANY_BYTES, + "directives", + "directive_value_many_bytes", + gungraun_setup = adversarial_map, +)] +#[bench::directive_value_many()] +fn directive_value_many_bytes(map: &'static HeaderMap) -> usize { + sum_value_bytes(map, ADVERSARIAL_VALUE_BYTES) +} + +#[metabench::benchmark( + DIRECTIVE_VALUE_MANY_STR, + "directives", + "directive_value_many_str", + gungraun_setup = adversarial_map, +)] +#[bench::directive_value_many()] +fn directive_value_many_str(map: &'static HeaderMap) -> usize { + sum_value_str(map, ADVERSARIAL_VALUE_BYTES) +} + +// ── group: structured ──────────────────────────────────────────────────────── + +fn inspect_content_type(view: &ContentTypeView<'_>) -> usize { + let mut total = expect_bytes(view.type_().expect("utf-8 type").as_bytes(), b"application"); + total += expect_bytes(view.subtype().expect("utf-8 subtype").as_bytes(), b"json"); + total += expect_bytes( + view.parameter("charset").expect("valid parameters").expect("charset present"), + CHARSET, + ); + expect_usize(total, 20) +} + +fn inspect_mime(mime: &headers::Mime) -> usize { + let mut total = expect_bytes(mime.type_().as_str().as_bytes(), b"application"); + total += expect_bytes(mime.subtype().as_str().as_bytes(), b"json"); + total += expect_bytes(mime.get_param("charset").expect("charset present").as_str().as_bytes(), CHARSET); + expect_usize(total, 20) +} + +#[metabench::benchmark( + CONTENT_TYPE_INSPECT_HTTP_HEADERS, + "structured", + "content_type_inspect_http_headers", + gungraun_setup = json_map, +)] +#[bench::content_type_inspect()] +fn content_type_inspect_http_headers(map: &'static HeaderMap) -> usize { + let view = ContentType::view(map).expect("valid content type").expect("present content type"); + inspect_content_type(&view) +} + +#[metabench::benchmark( + CONTENT_TYPE_INSPECT_HEADERS, + "structured", + "content_type_inspect_headers", + gungraun_setup = json_map, +)] +#[bench::content_type_inspect()] +fn content_type_inspect_headers(map: &'static HeaderMap) -> usize { + let content_type = TheirMapExt::typed_try_get::(map) + .expect("valid content type") + .expect("present content type"); + let mime = headers::Mime::from(content_type); + let total = inspect_mime(&mime); + consume(mime); + total +} + +#[metabench::benchmark( + CONTENT_TYPE_PARAMETERS_HTTP_HEADERS, + "structured", + "content_type_parameters_http_headers", + gungraun_setup = content_type_many_map, +)] +#[bench::content_type_parameters()] +fn content_type_parameters_http_headers(map: &'static HeaderMap) -> usize { + let view = ContentType::view(map).expect("valid content type").expect("present content type"); + let mut total = 0; + for parameter in view.parameters() { + total += parameter.expect("valid parameter").name().len(); + } + let total = expect_usize(total, PARAMETER_NAME_BYTES); + total + + expect_bytes( + view.parameter("charset").expect("valid parameters").expect("charset present"), + CHARSET, + ) +} + +#[metabench::benchmark( + CONTENT_TYPE_PARAMETERS_HEADERS, + "structured", + "content_type_parameters_headers", + gungraun_setup = content_type_many_map, +)] +#[bench::content_type_parameters()] +fn content_type_parameters_headers(map: &'static HeaderMap) -> usize { + let content_type = TheirMapExt::typed_try_get::(map) + .expect("valid content type") + .expect("present content type"); + let mime = headers::Mime::from(content_type); + let mut total = 0; + for (name, _parameter) in mime.params() { + total += name.as_str().len(); + } + let total = expect_usize(total, PARAMETER_NAME_BYTES) + + expect_bytes(mime.get_param("charset").expect("charset present").as_str().as_bytes(), CHARSET); + consume(mime); + total +} + +#[metabench::benchmark( + CONTENT_TYPE_OWNED_HTTP_HEADERS, + "structured", + "content_type_owned_http_headers", + gungraun_setup = json_map, +)] +#[bench::content_type_owned()] +fn content_type_owned_http_headers(map: &'static HeaderMap) -> usize { + let content_type = ContentType::owned(map).expect("valid content type").expect("present content type"); + let mut total = expect_bytes(content_type.type_().expect("utf-8 type").as_bytes(), b"application"); + total += expect_bytes(content_type.subtype().expect("utf-8 subtype").as_bytes(), b"json"); + total += expect_bytes( + content_type + .parameter("charset") + .expect("valid parameters") + .expect("charset present"), + CHARSET, + ); + let total = expect_usize(total, 20); + consume(content_type); + total +} + +#[metabench::benchmark( + CONTENT_TYPE_OWNED_HEADERS, + "structured", + "content_type_owned_headers", + gungraun_setup = json_map, +)] +#[bench::content_type_owned()] +fn content_type_owned_headers(map: &'static HeaderMap) -> usize { + let content_type = TheirMapExt::typed_try_get::(map) + .expect("valid content type") + .expect("present content type"); + let mime = headers::Mime::from(content_type); + let total = inspect_mime(&mime); + consume(mime); + total +} + +#[metabench::benchmark( + CONTENT_TYPE_MALFORMED_HTTP_HEADERS, + "structured", + "content_type_malformed_http_headers", + gungraun_setup = degraded_map, +)] +#[bench::content_type_malformed()] +fn content_type_malformed_http_headers(map: &'static HeaderMap) -> usize { + let rejected = ContentType::view(map).is_err(); + expect_usize(usize::from(rejected), 1) +} + +#[metabench::benchmark( + CONTENT_TYPE_MALFORMED_HEADERS, + "structured", + "content_type_malformed_headers", + gungraun_setup = degraded_map, +)] +#[bench::content_type_malformed()] +fn content_type_malformed_headers(map: &'static HeaderMap) -> usize { + let rejected = TheirMapExt::typed_try_get::(map).is_err(); + expect_usize(usize::from(rejected), 1) +} + +// ── group: authorization ───────────────────────────────────────────────────── + +fn basic_credentials() -> (&'static HeaderMap, BasicCredentials) { + (&BASIC_REQUEST, BasicCredentials::new()) +} + +fn extract_basic<'a>(map: &HeaderMap, output: &'a mut BasicCredentials) -> &'a BasicCredentials { + Authorization::::view(map) + .expect("valid credentials") + .expect("present credentials") + .extract(output) + .expect("valid decoded credentials") +} + +#[metabench::benchmark( + BASIC_DECODE_HTTP_HEADERS, + "authorization", + "basic_decode_http_headers", + gungraun_setup = basic_credentials, + gungraun_teardown = drop_it, +)] +#[bench::basic_decode()] +fn basic_decode_http_headers(state: (&'static HeaderMap, BasicCredentials)) -> (usize, BasicCredentials) { + let (map, mut output) = state; + let total = { + let credentials = extract_basic(map, &mut output); + expect_bytes(credentials.username(), BASIC_USERNAME) + expect_bytes(credentials.password(), BASIC_PASSWORD) + }; + (expect_usize(total, 17), output) +} + +#[metabench::benchmark( + BASIC_DECODE_HEADERS, + "authorization", + "basic_decode_headers", + gungraun_setup = basic_credentials, + gungraun_teardown = drop_it, +)] +#[bench::basic_decode()] +fn basic_decode_headers(state: (&'static HeaderMap, BasicCredentials)) -> (usize, BasicCredentials) { + let (map, output) = state; + let credentials = TheirMapExt::typed_try_get::>(map) + .expect("valid credentials") + .expect("present credentials"); + let total = + expect_bytes(credentials.username().as_bytes(), BASIC_USERNAME) + expect_bytes(credentials.password().as_bytes(), BASIC_PASSWORD); + let total = expect_usize(total, 17); + consume(credentials); + (total, output) +} + +#[metabench::benchmark( + BASIC_VALIDATE_HTTP_HEADERS, + "authorization", + "basic_validate_http_headers", + gungraun_setup = basic_map, +)] +#[bench::basic_validate()] +fn basic_validate_http_headers(map: &'static HeaderMap) -> usize { + let view = Authorization::::view(map) + .expect("valid credentials") + .expect("present credentials"); + expect_bytes(view.credentials(), BASIC_ENCODED) +} + +#[metabench::benchmark( + BASIC_VALIDATE_HEADERS, + "authorization", + "basic_validate_headers", + gungraun_setup = basic_map, +)] +#[bench::basic_validate()] +fn basic_validate_headers(map: &'static HeaderMap) -> usize { + let credentials = TheirMapExt::typed_try_get::>(map) + .expect("valid credentials") + .expect("present credentials"); + let total = expect_usize(credentials.username().len() + credentials.password().len(), 17); + consume(credentials); + total +} + +#[metabench::benchmark( + BEARER_BORROWED_HTTP_HEADERS, + "authorization", + "bearer_borrowed_http_headers", + gungraun_setup = bearer_map, +)] +#[bench::bearer_borrowed()] +fn bearer_borrowed_http_headers(map: &'static HeaderMap) -> usize { + let view = Authorization::::view(map) + .expect("valid credentials") + .expect("present credentials"); + expect_bytes(view.token(), BEARER_TOKEN) +} + +#[metabench::benchmark( + BEARER_BORROWED_HEADERS, + "authorization", + "bearer_borrowed_headers", + gungraun_setup = bearer_map, +)] +#[bench::bearer_borrowed()] +fn bearer_borrowed_headers(map: &'static HeaderMap) -> usize { + let credentials = TheirMapExt::typed_try_get::>(map) + .expect("valid credentials") + .expect("present credentials"); + let total = expect_bytes(credentials.token().as_bytes(), BEARER_TOKEN); + consume(credentials); + total +} + +#[metabench::benchmark( + BASIC_ENCODE_HTTP_HEADERS, + "authorization", + "basic_encode_http_headers", + gungraun_setup = empty_map, + gungraun_teardown = drop_it, +)] +#[bench::basic_encode()] +fn basic_encode_http_headers(mut map: HeaderMap) -> (usize, HeaderMap) { + let credentials = AuthorizationOwned::::basic(BASIC_USERNAME, BASIC_PASSWORD).expect("legal credentials"); + expect_inserted(Authorization::::insert(&mut map, credentials)); + let stored = map.get(&AUTHORIZATION).expect("inserted credentials"); + let length = expect_bytes(stored.as_bytes(), fixtures::BASIC_VALUE); + (length, map) +} + +#[metabench::benchmark( + BASIC_ENCODE_HEADERS, + "authorization", + "basic_encode_headers", + gungraun_setup = empty_map, + gungraun_teardown = drop_it, +)] +#[bench::basic_encode()] +fn basic_encode_headers(mut map: HeaderMap) -> (usize, HeaderMap) { + let credentials = headers::Authorization::basic("aladdin", "opensesame"); + TheirMapExt::typed_insert(&mut map, credentials); + let stored = map.get(&AUTHORIZATION).expect("inserted credentials"); + let length = expect_bytes(stored.as_bytes(), fixtures::BASIC_VALUE); + (length, map) +} + +// ── group: insertion ───────────────────────────────────────────────────────── + +fn our_user_agent() -> (HeaderMap, UserAgentOwned) { + ( + empty_map(), + UserAgentOwned::try_from(field_value(USER_AGENT_VALUE)).expect("legal user agent"), + ) +} + +fn their_user_agent() -> (HeaderMap, headers::UserAgent) { + let text = str::from_utf8(USER_AGENT_VALUE).expect("utf-8 corpus"); + (empty_map(), headers::UserAgent::from_str(text).expect("legal user agent")) +} + +#[metabench::benchmark( + INSERT_USER_AGENT_HTTP_HEADERS, + "insertion", + "insert_user_agent_http_headers", + gungraun_setup = our_user_agent, + gungraun_teardown = drop_it, +)] +#[bench::insert_user_agent()] +fn insert_user_agent_http_headers(state: (HeaderMap, UserAgentOwned)) -> (usize, HeaderMap) { + let (mut map, agent) = state; + expect_inserted(UserAgent::insert(&mut map, agent)); + let stored = map.get(&USER_AGENT).expect("inserted user agent"); + let length = expect_bytes(stored.as_bytes(), USER_AGENT_VALUE); + (length, map) +} + +#[metabench::benchmark( + INSERT_USER_AGENT_HEADERS, + "insertion", + "insert_user_agent_headers", + gungraun_setup = their_user_agent, + gungraun_teardown = drop_it, +)] +#[bench::insert_user_agent()] +fn insert_user_agent_headers(state: (HeaderMap, headers::UserAgent)) -> (usize, HeaderMap) { + let (mut map, agent) = state; + TheirMapExt::typed_insert(&mut map, agent); + let stored = map.get(&USER_AGENT).expect("inserted user agent"); + let length = expect_bytes(stored.as_bytes(), USER_AGENT_VALUE); + (length, map) +} + +fn our_content_type() -> (HeaderMap, ContentTypeOwned) { + (empty_map(), ContentTypeOwned::json()) +} + +fn their_content_type() -> (HeaderMap, headers::ContentType) { + (empty_map(), headers::ContentType::json()) +} + +#[metabench::benchmark( + INSERT_CONTENT_TYPE_HTTP_HEADERS, + "insertion", + "insert_content_type_http_headers", + gungraun_setup = our_content_type, + gungraun_teardown = drop_it, +)] +#[bench::insert_content_type()] +fn insert_content_type_http_headers(state: (HeaderMap, ContentTypeOwned)) -> (usize, HeaderMap) { + let (mut map, content_type) = state; + expect_inserted(ContentType::insert(&mut map, content_type)); + let stored = map.get(&CONTENT_TYPE).expect("inserted content type"); + let length = expect_bytes(stored.as_bytes(), b"application/json"); + (length, map) +} + +#[metabench::benchmark( + INSERT_CONTENT_TYPE_HEADERS, + "insertion", + "insert_content_type_headers", + gungraun_setup = their_content_type, + gungraun_teardown = drop_it, +)] +#[bench::insert_content_type()] +fn insert_content_type_headers(state: (HeaderMap, headers::ContentType)) -> (usize, HeaderMap) { + let (mut map, content_type) = state; + TheirMapExt::typed_insert(&mut map, content_type); + let stored = map.get(&CONTENT_TYPE).expect("inserted content type"); + let length = expect_bytes(stored.as_bytes(), b"application/json"); + (length, map) +} + +fn our_cache_control() -> (HeaderMap, CacheControlOwned) { + let control = CacheControlOwned::builder() + .max_age(MAX_AGE) + .private() + .build() + .expect("legal directives"); + (empty_map(), control) +} + +fn their_cache_control() -> (HeaderMap, headers::CacheControl) { + let control = headers::CacheControl::new().with_max_age(MAX_AGE).with_private(); + (empty_map(), control) +} + +#[metabench::benchmark( + INSERT_CACHE_CONTROL_HTTP_HEADERS, + "insertion", + "insert_cache_control_http_headers", + gungraun_setup = our_cache_control, + gungraun_teardown = drop_it, +)] +#[bench::insert_cache_control()] +fn insert_cache_control_http_headers(state: (HeaderMap, CacheControlOwned)) -> (usize, HeaderMap) { + let (mut map, control) = state; + expect_inserted(CacheControl::insert(&mut map, control)); + let stored = map.get(&CACHE_CONTROL).expect("inserted cache control"); + let length = expect_bytes(stored.as_bytes(), b"max-age=3600, private"); + (length, map) +} + +#[metabench::benchmark( + INSERT_CACHE_CONTROL_HEADERS, + "insertion", + "insert_cache_control_headers", + gungraun_setup = their_cache_control, + gungraun_teardown = drop_it, +)] +#[bench::insert_cache_control()] +fn insert_cache_control_headers(state: (HeaderMap, headers::CacheControl)) -> (usize, HeaderMap) { + let (mut map, control) = state; + TheirMapExt::typed_insert(&mut map, control); + let stored = map.get(&CACHE_CONTROL).expect("inserted cache control"); + let length = expect_bytes(stored.as_bytes(), b"private, max-age=3600"); + (length, map) +} + +// ── group: ownership ───────────────────────────────────────────────────────── + +fn our_cookies() -> (HeaderMap, SetCookieOwned) { + let mut cookies = SetCookieOwned::new(); + for cookie in SET_COOKIE_VALUES { + cookies.push(field_value(cookie)).expect("legal cookie"); + } + (empty_map(), cookies) +} + +#[metabench::benchmark( + OWNERSHIP_USER_AGENT_MOVED, + "ownership", + "ownership_user_agent_moved", + gungraun_setup = our_user_agent, + gungraun_teardown = drop_it, +)] +#[bench::ownership_user_agent()] +fn ownership_user_agent_moved(state: (HeaderMap, UserAgentOwned)) -> (usize, HeaderMap) { + let (mut map, agent) = state; + expect_inserted(UserAgent::insert(&mut map, agent)); + let stored = map.get(&USER_AGENT).expect("inserted user agent"); + let length = expect_bytes(stored.as_bytes(), USER_AGENT_VALUE); + (length, map) +} + +#[metabench::benchmark( + OWNERSHIP_USER_AGENT_CLONED, + "ownership", + "ownership_user_agent_cloned", + gungraun_setup = our_user_agent, + gungraun_teardown = drop_it, +)] +#[bench::ownership_user_agent()] +fn ownership_user_agent_cloned(state: (HeaderMap, UserAgentOwned)) -> (usize, HeaderMap, UserAgentOwned) { + let (mut map, agent) = state; + expect_inserted(UserAgent::insert(&mut map, agent.clone())); + let stored = map.get(&USER_AGENT).expect("inserted user agent"); + let length = expect_bytes(stored.as_bytes(), USER_AGENT_VALUE); + (length, map, agent) +} + +#[metabench::benchmark( + OWNERSHIP_SET_COOKIE_MOVED, + "ownership", + "ownership_set_cookie_moved", + gungraun_setup = our_cookies, + gungraun_teardown = drop_it, +)] +#[bench::ownership_set_cookie()] +fn ownership_set_cookie_moved(state: (HeaderMap, SetCookieOwned)) -> (usize, HeaderMap) { + let (mut map, cookies) = state; + expect_inserted(SetCookie::insert(&mut map, cookies)); + let stored = map.get_all(&SET_COOKIE).iter().count(); + (expect_usize(stored, 4), map) +} + +#[metabench::benchmark( + OWNERSHIP_SET_COOKIE_CLONED, + "ownership", + "ownership_set_cookie_cloned", + gungraun_setup = our_cookies, + gungraun_teardown = drop_it, +)] +#[bench::ownership_set_cookie()] +fn ownership_set_cookie_cloned(state: (HeaderMap, SetCookieOwned)) -> (usize, HeaderMap, SetCookieOwned) { + let (mut map, cookies) = state; + expect_inserted(SetCookie::insert(&mut map, cookies.clone())); + let stored = map.get_all(&SET_COOKIE).iter().count(); + (expect_usize(stored, 4), map, cookies) +} + +// ── group: credentials ─────────────────────────────────────────────────────── + +fn warm_credentials() -> (&'static HeaderMap, BasicCredentials) { + let map: &'static HeaderMap = &BASIC_REQUEST; + let mut output = BasicCredentials::new(); + assert!(!extract_basic(map, &mut output).username().is_empty(), "credentials must decode"); + (map, output) +} + +fn large_warm_credentials() -> (&'static HeaderMap, BasicCredentials) { + let map: &'static HeaderMap = &LARGE_BASIC; + let mut output = BasicCredentials::new(); + assert!(!extract_basic(map, &mut output).password().is_empty(), "credentials must decode"); + (map, output) +} + +fn high_water_warm_credentials() -> (&'static HeaderMap, BasicCredentials) { + let mut output = BasicCredentials::new(); + assert!( + !extract_basic(&LARGE_BASIC, &mut output).password().is_empty(), + "credentials must decode" + ); + (&BASIC_REQUEST, output) +} + +fn decode_rounds(map: &HeaderMap, output: &mut BasicCredentials) -> usize { + let mut total = 0; + for _ in 0..CREDENTIAL_ROUNDS { + let credentials = extract_basic(map, output); + total += credentials.username().len() + credentials.password().len(); + } + total +} + +fn decode_rounds_fresh(map: &HeaderMap) -> usize { + let mut total = 0; + for _ in 0..CREDENTIAL_ROUNDS { + let mut output = BasicCredentials::new(); + let credentials = extract_basic(map, &mut output); + total += credentials.username().len() + credentials.password().len(); + } + total +} + +#[metabench::benchmark( + CREDENTIALS_BASIC_REUSED, + "credentials", + "credentials_basic_reused", + gungraun_setup = warm_credentials, + gungraun_teardown = drop_it, +)] +#[bench::credentials_basic()] +fn credentials_basic_reused(state: (&'static HeaderMap, BasicCredentials)) -> (usize, BasicCredentials) { + let (map, mut output) = state; + let total = decode_rounds(map, &mut output); + (expect_usize(total, 17 * CREDENTIAL_ROUNDS), output) +} + +#[metabench::benchmark( + CREDENTIALS_BASIC_HIGH_WATER_REUSED, + "credentials", + "credentials_basic_high_water_reused", + gungraun_setup = high_water_warm_credentials, + gungraun_teardown = drop_it, +)] +#[bench::credentials_basic_high_water()] +fn credentials_basic_high_water_reused(state: (&'static HeaderMap, BasicCredentials)) -> (usize, BasicCredentials) { + let (map, mut output) = state; + let total = decode_rounds(map, &mut output); + (expect_usize(total, 17 * CREDENTIAL_ROUNDS), output) +} + +#[metabench::benchmark( + CREDENTIALS_BASIC_FRESH, + "credentials", + "credentials_basic_fresh", + gungraun_setup = basic_map, +)] +#[bench::credentials_basic()] +fn credentials_basic_fresh(map: &'static HeaderMap) -> usize { + expect_usize(decode_rounds_fresh(map), 17 * CREDENTIAL_ROUNDS) +} + +#[metabench::benchmark( + CREDENTIALS_LARGE_REUSED, + "credentials", + "credentials_large_reused", + gungraun_setup = large_warm_credentials, + gungraun_teardown = drop_it, +)] +#[bench::credentials_large()] +fn credentials_large_reused(state: (&'static HeaderMap, BasicCredentials)) -> (usize, BasicCredentials) { + let (map, mut output) = state; + let total = decode_rounds(map, &mut output); + (expect_usize(total, 2999 * CREDENTIAL_ROUNDS), output) +} + +#[metabench::benchmark( + CREDENTIALS_LARGE_FRESH, + "credentials", + "credentials_large_fresh", + gungraun_setup = large_map, +)] +#[bench::credentials_large()] +fn credentials_large_fresh(map: &'static HeaderMap) -> usize { + expect_usize(decode_rounds_fresh(map), 2999 * CREDENTIAL_ROUNDS) +} + +// ── group: scanning ────────────────────────────────────────────────────────── + +/// Historical scalar code retained for one before-and-after comparison. +/// +/// Ordinary scalar comparisons call the production scalar backend directly. +/// `is_token68_legacy` is different in kind: it is the per-byte `matches!` +/// chain the `Authorization` parser used *before* `http_headers_simd::is_token68` existed. +/// It is kept so the report can price what that change bought, and it is not +/// what any production path runs today. +mod reference { + /// The per-byte `matches!` chain the `Authorization` parser ran before + /// `http_headers_simd::is_token68` existed. Retained only as a before-and-after + /// reference; no production path executes this today. + pub(super) fn is_token68_legacy(bytes: &[u8]) -> bool { + if bytes.is_empty() { + return false; + } + let data_end = bytes.iter().position(|byte| *byte == b'=').unwrap_or(bytes.len()); + data_end > 0 + && bytes[..data_end].iter().all(|byte| { + matches!( + byte, + b'A'..=b'Z' + | b'a'..=b'z' + | b'0'..=b'9' + | b'-' + | b'.' + | b'_' + | b'~' + | b'+' + | b'/' + ) + }) + && bytes[data_end..].iter().all(|byte| *byte == b'=') + } +} + +fn token_of(length: usize) -> Vec { + let mut bytes = vec![b'a'; length]; + if let Some(last) = bytes.last_mut() { + *last = b'z'; + } + bytes +} + +fn token_15() -> Vec { + token_of(15) +} + +fn token_16() -> Vec { + token_of(16) +} + +fn token_31() -> Vec { + token_of(31) +} + +fn token_32() -> Vec { + token_of(32) +} + +fn token_33() -> Vec { + token_of(33) +} + +fn token_512() -> Vec { + token_of(512) +} + +fn field_value_bytes() -> Vec { + let mut bytes = Vec::with_capacity(512); + while bytes.len() < 512 { + bytes.extend_from_slice(b"text/html; charset=utf-8, application/json; q=0.9, "); + } + bytes.truncate(512); + bytes +} + +/// 512 bytes holding one delimiter every 32, so a scan finds work everywhere. +fn single_line_list() -> Vec { + let mut bytes = Vec::with_capacity(512); + while bytes.len() < 512 { + bytes.extend_from_slice(b"directiveabcdefghijklmnop=value,"); + } + bytes.truncate(512); + bytes +} + +/// The same 512 bytes split into eight short field lines. +fn multi_line_list() -> Vec> { + let line = single_line_list(); + (0..8).map(|index| line[index * 64..][..64].to_vec()).collect() +} + +fn comma_items_with_irrelevant_bytes() -> Vec { + let mut bytes = Vec::with_capacity(512); + while bytes.len() < 512 { + bytes.extend_from_slice(b" alpha=bravo ; charlie = delta , "); + } + bytes.truncate(510); + bytes.extend_from_slice(b",x"); + bytes +} + +fn scan_all(bytes: &[u8], find: fn(&[u8]) -> Option) -> usize { + let mut position = 0; + let mut found = 0; + while position < bytes.len() { + let Some(offset) = find(&bytes[position..]) else { + break; + }; + position += offset + 1; + found += 1; + } + found +} + +fn token_answer(ok: bool, bytes: &[u8], length: usize) -> usize { + expect_usize(usize::from(ok) + bytes.len(), 1 + length) +} + +#[metabench::benchmark( + TOKEN_15_DISPATCHED, + "scanning", + "token_15_dispatched", + gungraun_setup = token_15, + gungraun_teardown = drop_it, +)] +#[bench::token_15()] +fn token_15_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::is_token(&bytes), &bytes, 15); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_15_REFERENCE, + "scanning", + "token_15_reference", + gungraun_setup = token_15, + gungraun_teardown = drop_it, +)] +#[bench::token_15()] +fn token_15_reference(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::benchmarking::is_token_scalar(&bytes), &bytes, 15); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_16_DISPATCHED, + "scanning", + "token_16_dispatched", + gungraun_setup = token_16, + gungraun_teardown = drop_it, +)] +#[bench::token_16()] +fn token_16_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::is_token(&bytes), &bytes, 16); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_16_REFERENCE, + "scanning", + "token_16_reference", + gungraun_setup = token_16, + gungraun_teardown = drop_it, +)] +#[bench::token_16()] +fn token_16_reference(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::benchmarking::is_token_scalar(&bytes), &bytes, 16); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_31_DISPATCHED, + "scanning", + "token_31_dispatched", + gungraun_setup = token_31, + gungraun_teardown = drop_it, +)] +#[bench::token_31()] +fn token_31_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::is_token(&bytes), &bytes, 31); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_31_REFERENCE, + "scanning", + "token_31_reference", + gungraun_setup = token_31, + gungraun_teardown = drop_it, +)] +#[bench::token_31()] +fn token_31_reference(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::benchmarking::is_token_scalar(&bytes), &bytes, 31); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_32_DISPATCHED, + "scanning", + "token_32_dispatched", + gungraun_setup = token_32, + gungraun_teardown = drop_it, +)] +#[bench::token_32()] +fn token_32_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::is_token(&bytes), &bytes, 32); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_32_REFERENCE, + "scanning", + "token_32_reference", + gungraun_setup = token_32, + gungraun_teardown = drop_it, +)] +#[bench::token_32()] +fn token_32_reference(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::benchmarking::is_token_scalar(&bytes), &bytes, 32); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_33_DISPATCHED, + "scanning", + "token_33_dispatched", + gungraun_setup = token_33, + gungraun_teardown = drop_it, +)] +#[bench::token_33()] +fn token_33_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::is_token(&bytes), &bytes, 33); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_33_REFERENCE, + "scanning", + "token_33_reference", + gungraun_setup = token_33, + gungraun_teardown = drop_it, +)] +#[bench::token_33()] +fn token_33_reference(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::benchmarking::is_token_scalar(&bytes), &bytes, 33); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_512_DISPATCHED, + "scanning", + "token_512_dispatched", + gungraun_setup = token_512, + gungraun_teardown = drop_it, +)] +#[bench::token_512()] +fn token_512_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::is_token(&bytes), &bytes, 512); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN_512_REFERENCE, + "scanning", + "token_512_reference", + gungraun_setup = token_512, + gungraun_teardown = drop_it, +)] +#[bench::token_512()] +fn token_512_reference(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::benchmarking::is_token_scalar(&bytes), &bytes, 512); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN68_150_DISPATCHED, + "scanning", + "token68_150_dispatched", + gungraun_setup = bearer_token_bytes, + gungraun_teardown = drop_it, +)] +#[bench::token68_150()] +fn token68_150_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::is_token68(&bytes), &bytes, BEARER_TOKEN.len()); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN68_150_REFERENCE, + "scanning", + "token68_150_reference", + gungraun_setup = bearer_token_bytes, + gungraun_teardown = drop_it, +)] +#[bench::token68_150()] +fn token68_150_reference(bytes: Vec) -> (usize, Vec) { + let answer = token_answer( + http_headers_simd::benchmarking::is_token68_scalar(&bytes), + &bytes, + BEARER_TOKEN.len(), + ); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN68_512_DISPATCHED, + "scanning", + "token68_512_dispatched", + gungraun_setup = token68_512_bytes, + gungraun_teardown = drop_it, +)] +#[bench::token68_512()] +fn token68_512_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::is_token68(&bytes), &bytes, 512); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN68_512_REFERENCE, + "scanning", + "token68_512_reference", + gungraun_setup = token68_512_bytes, + gungraun_teardown = drop_it, +)] +#[bench::token68_512()] +fn token68_512_reference(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::benchmarking::is_token68_scalar(&bytes), &bytes, 512); + (answer, bytes) +} + +// The same 150 bytes as `token68_150`, through a different byte-class kernel: +// identical input and harness, so the two rows differ only by the kernel. +#[metabench::benchmark( + FIELD_VALUE_150_DISPATCHED, + "scanning", + "field_value_150_dispatched", + gungraun_setup = bearer_token_bytes, + gungraun_teardown = drop_it, +)] +#[bench::field_value_150()] +fn field_value_150_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = expect_usize(usize::from(http_headers_simd::is_field_value(&bytes)), 1); + (answer, bytes) +} + +#[metabench::benchmark( + FIELD_VALUE_150_REFERENCE, + "scanning", + "field_value_150_reference", + gungraun_setup = bearer_token_bytes, + gungraun_teardown = drop_it, +)] +#[bench::field_value_150()] +fn field_value_150_reference(bytes: Vec) -> (usize, Vec) { + let answer = expect_usize(usize::from(http_headers_simd::benchmarking::is_field_value_scalar(&bytes)), 1); + (answer, bytes) +} + +#[metabench::benchmark( + FIELD_VALUE_512_DISPATCHED, + "scanning", + "field_value_512_dispatched", + gungraun_setup = field_value_bytes, + gungraun_teardown = drop_it, +)] +#[bench::field_value_512()] +fn field_value_512_dispatched(bytes: Vec) -> (usize, Vec) { + let answer = expect_usize(usize::from(http_headers_simd::is_field_value(&bytes)), 1); + (answer, bytes) +} + +#[metabench::benchmark( + FIELD_VALUE_512_REFERENCE, + "scanning", + "field_value_512_reference", + gungraun_setup = field_value_bytes, + gungraun_teardown = drop_it, +)] +#[bench::field_value_512()] +fn field_value_512_reference(bytes: Vec) -> (usize, Vec) { + let answer = expect_usize(usize::from(http_headers_simd::benchmarking::is_field_value_scalar(&bytes)), 1); + (answer, bytes) +} + +#[metabench::benchmark( + DELIMITERS_SINGLE_LINE_DISPATCHED, + "scanning", + "delimiters_single_line_dispatched", + gungraun_setup = single_line_list, + gungraun_teardown = drop_it, +)] +#[bench::delimiters_single_line()] +fn delimiters_single_line_dispatched(bytes: Vec) -> (usize, Vec) { + let found = expect_usize(scan_all(&bytes, http_headers_simd::find_interesting), 16); + (found, bytes) +} + +#[metabench::benchmark( + DELIMITERS_SINGLE_LINE_REFERENCE, + "scanning", + "delimiters_single_line_reference", + gungraun_setup = single_line_list, + gungraun_teardown = drop_it, +)] +#[bench::delimiters_single_line()] +fn delimiters_single_line_reference(bytes: Vec) -> (usize, Vec) { + let found = expect_usize(scan_all(&bytes, http_headers_simd::benchmarking::find_interesting_scalar), 16); + (found, bytes) +} + +#[metabench::benchmark( + DELIMITERS_MULTI_LINE_DISPATCHED, + "scanning", + "delimiters_multi_line_dispatched", + gungraun_setup = multi_line_list, + gungraun_teardown = drop_it, +)] +#[bench::delimiters_multi_line()] +fn delimiters_multi_line_dispatched(lines: Vec>) -> (usize, Vec>) { + let mut found = 0; + for line in &lines { + found += scan_all(line, http_headers_simd::find_interesting); + } + (expect_usize(found, 16), lines) +} + +#[metabench::benchmark( + DELIMITERS_MULTI_LINE_REFERENCE, + "scanning", + "delimiters_multi_line_reference", + gungraun_setup = multi_line_list, + gungraun_teardown = drop_it, +)] +#[bench::delimiters_multi_line()] +fn delimiters_multi_line_reference(lines: Vec>) -> (usize, Vec>) { + let mut found = 0; + for line in &lines { + found += scan_all(line, http_headers_simd::benchmarking::find_interesting_scalar); + } + (expect_usize(found, 16), lines) +} + +#[metabench::benchmark( + COMMA_ITEMS_IRRELEVANT_BYTES, + "scanning", + "comma_items_irrelevant_bytes", + gungraun_setup = comma_items_with_irrelevant_bytes, + gungraun_teardown = drop_it, +)] +#[bench::comma_items_irrelevant_bytes()] +fn comma_items_irrelevant_bytes(bytes: Vec) -> (usize, Vec) { + let lines = FieldLines::single(&http_headers::FieldName::Vary, &bytes); + let total = lines.comma_items().map(|item| item.expect("valid item").len()).sum::(); + (black_box(total), bytes) +} + +// ── group: list_scanning ───────────────────────────────────────────────────── +// +// Cases straddling the two list-scan thresholds that `scanning` above does +// not cover: `LIST_SIMD_THRESHOLD` (`WIDTH`, 16 bytes — also +// `cors::LIST_SCAN_MIN_LEN`) and `LIST_SHORT_LIMIT` (32 bytes). + +/// A comma-separated list of single-byte tokens totalling exactly `length` +/// bytes, e.g. `a,a,a` — long enough to exercise `scan_token_list` at a +/// chosen length without ever producing a trailing empty member. +fn token_list_of(length: usize) -> Vec { + let mut bytes = Vec::with_capacity(length); + while bytes.len() < length { + if !bytes.is_empty() { + bytes.push(b','); + } + bytes.push(b'a'); + } + bytes.truncate(length); + if bytes.last() == Some(&b',') + && let Some(last) = bytes.last_mut() + { + *last = b'a'; + } + bytes +} + +fn list_scan_15() -> Vec { + token_list_of(15) +} + +fn list_scan_16() -> Vec { + token_list_of(16) +} + +fn list_scan_17() -> Vec { + token_list_of(17) +} + +fn list_scan_31() -> Vec { + token_list_of(31) +} + +fn list_scan_32() -> Vec { + token_list_of(32) +} + +fn list_scan_33() -> Vec { + token_list_of(33) +} + +fn list_scan_answer(scan: http_headers_simd::TokenListScan, length: usize) -> usize { + let ok = matches!(scan, http_headers_simd::TokenListScan::Members); + expect_usize(usize::from(ok) + length, 1 + length) +} + +macro_rules! list_scan_case { + ( + $dispatched:ident, + $dispatched_name:literal, + $reference:ident, + $reference_name:literal, + $bench:ident, + $setup:ident + ) => { + paste::paste! { + #[metabench::benchmark( + [<$dispatched:upper>], + "list_scanning", + $dispatched_name, + gungraun_setup = $setup, + gungraun_teardown = drop_it, + )] + #[bench::$bench()] + fn $dispatched(bytes: Vec) -> (usize, Vec) { + let scan = http_headers_simd::scan_token_list( + &bytes, + http_headers_simd::EmptyMembers::Skip, + ); + let answer = list_scan_answer(scan, bytes.len()); + (answer, bytes) + } + + #[metabench::benchmark( + [<$reference:upper>], + "list_scanning", + $reference_name, + gungraun_setup = $setup, + gungraun_teardown = drop_it, + )] + #[bench::$bench()] + fn $reference(bytes: Vec) -> (usize, Vec) { + let scan = http_headers_simd::benchmarking::scan_token_list_scalar( + &bytes, + http_headers_simd::EmptyMembers::Skip, + ); + let answer = list_scan_answer(scan, bytes.len()); + (answer, bytes) + } + } + }; +} + +list_scan_case!( + list_scan_15_dispatched, + "list_scan_15_dispatched", + list_scan_15_reference, + "list_scan_15_reference", + list_scan_15, + list_scan_15 +); +list_scan_case!( + list_scan_16_dispatched, + "list_scan_16_dispatched", + list_scan_16_reference, + "list_scan_16_reference", + list_scan_16, + list_scan_16 +); +list_scan_case!( + list_scan_17_dispatched, + "list_scan_17_dispatched", + list_scan_17_reference, + "list_scan_17_reference", + list_scan_17, + list_scan_17 +); +list_scan_case!( + list_scan_31_dispatched, + "list_scan_31_dispatched", + list_scan_31_reference, + "list_scan_31_reference", + list_scan_31, + list_scan_31 +); +list_scan_case!( + list_scan_32_dispatched, + "list_scan_32_dispatched", + list_scan_32_reference, + "list_scan_32_reference", + list_scan_32, + list_scan_32 +); +list_scan_case!( + list_scan_33_dispatched, + "list_scan_33_dispatched", + list_scan_33_reference, + "list_scan_33_reference", + list_scan_33, + list_scan_33 +); + +// ── group: uri_scanning ────────────────────────────────────────────────────── +// +// Cases straddling `URI_SIMD_THRESHOLD` (16 bytes). + +/// A `path-absolute` URI reference of exactly `length` bytes. +fn uri_path_of(length: usize) -> Vec { + let mut bytes = vec![b'a'; length.max(1)]; + bytes[0] = b'/'; + bytes +} + +fn uri_path_15() -> Vec { + uri_path_of(15) +} + +fn uri_path_16() -> Vec { + uri_path_of(16) +} + +fn uri_path_17() -> Vec { + uri_path_of(17) +} + +fn uri_path_answer(ok: bool, length: usize) -> usize { + expect_usize(usize::from(ok) + length, 1 + length) +} + +macro_rules! uri_path_case { + ( + $dispatched:ident, + $dispatched_name:literal, + $reference:ident, + $reference_name:literal, + $bench:ident, + $setup:ident + ) => { + paste::paste! { + #[metabench::benchmark( + [<$dispatched:upper>], + "uri_scanning", + $dispatched_name, + gungraun_setup = $setup, + gungraun_teardown = drop_it, + )] + #[bench::$bench()] + fn $dispatched(bytes: Vec) -> (usize, Vec) { + let answer = + uri_path_answer(http_headers_simd::is_simple_uri_path(&bytes), bytes.len()); + (answer, bytes) + } + + #[metabench::benchmark( + [<$reference:upper>], + "uri_scanning", + $reference_name, + gungraun_setup = $setup, + gungraun_teardown = drop_it, + )] + #[bench::$bench()] + fn $reference(bytes: Vec) -> (usize, Vec) { + let answer = uri_path_answer( + http_headers_simd::benchmarking::is_simple_uri_path_scalar(&bytes), + bytes.len(), + ); + (answer, bytes) + } + } + }; +} + +uri_path_case!( + uri_path_15_dispatched, + "uri_path_15_dispatched", + uri_path_15_reference, + "uri_path_15_reference", + uri_path_15, + uri_path_15 +); +uri_path_case!( + uri_path_16_dispatched, + "uri_path_16_dispatched", + uri_path_16_reference, + "uri_path_16_reference", + uri_path_16, + uri_path_16 +); +uri_path_case!( + uri_path_17_dispatched, + "uri_path_17_dispatched", + uri_path_17_reference, + "uri_path_17_reference", + uri_path_17, + uri_path_17 +); + +// ── group: forced_uri_backends ────────────────────────────────────────────── +// +// The same input is sent directly to every compiled URI backend. Unsupported +// backends return `None`, so the x86 and AArch64 scheduled runners together +// produce a real count for SSE2, SSSE3, SSE4.2, and NEON without relying on +// the dispatcher's choice for the host. + +fn uri_path_64() -> Vec { + uri_path_of(64) +} + +fn forced_uri_answer(result: Option, length: usize) -> usize { + black_box(usize::from(result.unwrap_or(false)) + length) +} + +macro_rules! forced_uri_backend_case { + ($name:ident, $benchmark_name:literal, $scanner:path) => { + paste::paste! { + #[metabench::benchmark( + [<$name:upper>], + "forced_uri_backends", + $benchmark_name, + gungraun_setup = uri_path_64, + gungraun_teardown = drop_it, + )] + #[bench::uri_path_64()] + fn $name(bytes: Vec) -> (usize, Vec) { + let answer = forced_uri_answer($scanner(&bytes), bytes.len()); + (answer, bytes) + } + } + }; +} + +forced_uri_backend_case!( + uri_path_64_sse2, + "uri_path_64_sse2", + http_headers_simd::benchmarking::is_simple_uri_path_sse2 +); +forced_uri_backend_case!( + uri_path_64_ssse3, + "uri_path_64_ssse3", + http_headers_simd::benchmarking::is_simple_uri_path_ssse3 +); +forced_uri_backend_case!( + uri_path_64_sse42, + "uri_path_64_sse42", + http_headers_simd::benchmarking::is_simple_uri_path_sse42 +); +forced_uri_backend_case!( + uri_path_64_neon, + "uri_path_64_neon", + http_headers_simd::benchmarking::is_simple_uri_path_neon +); + +// ── group: credential_scan ─────────────────────────────────────────────────── +// +// Measures the shared `token68` scanner on the credential in the corpus. + +fn bearer_token_bytes() -> Vec { + let bytes = BEARER_TOKEN.to_vec(); + black_box(http_headers_simd::is_token68(&bytes)); + bytes +} + +fn token68_512_bytes() -> Vec { + let bytes = token_512(); + black_box(http_headers_simd::is_token68(&bytes)); + bytes +} + +#[metabench::benchmark( + TOKEN68_CREDENTIAL_SHARED_SCANNER, + "credential_scan", + "token68_credential_shared_scanner", + gungraun_setup = bearer_token_bytes, + gungraun_teardown = drop_it, +)] +#[bench::token68_credential()] +fn token68_credential_shared_scanner(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(http_headers_simd::is_token68(&bytes), &bytes, BEARER_TOKEN.len()); + (answer, bytes) +} + +#[metabench::benchmark( + TOKEN68_CREDENTIAL_MATCH_CHAIN, + "credential_scan", + "token68_credential_match_chain", + gungraun_setup = bearer_token_bytes, + gungraun_teardown = drop_it, +)] +#[bench::token68_credential()] +fn token68_credential_match_chain(bytes: Vec) -> (usize, Vec) { + let answer = token_answer(reference::is_token68_legacy(&bytes), &bytes, BEARER_TOKEN.len()); + (answer, bytes) +} + +// ── group: custom ──────────────────────────────────────────────────────────── + +/// A downstream header written with nothing but the public extension API. +/// +/// It names `User-Agent` and validates exactly what the built-in `UserAgent` +/// validates, so the pair prices framework overhead rather than two different +/// grammars over two different values. +struct DownstreamAgent(FieldValue); + +/// The borrowed view of [`DownstreamAgent`]. +#[derive(Clone, Copy)] +struct DownstreamAgentView<'a>(FieldValueRef<'a>); + +impl<'a> DownstreamAgentView<'a> { + fn as_bytes(self) -> &'a [u8] { + self.0.as_bytes() + } +} + +impl Field for DownstreamAgent { + type View<'a> = DownstreamAgentView<'a>; + type Owned = Self; + + fn name() -> &'static http_headers::FieldName { + &http_headers::FieldName::UserAgent + } + + fn view_with(source: &S, _mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(values) = source.lines(Self::name()) else { + return Ok(None); + }; + let value = values.exactly_one()?; + if value.as_bytes().is_empty() || !http_headers_simd::is_field_value(value.as_bytes()) { + return Err(DecodeError::new( + &http_headers::FieldName::UserAgent, + DecodeErrorKind::InvalidSyntax, + )); + } + Ok(Some(DownstreamAgentView(value))) + } + + fn owned_with(source: &S, mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + Self::view_with(source, mode).map(|view| { + view.map(|view| { + Self( + view.0 + .try_to_field_value() + .expect("view_with validated the HTTP field-value grammar"), + ) + }) + }) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::single(value.0)) + } +} + +fn downstream_agent() -> (HeaderMap, DownstreamAgent) { + (empty_map(), DownstreamAgent(field_value(USER_AGENT_VALUE))) +} + +struct RawUserAgentSource; + +impl FieldSource for RawUserAgentSource { + fn lines(&self, name: &'static http_headers::FieldName) -> Option> { + (name == &http_headers::FieldName::UserAgent).then(|| FieldLines::single(name, USER_AGENT_VALUE)) + } +} + +fn raw_user_agent_source() -> RawUserAgentSource { + RawUserAgentSource +} + +#[metabench::benchmark( + CUSTOM_BORROWED_BUILTIN, + "custom", + "custom_borrowed_builtin", + gungraun_setup = json_map, +)] +#[bench::custom_borrowed()] +fn custom_borrowed_builtin(map: &'static HeaderMap) -> usize { + let view = UserAgent::view(map).expect("valid user agent").expect("present user agent"); + expect_bytes(view.as_bytes(), USER_AGENT_VALUE) +} + +#[metabench::benchmark( + CUSTOM_BORROWED_CUSTOM, + "custom", + "custom_borrowed_custom", + gungraun_setup = json_map, +)] +#[bench::custom_borrowed()] +fn custom_borrowed_custom(map: &'static HeaderMap) -> usize { + let view = DownstreamAgent::view(map).expect("valid user agent").expect("present user agent"); + expect_bytes(view.as_bytes(), USER_AGENT_VALUE) +} + +#[metabench::benchmark( + CUSTOM_OWNED_BUILTIN, + "custom", + "custom_owned_builtin", + gungraun_setup = json_map, +)] +#[bench::custom_owned()] +fn custom_owned_builtin(map: &'static HeaderMap) -> usize { + let agent = UserAgent::owned(map).expect("valid user agent").expect("present user agent"); + let length = expect_bytes(agent.as_bytes(), USER_AGENT_VALUE); + consume(agent); + length +} + +#[metabench::benchmark( + CUSTOM_OWNED_CUSTOM, + "custom", + "custom_owned_custom", + gungraun_setup = json_map, +)] +#[bench::custom_owned()] +fn custom_owned_custom(map: &'static HeaderMap) -> usize { + let agent = DownstreamAgent::owned(map).expect("valid user agent").expect("present user agent"); + let length = expect_bytes(agent.0.as_bytes(), USER_AGENT_VALUE); + consume(agent); + length +} + +#[metabench::benchmark( + CUSTOM_OWNED_RAW_SOURCE, + "custom", + "custom_owned_raw_source", + gungraun_setup = raw_user_agent_source, +)] +#[bench::custom_owned_raw_source()] +fn custom_owned_raw_source(source: RawUserAgentSource) -> usize { + let agent = UserAgent::owned(&source).expect("valid user agent").expect("present user agent"); + let length = expect_bytes(agent.as_bytes(), USER_AGENT_VALUE); + consume(agent); + length +} + +#[metabench::benchmark( + CUSTOM_INSERT_BUILTIN, + "custom", + "custom_insert_builtin", + gungraun_setup = our_user_agent, + gungraun_teardown = drop_it, +)] +#[bench::custom_insert()] +fn custom_insert_builtin(state: (HeaderMap, UserAgentOwned)) -> (usize, HeaderMap) { + let (mut map, agent) = state; + expect_inserted(UserAgent::insert(&mut map, agent)); + let stored = map.get(&USER_AGENT).expect("inserted user agent"); + let length = expect_bytes(stored.as_bytes(), USER_AGENT_VALUE); + (length, map) +} + +#[metabench::benchmark( + CUSTOM_INSERT_CUSTOM, + "custom", + "custom_insert_custom", + gungraun_setup = downstream_agent, + gungraun_teardown = drop_it, +)] +#[bench::custom_insert()] +fn custom_insert_custom(state: (HeaderMap, DownstreamAgent)) -> (usize, HeaderMap) { + let (mut map, agent) = state; + expect_inserted(DownstreamAgent::insert(&mut map, agent)); + let stored = map.get(&USER_AGENT).expect("inserted user agent"); + let length = expect_bytes(stored.as_bytes(), USER_AGENT_VALUE); + (length, map) +} + +const WARM_UP: Duration = Duration::from_secs(1); +const MEASUREMENT: Duration = Duration::from_secs(3); +const SAMPLES: usize = 60; + +fn tuned_criterion() -> Criterion { + Criterion::default() + .warm_up_time(WARM_UP) + .measurement_time(MEASUREMENT) + .sample_size(SAMPLES) +} + +macro_rules! register_case { + ($group:ident, $identity:ident, $case:literal, $setup:path, $benchmark:path) => { + $group.bench_with_input(BenchmarkId::new($identity.benchmark_name(), $case), &(), |bencher, &()| { + bencher.iter_batched($setup, $benchmark, BatchSize::SmallInput); + }); + }; +} + +macro_rules! register_group { + ( + $criterion:ident, + $name:literal, + [$(($identity:ident, $case:literal, $setup:path, $benchmark:path)),+ $(,)?] + ) => { + let mut group = $criterion.benchmark_group(concat!("http_headers_micro/", $name)); + register_cases!( + group, + [$(($identity, $case, $setup, $benchmark)),+] + ); + group.finish(); + }; +} + +macro_rules! register_cases { + ( + $group:ident, + [$(($identity:ident, $case:literal, $setup:path, $benchmark:path)),+ $(,)?] + ) => { + $(register_case!($group, $identity, $case, $setup, $benchmark);)+ + }; +} + +fn criterion_lookup(criterion: &mut Criterion) { + register_group!( + criterion, + "lookup", + [ + ( + USER_AGENT_BORROWED_HTTP_HEADERS, + "user_agent_borrowed", + json_map, + user_agent_borrowed_http_headers + ), + ( + USER_AGENT_BORROWED_HEADERS, + "user_agent_borrowed", + json_map, + user_agent_borrowed_headers + ), + ( + USER_AGENT_OWNED_HTTP_HEADERS, + "user_agent_owned", + json_map, + user_agent_owned_http_headers + ), + (USER_AGENT_OWNED_HEADERS, "user_agent_owned", json_map, user_agent_owned_headers), + ( + USER_AGENT_ABSENT_HTTP_HEADERS, + "user_agent_absent", + cookies_4, + user_agent_absent_http_headers + ), + (USER_AGENT_ABSENT_HEADERS, "user_agent_absent", cookies_4, user_agent_absent_headers), + ] + ); +} + +fn criterion_repeated(criterion: &mut Criterion) { + register_group!( + criterion, + "repeated", + [ + ( + SET_COOKIE_BORROWED_1_HTTP_HEADERS, + "set_cookie_borrowed_1", + cookies_1, + set_cookie_borrowed_1_http_headers + ), + ( + SET_COOKIE_BORROWED_1_HEADERS, + "set_cookie_borrowed_1", + cookies_1, + set_cookie_borrowed_1_headers + ), + ( + SET_COOKIE_BORROWED_4_HTTP_HEADERS, + "set_cookie_borrowed_4", + cookies_4, + set_cookie_borrowed_4_http_headers + ), + ( + SET_COOKIE_BORROWED_4_HEADERS, + "set_cookie_borrowed_4", + cookies_4, + set_cookie_borrowed_4_headers + ), + ( + SET_COOKIE_BORROWED_12_HTTP_HEADERS, + "set_cookie_borrowed_12", + cookies_12, + set_cookie_borrowed_12_http_headers + ), + ( + SET_COOKIE_BORROWED_12_HEADERS, + "set_cookie_borrowed_12", + cookies_12, + set_cookie_borrowed_12_headers + ), + ( + SET_COOKIE_OWNED_4_HTTP_HEADERS, + "set_cookie_owned_4", + cookies_4, + set_cookie_owned_4_http_headers + ), + ( + SET_COOKIE_OWNED_4_HEADERS, + "set_cookie_owned_4", + cookies_4, + set_cookie_owned_4_headers + ), + ] + ); +} + +fn criterion_list(criterion: &mut Criterion) { + register_group!( + criterion, + "list", + [ + ( + CACHE_CONTROL_SINGLE_LINE_HTTP_HEADERS, + "cache_control_single_line", + json_map, + cache_control_single_line_http_headers + ), + ( + CACHE_CONTROL_SINGLE_LINE_HEADERS, + "cache_control_single_line", + json_map, + cache_control_single_line_headers + ), + ( + CACHE_CONTROL_MULTI_LINE_HTTP_HEADERS, + "cache_control_multi_line", + cache_multi_map, + cache_control_multi_line_http_headers + ), + ( + CACHE_CONTROL_MULTI_LINE_HEADERS, + "cache_control_multi_line", + cache_multi_map, + cache_control_multi_line_headers + ), + ( + CACHE_CONTROL_OWNED_HTTP_HEADERS, + "cache_control_owned", + json_map, + cache_control_owned_http_headers + ), + ( + CACHE_CONTROL_OWNED_HEADERS, + "cache_control_owned", + json_map, + cache_control_owned_headers + ), + ( + CACHE_CONTROL_ADVERSARIAL_HTTP_HEADERS, + "cache_control_adversarial", + adversarial_map, + cache_control_adversarial_http_headers + ), + ( + CACHE_CONTROL_ADVERSARIAL_HEADERS, + "cache_control_adversarial", + adversarial_map, + cache_control_adversarial_headers + ), + ( + CACHE_CONTROL_MAX_AGE_ONLY_HTTP_HEADERS, + "cache_control_max_age_only", + json_map, + cache_control_max_age_only_http_headers + ), + ( + CACHE_CONTROL_MAX_AGE_ONLY_HEADERS, + "cache_control_max_age_only", + json_map, + cache_control_max_age_only_headers + ), + ( + CACHE_CONTROL_ONE_DIRECTIVE_HTTP_HEADERS, + "cache_control_one_directive", + cache_one_map, + cache_control_one_directive_http_headers + ), + ( + CACHE_CONTROL_ONE_DIRECTIVE_HEADERS, + "cache_control_one_directive", + cache_one_map, + cache_control_one_directive_headers + ), + ] + ); +} + +fn criterion_directives(criterion: &mut Criterion) { + register_group!( + criterion, + "directives", + [ + (DIRECTIVE_VALUE_BYTES, "directive_value", json_map, directive_value_bytes), + (DIRECTIVE_VALUE_STR, "directive_value", json_map, directive_value_str), + ( + DIRECTIVE_VALUE_MANY_BYTES, + "directive_value_many", + adversarial_map, + directive_value_many_bytes + ), + ( + DIRECTIVE_VALUE_MANY_STR, + "directive_value_many", + adversarial_map, + directive_value_many_str + ), + ] + ); +} + +fn criterion_structured(criterion: &mut Criterion) { + register_group!( + criterion, + "structured", + [ + ( + CONTENT_TYPE_INSPECT_HTTP_HEADERS, + "content_type_inspect", + json_map, + content_type_inspect_http_headers + ), + ( + CONTENT_TYPE_INSPECT_HEADERS, + "content_type_inspect", + json_map, + content_type_inspect_headers + ), + ( + CONTENT_TYPE_PARAMETERS_HTTP_HEADERS, + "content_type_parameters", + content_type_many_map, + content_type_parameters_http_headers + ), + ( + CONTENT_TYPE_PARAMETERS_HEADERS, + "content_type_parameters", + content_type_many_map, + content_type_parameters_headers + ), + ( + CONTENT_TYPE_OWNED_HTTP_HEADERS, + "content_type_owned", + json_map, + content_type_owned_http_headers + ), + ( + CONTENT_TYPE_OWNED_HEADERS, + "content_type_owned", + json_map, + content_type_owned_headers + ), + ( + CONTENT_TYPE_MALFORMED_HTTP_HEADERS, + "content_type_malformed", + degraded_map, + content_type_malformed_http_headers + ), + ( + CONTENT_TYPE_MALFORMED_HEADERS, + "content_type_malformed", + degraded_map, + content_type_malformed_headers + ), + ] + ); +} + +fn criterion_authorization(criterion: &mut Criterion) { + register_group!( + criterion, + "authorization", + [ + ( + BASIC_DECODE_HTTP_HEADERS, + "basic_decode", + basic_credentials, + basic_decode_http_headers + ), + (BASIC_DECODE_HEADERS, "basic_decode", basic_credentials, basic_decode_headers), + ( + BASIC_VALIDATE_HTTP_HEADERS, + "basic_validate", + basic_map, + basic_validate_http_headers + ), + (BASIC_VALIDATE_HEADERS, "basic_validate", basic_map, basic_validate_headers), + ( + BEARER_BORROWED_HTTP_HEADERS, + "bearer_borrowed", + bearer_map, + bearer_borrowed_http_headers + ), + (BEARER_BORROWED_HEADERS, "bearer_borrowed", bearer_map, bearer_borrowed_headers), + (BASIC_ENCODE_HTTP_HEADERS, "basic_encode", empty_map, basic_encode_http_headers), + (BASIC_ENCODE_HEADERS, "basic_encode", empty_map, basic_encode_headers), + ] + ); +} + +fn criterion_insertion(criterion: &mut Criterion) { + register_group!( + criterion, + "insertion", + [ + ( + INSERT_USER_AGENT_HTTP_HEADERS, + "insert_user_agent", + our_user_agent, + insert_user_agent_http_headers + ), + ( + INSERT_USER_AGENT_HEADERS, + "insert_user_agent", + their_user_agent, + insert_user_agent_headers + ), + ( + INSERT_CONTENT_TYPE_HTTP_HEADERS, + "insert_content_type", + our_content_type, + insert_content_type_http_headers + ), + ( + INSERT_CONTENT_TYPE_HEADERS, + "insert_content_type", + their_content_type, + insert_content_type_headers + ), + ( + INSERT_CACHE_CONTROL_HTTP_HEADERS, + "insert_cache_control", + our_cache_control, + insert_cache_control_http_headers + ), + ( + INSERT_CACHE_CONTROL_HEADERS, + "insert_cache_control", + their_cache_control, + insert_cache_control_headers + ), + ] + ); +} + +fn criterion_ownership(criterion: &mut Criterion) { + register_group!( + criterion, + "ownership", + [ + ( + OWNERSHIP_USER_AGENT_MOVED, + "ownership_user_agent", + our_user_agent, + ownership_user_agent_moved + ), + ( + OWNERSHIP_USER_AGENT_CLONED, + "ownership_user_agent", + our_user_agent, + ownership_user_agent_cloned + ), + ( + OWNERSHIP_SET_COOKIE_MOVED, + "ownership_set_cookie", + our_cookies, + ownership_set_cookie_moved + ), + ( + OWNERSHIP_SET_COOKIE_CLONED, + "ownership_set_cookie", + our_cookies, + ownership_set_cookie_cloned + ), + ] + ); +} + +fn criterion_credentials(criterion: &mut Criterion) { + register_group!( + criterion, + "credentials", + [ + ( + CREDENTIALS_BASIC_REUSED, + "credentials_basic", + warm_credentials, + credentials_basic_reused + ), + ( + CREDENTIALS_BASIC_HIGH_WATER_REUSED, + "credentials_basic_high_water", + high_water_warm_credentials, + credentials_basic_high_water_reused + ), + (CREDENTIALS_BASIC_FRESH, "credentials_basic", basic_map, credentials_basic_fresh), + ( + CREDENTIALS_LARGE_REUSED, + "credentials_large", + large_warm_credentials, + credentials_large_reused + ), + (CREDENTIALS_LARGE_FRESH, "credentials_large", large_map, credentials_large_fresh), + ] + ); +} + +fn criterion_scanning_tokens(group: &mut BenchmarkGroup<'_, WallTime>) { + register_cases!( + group, + [ + (TOKEN_15_DISPATCHED, "token_15", token_15, token_15_dispatched), + (TOKEN_15_REFERENCE, "token_15", token_15, token_15_reference), + (TOKEN_16_DISPATCHED, "token_16", token_16, token_16_dispatched), + (TOKEN_16_REFERENCE, "token_16", token_16, token_16_reference), + (TOKEN_31_DISPATCHED, "token_31", token_31, token_31_dispatched), + (TOKEN_31_REFERENCE, "token_31", token_31, token_31_reference), + (TOKEN_32_DISPATCHED, "token_32", token_32, token_32_dispatched), + (TOKEN_32_REFERENCE, "token_32", token_32, token_32_reference), + (TOKEN_33_DISPATCHED, "token_33", token_33, token_33_dispatched), + (TOKEN_33_REFERENCE, "token_33", token_33, token_33_reference), + (TOKEN_512_DISPATCHED, "token_512", token_512, token_512_dispatched), + (TOKEN_512_REFERENCE, "token_512", token_512, token_512_reference), + ] + ); +} + +fn criterion_scanning_values(group: &mut BenchmarkGroup<'_, WallTime>) { + register_cases!( + group, + [ + (TOKEN68_150_DISPATCHED, "token68_150", bearer_token_bytes, token68_150_dispatched), + (TOKEN68_150_REFERENCE, "token68_150", bearer_token_bytes, token68_150_reference), + (TOKEN68_512_DISPATCHED, "token68_512", token68_512_bytes, token68_512_dispatched), + (TOKEN68_512_REFERENCE, "token68_512", token68_512_bytes, token68_512_reference), + ( + FIELD_VALUE_150_DISPATCHED, + "field_value_150", + bearer_token_bytes, + field_value_150_dispatched + ), + ( + FIELD_VALUE_150_REFERENCE, + "field_value_150", + bearer_token_bytes, + field_value_150_reference + ), + ( + FIELD_VALUE_512_DISPATCHED, + "field_value_512", + field_value_bytes, + field_value_512_dispatched + ), + ( + FIELD_VALUE_512_REFERENCE, + "field_value_512", + field_value_bytes, + field_value_512_reference + ), + ( + DELIMITERS_SINGLE_LINE_DISPATCHED, + "delimiters_single_line", + single_line_list, + delimiters_single_line_dispatched + ), + ( + DELIMITERS_SINGLE_LINE_REFERENCE, + "delimiters_single_line", + single_line_list, + delimiters_single_line_reference + ), + ( + DELIMITERS_MULTI_LINE_DISPATCHED, + "delimiters_multi_line", + multi_line_list, + delimiters_multi_line_dispatched + ), + ( + DELIMITERS_MULTI_LINE_REFERENCE, + "delimiters_multi_line", + multi_line_list, + delimiters_multi_line_reference + ), + ] + ); +} + +fn criterion_scanning(criterion: &mut Criterion) { + let mut group = criterion.benchmark_group("http_headers_micro/scanning"); + criterion_scanning_tokens(&mut group); + criterion_scanning_values(&mut group); + group.finish(); +} + +fn criterion_list_scanning(criterion: &mut Criterion) { + register_group!( + criterion, + "list_scanning", + [ + (LIST_SCAN_15_DISPATCHED, "list_scan_15", list_scan_15, list_scan_15_dispatched), + (LIST_SCAN_15_REFERENCE, "list_scan_15", list_scan_15, list_scan_15_reference), + (LIST_SCAN_16_DISPATCHED, "list_scan_16", list_scan_16, list_scan_16_dispatched), + (LIST_SCAN_16_REFERENCE, "list_scan_16", list_scan_16, list_scan_16_reference), + (LIST_SCAN_17_DISPATCHED, "list_scan_17", list_scan_17, list_scan_17_dispatched), + (LIST_SCAN_17_REFERENCE, "list_scan_17", list_scan_17, list_scan_17_reference), + (LIST_SCAN_31_DISPATCHED, "list_scan_31", list_scan_31, list_scan_31_dispatched), + (LIST_SCAN_31_REFERENCE, "list_scan_31", list_scan_31, list_scan_31_reference), + (LIST_SCAN_32_DISPATCHED, "list_scan_32", list_scan_32, list_scan_32_dispatched), + (LIST_SCAN_32_REFERENCE, "list_scan_32", list_scan_32, list_scan_32_reference), + (LIST_SCAN_33_DISPATCHED, "list_scan_33", list_scan_33, list_scan_33_dispatched), + (LIST_SCAN_33_REFERENCE, "list_scan_33", list_scan_33, list_scan_33_reference), + ] + ); +} + +fn criterion_uri_scanning(criterion: &mut Criterion) { + register_group!( + criterion, + "uri_scanning", + [ + (URI_PATH_15_DISPATCHED, "uri_path_15", uri_path_15, uri_path_15_dispatched), + (URI_PATH_15_REFERENCE, "uri_path_15", uri_path_15, uri_path_15_reference), + (URI_PATH_16_DISPATCHED, "uri_path_16", uri_path_16, uri_path_16_dispatched), + (URI_PATH_16_REFERENCE, "uri_path_16", uri_path_16, uri_path_16_reference), + (URI_PATH_17_DISPATCHED, "uri_path_17", uri_path_17, uri_path_17_dispatched), + (URI_PATH_17_REFERENCE, "uri_path_17", uri_path_17, uri_path_17_reference), + ] + ); +} + +fn criterion_forced_uri_backends(criterion: &mut Criterion) { + register_group!( + criterion, + "forced_uri_backends", + [ + (URI_PATH_64_SSE2, "uri_path_64", uri_path_64, uri_path_64_sse2), + (URI_PATH_64_SSSE3, "uri_path_64", uri_path_64, uri_path_64_ssse3), + (URI_PATH_64_SSE42, "uri_path_64", uri_path_64, uri_path_64_sse42), + (URI_PATH_64_NEON, "uri_path_64", uri_path_64, uri_path_64_neon), + ] + ); +} + +fn criterion_credential_scan(criterion: &mut Criterion) { + register_group!( + criterion, + "credential_scan", + [ + ( + TOKEN68_CREDENTIAL_SHARED_SCANNER, + "token68_credential", + bearer_token_bytes, + token68_credential_shared_scanner + ), + ( + TOKEN68_CREDENTIAL_MATCH_CHAIN, + "token68_credential", + bearer_token_bytes, + token68_credential_match_chain + ), + ] + ); +} + +fn criterion_custom(criterion: &mut Criterion) { + register_group!( + criterion, + "custom", + [ + (CUSTOM_BORROWED_BUILTIN, "custom_borrowed", json_map, custom_borrowed_builtin), + (CUSTOM_BORROWED_CUSTOM, "custom_borrowed", json_map, custom_borrowed_custom), + (CUSTOM_OWNED_BUILTIN, "custom_owned", json_map, custom_owned_builtin), + (CUSTOM_OWNED_CUSTOM, "custom_owned", json_map, custom_owned_custom), + (CUSTOM_INSERT_BUILTIN, "custom_insert", our_user_agent, custom_insert_builtin), + (CUSTOM_INSERT_CUSTOM, "custom_insert", downstream_agent, custom_insert_custom), + ] + ); +} + +fn criterion_benchmarks(criterion: &mut Criterion) { + criterion_lookup(criterion); + criterion_repeated(criterion); + criterion_list(criterion); + criterion_directives(criterion); + criterion_structured(criterion); + criterion_authorization(criterion); + criterion_insertion(criterion); + criterion_ownership(criterion); + criterion_credentials(criterion); + criterion_scanning(criterion); + criterion_list_scanning(criterion); + criterion_uri_scanning(criterion); + criterion_forced_uri_backends(criterion); + criterion_credential_scan(criterion); + criterion_custom(criterion); +} + +metabench::main!( + criterion = { + factory = tuned_criterion, + benchmarks = criterion_benchmarks, + unit = "ns", + }, + groups = { + LOOKUP { + benchmarks = [ + USER_AGENT_BORROWED_HTTP_HEADERS, + USER_AGENT_BORROWED_HEADERS, + USER_AGENT_OWNED_HTTP_HEADERS, + USER_AGENT_OWNED_HEADERS, + USER_AGENT_ABSENT_HTTP_HEADERS, + USER_AGENT_ABSENT_HEADERS, + ], + gungraun_compare_by_id = true, + }, + REPEATED { + benchmarks = [ + SET_COOKIE_BORROWED_1_HTTP_HEADERS, + SET_COOKIE_BORROWED_1_HEADERS, + SET_COOKIE_BORROWED_4_HTTP_HEADERS, + SET_COOKIE_BORROWED_4_HEADERS, + SET_COOKIE_BORROWED_12_HTTP_HEADERS, + SET_COOKIE_BORROWED_12_HEADERS, + SET_COOKIE_OWNED_4_HTTP_HEADERS, + SET_COOKIE_OWNED_4_HEADERS, + ], + gungraun_compare_by_id = true, + }, + LIST { + benchmarks = [ + CACHE_CONTROL_SINGLE_LINE_HTTP_HEADERS, + CACHE_CONTROL_SINGLE_LINE_HEADERS, + CACHE_CONTROL_MULTI_LINE_HTTP_HEADERS, + CACHE_CONTROL_MULTI_LINE_HEADERS, + CACHE_CONTROL_OWNED_HTTP_HEADERS, + CACHE_CONTROL_OWNED_HEADERS, + CACHE_CONTROL_ADVERSARIAL_HTTP_HEADERS, + CACHE_CONTROL_ADVERSARIAL_HEADERS, + CACHE_CONTROL_MAX_AGE_ONLY_HTTP_HEADERS, + CACHE_CONTROL_MAX_AGE_ONLY_HEADERS, + CACHE_CONTROL_ONE_DIRECTIVE_HTTP_HEADERS, + CACHE_CONTROL_ONE_DIRECTIVE_HEADERS, + ], + gungraun_compare_by_id = true, + }, + DIRECTIVES { + benchmarks = [ + DIRECTIVE_VALUE_BYTES, + DIRECTIVE_VALUE_STR, + DIRECTIVE_VALUE_MANY_BYTES, + DIRECTIVE_VALUE_MANY_STR, + ], + gungraun_compare_by_id = true, + }, + STRUCTURED { + benchmarks = [ + CONTENT_TYPE_INSPECT_HTTP_HEADERS, + CONTENT_TYPE_INSPECT_HEADERS, + CONTENT_TYPE_PARAMETERS_HTTP_HEADERS, + CONTENT_TYPE_PARAMETERS_HEADERS, + CONTENT_TYPE_OWNED_HTTP_HEADERS, + CONTENT_TYPE_OWNED_HEADERS, + CONTENT_TYPE_MALFORMED_HTTP_HEADERS, + CONTENT_TYPE_MALFORMED_HEADERS, + ], + gungraun_compare_by_id = true, + }, + AUTHORIZATION { + benchmarks = [ + BASIC_DECODE_HTTP_HEADERS, + BASIC_DECODE_HEADERS, + BASIC_VALIDATE_HTTP_HEADERS, + BASIC_VALIDATE_HEADERS, + BEARER_BORROWED_HTTP_HEADERS, + BEARER_BORROWED_HEADERS, + BASIC_ENCODE_HTTP_HEADERS, + BASIC_ENCODE_HEADERS, + ], + gungraun_compare_by_id = true, + }, + INSERTION { + benchmarks = [ + INSERT_USER_AGENT_HTTP_HEADERS, + INSERT_USER_AGENT_HEADERS, + INSERT_CONTENT_TYPE_HTTP_HEADERS, + INSERT_CONTENT_TYPE_HEADERS, + INSERT_CACHE_CONTROL_HTTP_HEADERS, + INSERT_CACHE_CONTROL_HEADERS, + ], + gungraun_compare_by_id = true, + }, + OWNERSHIP { + benchmarks = [ + OWNERSHIP_USER_AGENT_MOVED, + OWNERSHIP_USER_AGENT_CLONED, + OWNERSHIP_SET_COOKIE_MOVED, + OWNERSHIP_SET_COOKIE_CLONED, + ], + gungraun_compare_by_id = true, + }, + CREDENTIALS { + benchmarks = [ + CREDENTIALS_BASIC_REUSED, + CREDENTIALS_BASIC_HIGH_WATER_REUSED, + CREDENTIALS_BASIC_FRESH, + CREDENTIALS_LARGE_REUSED, + CREDENTIALS_LARGE_FRESH, + ], + gungraun_compare_by_id = true, + }, + SCANNING { + benchmarks = [ + TOKEN_15_DISPATCHED, + TOKEN_15_REFERENCE, + TOKEN_16_DISPATCHED, + TOKEN_16_REFERENCE, + TOKEN_31_DISPATCHED, + TOKEN_31_REFERENCE, + TOKEN_32_DISPATCHED, + TOKEN_32_REFERENCE, + TOKEN_33_DISPATCHED, + TOKEN_33_REFERENCE, + TOKEN_512_DISPATCHED, + TOKEN_512_REFERENCE, + TOKEN68_150_DISPATCHED, + TOKEN68_150_REFERENCE, + TOKEN68_512_DISPATCHED, + TOKEN68_512_REFERENCE, + FIELD_VALUE_150_DISPATCHED, + FIELD_VALUE_150_REFERENCE, + FIELD_VALUE_512_DISPATCHED, + FIELD_VALUE_512_REFERENCE, + DELIMITERS_SINGLE_LINE_DISPATCHED, + DELIMITERS_SINGLE_LINE_REFERENCE, + DELIMITERS_MULTI_LINE_DISPATCHED, + DELIMITERS_MULTI_LINE_REFERENCE, + COMMA_ITEMS_IRRELEVANT_BYTES, + ], + gungraun_compare_by_id = true, + }, + LIST_SCANNING { + benchmarks = [ + LIST_SCAN_15_DISPATCHED, + LIST_SCAN_15_REFERENCE, + LIST_SCAN_16_DISPATCHED, + LIST_SCAN_16_REFERENCE, + LIST_SCAN_17_DISPATCHED, + LIST_SCAN_17_REFERENCE, + LIST_SCAN_31_DISPATCHED, + LIST_SCAN_31_REFERENCE, + LIST_SCAN_32_DISPATCHED, + LIST_SCAN_32_REFERENCE, + LIST_SCAN_33_DISPATCHED, + LIST_SCAN_33_REFERENCE, + ], + gungraun_compare_by_id = true, + }, + URI_SCANNING { + benchmarks = [ + URI_PATH_15_DISPATCHED, + URI_PATH_15_REFERENCE, + URI_PATH_16_DISPATCHED, + URI_PATH_16_REFERENCE, + URI_PATH_17_DISPATCHED, + URI_PATH_17_REFERENCE, + ], + gungraun_compare_by_id = true, + }, + FORCED_URI_BACKENDS { + benchmarks = [ + URI_PATH_64_SSE2, + URI_PATH_64_SSSE3, + URI_PATH_64_SSE42, + URI_PATH_64_NEON, + ], + }, + CREDENTIAL_SCAN { + benchmarks = [ + TOKEN68_CREDENTIAL_SHARED_SCANNER, + TOKEN68_CREDENTIAL_MATCH_CHAIN, + ], + gungraun_compare_by_id = true, + }, + CUSTOM { + benchmarks = [ + CUSTOM_BORROWED_BUILTIN, + CUSTOM_BORROWED_CUSTOM, + CUSTOM_OWNED_BUILTIN, + CUSTOM_OWNED_CUSTOM, + CUSTOM_OWNED_RAW_SOURCE, + CUSTOM_INSERT_BUILTIN, + CUSTOM_INSERT_CUSTOM, + ], + gungraun_compare_by_id = true, + }, + }, +); diff --git a/crates/http_headers/benches/http_headers_name_recognition.rs b/crates/http_headers/benches/http_headers_name_recognition.rs new file mode 100644 index 000000000..d305bb998 --- /dev/null +++ b/crates/http_headers/benches/http_headers_name_recognition.rs @@ -0,0 +1,151 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Unified `http_headers` benchmarks for field-name recognition and `http` name conversion. +//! +//! Every other benchmark in this crate starts from a `&'static FieldName` +//! constant or a prebuilt map, so the recognition path that turns wire bytes +//! into a [`FieldName`] is invisible to them. This file measures it directly: +//! +//! * `parse_known_names` — the field-name set of an ordinary browser request, +//! in wire case, across the whole length range. Every name is a hit. +//! * `parse_custom_names_lowercase` / `parse_custom_names_mixed_case` — +//! matched vendor and tracing names no table entry recognizes, separating +//! already-normalized input from HTTP/1-style normalization. +//! * `custom_name_index` — the dense-index lookup a `Custom` name performs, +//! which recognizes the name a second time. +//! * `http_names_into_crate_names` / `crate_names_into_http_names` — the bulk +//! conversion an adapter performs when it moves a whole `http::HeaderMap` +//! across the crate boundary. +//! +//! Setup builds the corpora outside the measured region, so each engine +//! attributes a case to recognition and nothing else. Each case consumes its +//! result, so a shortcut that skipped the work fails instead of posting a +//! better number. + +use std::hint::black_box; + +use criterion::{BatchSize, BenchmarkId, Criterion}; +use http_headers::FieldName; + +use self::name_corpus::{crate_names, custom_header_names, custom_names_lowercase, custom_names_mixed_case, http_names, known_names}; + +#[path = "../tests/common/http_headers_name_corpus.rs"] +mod name_corpus; + +const RECOGNITION: &str = "http_headers_name_recognition/recognition"; +const HTTP_CONVERSION: &str = "http_headers_name_recognition/http_conversion"; + +#[metabench::benchmark(PARSE_KNOWN_NAMES, RECOGNITION, "parse_known_names")] +#[bench::request(setup = known_names)] +fn parse_known_names(names: &'static [&'static [u8]]) -> usize { + let mut recognized = 0; + for name in names { + let parsed = FieldName::try_from_bytes(black_box(name)).expect("valid field name"); + recognized += usize::from(parsed.index().is_some()); + } + black_box(recognized) +} + +#[metabench::benchmark(PARSE_CUSTOM_NAMES_LOWERCASE, RECOGNITION, "parse_custom_names_lowercase")] +#[bench::request(setup = custom_names_lowercase)] +fn parse_custom_names_lowercase(names: &'static [&'static [u8]]) -> usize { + parse_custom_names(names) +} + +#[metabench::benchmark(PARSE_CUSTOM_NAMES_MIXED_CASE, RECOGNITION, "parse_custom_names_mixed_case")] +#[bench::request(setup = custom_names_mixed_case)] +fn parse_custom_names_mixed_case(names: &'static [&'static [u8]]) -> usize { + parse_custom_names(names) +} + +fn parse_custom_names(names: &[&[u8]]) -> usize { + let mut length = 0; + for name in names { + let parsed = FieldName::try_from_bytes(black_box(name)).expect("valid field name"); + length += parsed.as_str().len(); + } + black_box(length) +} + +#[metabench::benchmark(CUSTOM_NAME_INDEX, RECOGNITION, "custom_name_index")] +#[bench::request(setup = custom_header_names)] +fn custom_name_index(names: &'static [FieldName]) -> usize { + let mut unknown = 0; + for name in names { + unknown += usize::from(black_box(name).index().is_none()); + } + black_box(unknown) +} + +#[metabench::benchmark(HTTP_NAMES_INTO_CRATE_NAMES, HTTP_CONVERSION, "http_names_into_crate_names")] +#[bench::request(setup = http_names)] +fn http_names_into_crate_names(names: &'static [http::HeaderName]) -> usize { + let mut length = 0; + for name in names { + let converted = FieldName::from(black_box(name)); + length += converted.as_str().len(); + } + black_box(length) +} + +#[metabench::benchmark(CRATE_NAMES_INTO_HTTP_NAMES, HTTP_CONVERSION, "crate_names_into_http_names")] +#[bench::request(setup = crate_names)] +fn crate_names_into_http_names(names: &'static [FieldName]) -> usize { + let mut length = 0; + for name in names { + let converted = http::HeaderName::from(black_box(name)); + length += converted.as_str().len(); + } + black_box(length) +} + +fn criterion_benchmarks(criterion: &mut Criterion) { + let mut recognition = criterion.benchmark_group(RECOGNITION); + recognition.bench_function(BenchmarkId::new(PARSE_KNOWN_NAMES.benchmark_name(), "request"), |bencher| { + bencher.iter_batched(known_names, parse_known_names, BatchSize::SmallInput); + }); + recognition.bench_function( + BenchmarkId::new(PARSE_CUSTOM_NAMES_LOWERCASE.benchmark_name(), "request"), + |bencher| { + bencher.iter_batched(custom_names_lowercase, parse_custom_names_lowercase, BatchSize::SmallInput); + }, + ); + recognition.bench_function( + BenchmarkId::new(PARSE_CUSTOM_NAMES_MIXED_CASE.benchmark_name(), "request"), + |bencher| { + bencher.iter_batched(custom_names_mixed_case, parse_custom_names_mixed_case, BatchSize::SmallInput); + }, + ); + recognition.bench_function(BenchmarkId::new(CUSTOM_NAME_INDEX.benchmark_name(), "request"), |bencher| { + bencher.iter_batched(custom_header_names, custom_name_index, BatchSize::SmallInput); + }); + recognition.finish(); + + let mut conversion = criterion.benchmark_group(HTTP_CONVERSION); + conversion.bench_function( + BenchmarkId::new(HTTP_NAMES_INTO_CRATE_NAMES.benchmark_name(), "request"), + |bencher| { + bencher.iter_batched(http_names, http_names_into_crate_names, BatchSize::SmallInput); + }, + ); + conversion.bench_function( + BenchmarkId::new(CRATE_NAMES_INTO_HTTP_NAMES.benchmark_name(), "request"), + |bencher| { + bencher.iter_batched(crate_names, crate_names_into_http_names, BatchSize::SmallInput); + }, + ); + conversion.finish(); +} + +metabench::main!( + criterion = criterion_benchmarks, + benchmarks = [ + PARSE_KNOWN_NAMES, + PARSE_CUSTOM_NAMES_LOWERCASE, + PARSE_CUSTOM_NAMES_MIXED_CASE, + CUSTOM_NAME_INDEX, + HTTP_NAMES_INTO_CRATE_NAMES, + CRATE_NAMES_INTO_HTTP_NAMES, + ], +); diff --git a/crates/http_headers/benches/http_headers_negotiation_semantics.rs b/crates/http_headers/benches/http_headers_negotiation_semantics.rs new file mode 100644 index 000000000..6f4595974 --- /dev/null +++ b/crates/http_headers/benches/http_headers_negotiation_semantics.rs @@ -0,0 +1,646 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Decode, repeated semantic reads and selection against fixed offer sets. +//! +//! The policy is explicit: more specific preferences override less specific +//! ones, the first duplicate wins, zero excludes an offer, and offer order +//! breaks equal-quality ties. This benchmark is not a negotiation policy API. +//! Raw cases include the caller-side parsing required by the original API. + +#![expect(clippy::unwrap_used, reason = "benchmark fixtures are independently validated during setup")] + +use std::cmp::Ordering; +use std::hint::black_box; + +use http_headers::headers::{ + Accept, AcceptEncoding, AcceptEncodingEntry, AcceptEntry, AcceptLanguage, AcceptLanguageEntry, ContentCodingKind, MediaRangeKind, + QualityView, +}; +use http_headers::sink::{EncodedValues, FieldSink, InsertError}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeMode, Field, FieldName}; + +#[derive(Clone, Copy)] +enum Family { + Media, + Coding, + Language, +} + +#[derive(Clone, Copy)] +struct Case { + family: Family, + reads: usize, +} + +struct Source { + name: &'static FieldName, + bytes: &'static [u8], +} + +impl FieldSource for Source { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::single(name, self.bytes)) + } +} + +impl Case { + fn source(self) -> Source { + match self.family { + Family::Media => Source { + name: &FieldName::Accept, + bytes: b"Text/HTML;level=\"1\";q=0.7;flag, application/json;q=0.8, text/*;q=0.9, */*;q=0.1, text/plain;q=0, application/json;q=0.8", + }, + Family::Coding => Source { + name: &FieldName::AcceptEncoding, + bytes: b"BR;q=0, gzip;q=.8000000000000000000001, identity;q=0.8, *;q=0, x-private;q=1", + }, + Family::Language => Source { + name: &FieldName::AcceptLanguage, + bytes: b"en;q=0.5, EN-us;q=0, fr;q=0.75, *;q=0.1, fr-CH;q=0.7500", + }, + } + } +} + +#[derive(Clone, Copy)] +struct Candidate { + quality: Q, + specificity: usize, + seen: bool, +} + +impl Candidate { + const fn new(zero: Q) -> Self { + Self { + quality: zero, + specificity: 0, + seen: false, + } + } + + fn consider(&mut self, quality: Q, specificity: usize) { + if !self.seen || specificity > self.specificity { + *self = Self { + quality, + specificity, + seen: true, + }; + } + } +} + +fn choose(candidates: [Candidate; 3], zero: Q) -> usize { + let mut selected = 3; + let mut quality = zero; + for (index, candidate) in candidates.into_iter().enumerate() { + if candidate.seen && candidate.quality > quality { + selected = index; + quality = candidate.quality; + } + } + selected +} + +const MEDIA: [(&str, &str); 3] = [("text", "html"), ("application", "json"), ("text", "plain")]; +const CODINGS: [&str; 3] = ["br", "gzip", "identity"]; +const LANGUAGES: [&str; 3] = ["en-US", "fr-CH", "de"]; + +fn typed_media<'a>(entries: impl Iterator>) -> usize { + let mut candidates = [Candidate::new(QualityView::ZERO); 3]; + for entry in entries { + let range = entry.range(); + let mut parameter_count = 0; + let mut parameters_match = true; + for parameter in entry.parameters() { + parameter_count += 1; + parameters_match &= parameter.name().eq_ignore_ascii_case("level") + && parameter + .value() + .is_some_and(|value| value.decoded_bytes().eq(b"1".iter().copied())); + } + for (index, (type_, subtype)) in MEDIA.into_iter().enumerate() { + if parameter_count != 0 && (index != 0 || !parameters_match) { + continue; + } + let specificity = match range.kind() { + MediaRangeKind::Any => 0, + MediaRangeKind::TypeWildcard if range.type_().eq_ignore_ascii_case(type_) => 1, + MediaRangeKind::Exact if range.type_().eq_ignore_ascii_case(type_) && range.subtype().eq_ignore_ascii_case(subtype) => { + 2 + parameter_count + } + _ => continue, + }; + candidates[index].consider(entry.quality(), specificity); + } + } + choose(candidates, QualityView::ZERO) +} + +fn language_matches(range: &str, offered: &str) -> bool { + offered.eq_ignore_ascii_case(range) + || (offered.len() > range.len() && offered.as_bytes()[range.len()] == b'-' && offered[..range.len()].eq_ignore_ascii_case(range)) +} + +fn typed_coding<'a>(entries: impl Iterator>) -> usize { + let mut candidates = [Candidate::new(QualityView::ZERO); 3]; + for entry in entries { + for (index, offer) in CODINGS.into_iter().enumerate() { + let specificity = if entry.coding().kind() == ContentCodingKind::Wildcard { + 0 + } else if entry.coding().token().eq_ignore_ascii_case(offer) { + 1 + } else { + continue; + }; + candidates[index].consider(entry.quality(), specificity); + } + } + choose(candidates, QualityView::ZERO) +} + +fn typed_language<'a>(entries: impl Iterator>) -> usize { + let mut candidates = [Candidate::new(QualityView::ZERO); 3]; + for entry in entries { + let range = entry.range(); + for (index, offer) in LANGUAGES.into_iter().enumerate() { + let specificity = if range.is_wildcard() { + 0 + } else if language_matches(range.as_str(), offer) { + 1 + range.subtags().count() + } else { + continue; + }; + candidates[index].consider(entry.quality(), specificity); + } + } + choose(candidates, QualityView::ZERO) +} + +fn typed(case: Case) -> usize { + let source = case.source(); + let mut checksum = 0; + match case.family { + Family::Media => { + let value = Accept::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap(); + for _ in 0..case.reads { + checksum += typed_media(black_box(&value).entries()); + } + black_box(value.values().next().unwrap().as_bytes()); + } + Family::Coding => { + let value = AcceptEncoding::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap(); + for _ in 0..case.reads { + checksum += typed_coding(black_box(&value).entries()); + } + black_box(value.values().next().unwrap().as_bytes()); + } + Family::Language => { + let value = AcceptLanguage::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap(); + for _ in 0..case.reads { + checksum += typed_language(black_box(&value).entries()); + } + black_box(value.values().next().unwrap().as_bytes()); + } + } + checksum +} + +#[derive(Clone, Copy, Eq)] +struct RawQuality<'a> { + whole: bool, + fraction: &'a [u8], +} + +impl<'a> RawQuality<'a> { + const ZERO: Self = Self { + whole: false, + fraction: b"", + }; + const ONE: Self = Self { + whole: true, + fraction: b"", + }; + + fn parse(bytes: &'a [u8]) -> Self { + let bytes = trim(bytes); + let whole = bytes[0] == b'1'; + let fraction = bytes + .iter() + .position(|byte| *byte == b'.') + .map_or(b"".as_slice(), |dot| &bytes[dot + 1..]); + Self { whole, fraction } + } +} + +impl PartialEq for RawQuality<'_> { + fn eq(&self, other: &Self) -> bool { + self.cmp(other) == Ordering::Equal + } +} + +impl Ord for RawQuality<'_> { + fn cmp(&self, other: &Self) -> Ordering { + let whole = self.whole.cmp(&other.whole); + if whole != Ordering::Equal { + return whole; + } + for index in 0..self.fraction.len().max(other.fraction.len()) { + let left = self.fraction.get(index).copied().unwrap_or(b'0'); + let right = other.fraction.get(index).copied().unwrap_or(b'0'); + let order = left.cmp(&right); + if order != Ordering::Equal { + return order; + } + } + Ordering::Equal + } +} + +impl PartialOrd for RawQuality<'_> { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +fn trim(bytes: &[u8]) -> &[u8] { + bytes.trim_ascii() +} + +fn parameters(bytes: &[u8]) -> impl Iterator { + let mut remaining = Some(bytes); + std::iter::from_fn(move || { + let bytes = remaining?; + let mut quoted = false; + let mut escaped = false; + for (index, byte) in bytes.iter().copied().enumerate() { + if escaped { + escaped = false; + } else if quoted && byte == b'\\' { + escaped = true; + } else if byte == b'"' { + quoted = !quoted; + } else if !quoted && byte == b';' { + remaining = Some(&bytes[index + 1..]); + return Some(trim(&bytes[..index])); + } + } + remaining = None; + Some(trim(bytes)) + }) +} + +fn split_parameter(bytes: &[u8]) -> (&[u8], Option<&[u8]>) { + bytes + .iter() + .position(|byte| *byte == b'=') + .map_or((bytes, None), |equals| (trim(&bytes[..equals]), Some(trim(&bytes[equals + 1..])))) +} + +fn decoded_equals(bytes: &[u8], expected: &[u8]) -> bool { + let quoted = bytes.first() == Some(&b'"'); + let bytes = if quoted { &bytes[1..bytes.len() - 1] } else { bytes }; + let mut bytes = bytes.iter().copied(); + std::iter::from_fn(move || { + let byte = bytes.next()?; + if quoted && byte == b'\\' { bytes.next() } else { Some(byte) } + }) + .eq(expected.iter().copied()) +} + +fn raw_selection<'a>(family: Family, items: impl Iterator) -> usize { + let mut candidates = [Candidate::new(RawQuality::ZERO); 3]; + for item in items { + let mut segments = parameters(item); + let head = segments.next().unwrap(); + let mut quality = RawQuality::ONE; + let mut parameter_count = 0; + let mut parameters_match = true; + for parameter in segments { + let (name, value) = split_parameter(parameter); + if name.eq_ignore_ascii_case(b"q") { + quality = RawQuality::parse(value.unwrap()); + break; + } + parameter_count += 1; + parameters_match &= name.eq_ignore_ascii_case(b"level") && value.is_some_and(|value| decoded_equals(value, b"1")); + } + match family { + Family::Media => { + let slash = head.iter().position(|byte| *byte == b'/').unwrap(); + let (type_, subtype) = (&head[..slash], &head[slash + 1..]); + for (index, (offered_type, offered_subtype)) in MEDIA.into_iter().enumerate() { + if parameter_count != 0 && (index != 0 || !parameters_match) { + continue; + } + let specificity = if type_ == b"*" { + 0 + } else if type_.eq_ignore_ascii_case(offered_type.as_bytes()) && subtype == b"*" { + 1 + } else if type_.eq_ignore_ascii_case(offered_type.as_bytes()) + && subtype.eq_ignore_ascii_case(offered_subtype.as_bytes()) + { + 2 + parameter_count + } else { + continue; + }; + candidates[index].consider(quality, specificity); + } + } + Family::Coding => { + for (index, offer) in CODINGS.into_iter().enumerate() { + let specificity = if head == b"*" { + 0 + } else if head.eq_ignore_ascii_case(offer.as_bytes()) { + 1 + } else { + continue; + }; + candidates[index].consider(quality, specificity); + } + } + Family::Language => { + let range = std::str::from_utf8(head).unwrap(); + for (index, offer) in LANGUAGES.into_iter().enumerate() { + let specificity = if range == "*" { + 0 + } else if language_matches(range, offer) { + range.split('-').count() + } else { + continue; + }; + candidates[index].consider(quality, specificity); + } + } + } + } + choose(candidates, RawQuality::ZERO) +} + +fn raw(case: Case) -> usize { + let source = case.source(); + let mut checksum = 0; + macro_rules! consume { + ($header:ty) => {{ + let value = <$header>::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap(); + for _ in 0..case.reads { + checksum += raw_selection(case.family, black_box(&value).items()); + } + black_box(value.values().next().unwrap().as_bytes()); + }}; + } + match case.family { + Family::Media => consume!(Accept), + Family::Coding => consume!(AcceptEncoding), + Family::Language => consume!(AcceptLanguage), + } + checksum +} + +fn retained(case: Case) -> usize { + let source = case.source(); + let mut checksum = 0; + match case.family { + Family::Media => { + let value = Accept::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap(); + for entry in value.entries() { + for _ in 0..case.reads { + let entry = black_box(entry); + checksum += entry.range().type_().as_str().len() + usize::from(!entry.quality().is_zero()); + } + } + } + Family::Coding => { + let value = AcceptEncoding::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap(); + for entry in value.entries() { + for _ in 0..case.reads { + let entry = black_box(entry); + checksum += entry.coding().as_str().len() + usize::from(!entry.quality().is_zero()); + } + } + } + Family::Language => { + let value = AcceptLanguage::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap(); + for entry in value.entries() { + for _ in 0..case.reads { + let entry = black_box(entry); + checksum += + entry.range().primary().map_or(0, |primary| primary.as_str().len()) + usize::from(!entry.quality().is_zero()); + } + } + } + } + checksum +} + +struct Forward; + +impl FieldSource for Forward { + fn lines(&self, _name: &'static FieldName) -> Option> { + None + } +} + +impl FieldSink for Forward { + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + black_box(name); + for value in values { + black_box(value); + } + Ok(()) + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.set_values(name, values) + } + + fn remove_values(&mut self, name: &'static FieldName) { + black_box(name); + } +} + +fn forward(case: Case, caller_parsed: bool) { + let source = case.source(); + macro_rules! forward { + ($header:ty, $select:ident) => {{ + let value = <$header>::owned_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap(); + let selected = if caller_parsed { + raw_selection(case.family, value.items()) + } else { + $select(value.entries()) + }; + black_box(selected); + <$header>::insert(&mut Forward, value).unwrap(); + }}; + } + match case.family { + Family::Media => forward!(Accept, typed_media), + Family::Coding => forward!(AcceptEncoding, typed_coding), + Family::Language => forward!(AcceptLanguage, typed_language), + } +} + +fn setup(family: Family, reads: usize) -> Case { + let case = Case { family, reads }; + assert_eq!(typed(case), reads, "offer one is independently expected for every fixture"); + assert_eq!(raw(case), reads, "caller-side parsing must select the same offer"); + case +} + +fn media_one() -> Case { + setup(Family::Media, 1) +} +fn media_two() -> Case { + setup(Family::Media, 2) +} +fn media_eight() -> Case { + setup(Family::Media, 8) +} +fn coding_one() -> Case { + setup(Family::Coding, 1) +} +fn coding_two() -> Case { + setup(Family::Coding, 2) +} +fn coding_eight() -> Case { + setup(Family::Coding, 8) +} +fn language_one() -> Case { + setup(Family::Language, 1) +} +fn language_two() -> Case { + setup(Family::Language, 2) +} +fn language_eight() -> Case { + setup(Family::Language, 8) +} + +#[metabench::benchmark(TYPED, "http_headers_negotiation_semantics/select", "typed")] +#[bench::media_one(setup = media_one)] +#[bench::media_two(setup = media_two)] +#[bench::media_eight(setup = media_eight)] +#[bench::coding_one(setup = coding_one)] +#[bench::coding_two(setup = coding_two)] +#[bench::coding_eight(setup = coding_eight)] +#[bench::language_one(setup = language_one)] +#[bench::language_two(setup = language_two)] +#[bench::language_eight(setup = language_eight)] +fn typed_selection(case: Case) { + black_box(typed(case)); +} + +#[metabench::benchmark(RAW, "http_headers_negotiation_semantics/select", "caller_parsed")] +#[bench::media_one(setup = media_one)] +#[bench::media_two(setup = media_two)] +#[bench::media_eight(setup = media_eight)] +#[bench::coding_one(setup = coding_one)] +#[bench::coding_two(setup = coding_two)] +#[bench::coding_eight(setup = coding_eight)] +#[bench::language_one(setup = language_one)] +#[bench::language_two(setup = language_two)] +#[bench::language_eight(setup = language_eight)] +fn raw_selection_workload(case: Case) { + black_box(raw(case)); +} + +#[metabench::benchmark(RETAINED, "http_headers_negotiation_semantics/retained", "scalar_reads")] +#[bench::media_eight(setup = media_eight)] +#[bench::coding_eight(setup = coding_eight)] +#[bench::language_eight(setup = language_eight)] +fn retained_reads(case: Case) { + black_box(retained(case)); +} + +#[metabench::benchmark(FORWARD, "http_headers_negotiation_semantics/forward", "select_and_forward")] +#[bench::media_one(setup = media_one)] +#[bench::coding_one(setup = coding_one)] +#[bench::language_one(setup = language_one)] +fn selection_forward(case: Case) { + forward(case, false); +} + +#[metabench::benchmark(RAW_FORWARD, "http_headers_negotiation_semantics/forward", "caller_parsed_and_forward")] +#[bench::media_one(setup = media_one)] +#[bench::coding_one(setup = coding_one)] +#[bench::language_one(setup = language_one)] +fn raw_selection_forward(case: Case) { + forward(case, true); +} + +#[metabench::benchmark(DECODE, "http_headers_negotiation_semantics/decode", "borrowed")] +#[bench::media_one(setup = media_one)] +#[bench::coding_one(setup = coding_one)] +#[bench::language_one(setup = language_one)] +fn decode_only(case: Case) { + let source = case.source(); + match case.family { + Family::Media => { + let _value = black_box(Accept::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap()); + } + Family::Coding => { + let _value = black_box(AcceptEncoding::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap()); + } + Family::Language => { + let _value = black_box(AcceptLanguage::view_with(black_box(&source), DecodeMode::Relaxed).unwrap().unwrap()); + } + } +} + +fn criterion_benchmarks(criterion: &mut criterion::Criterion) { + type Setup = fn() -> Case; + let cases: [(&str, Setup); 9] = [ + ("media_one", media_one), + ("media_two", media_two), + ("media_eight", media_eight), + ("coding_one", coding_one), + ("coding_two", coding_two), + ("coding_eight", coding_eight), + ("language_one", language_one), + ("language_two", language_two), + ("language_eight", language_eight), + ]; + let mut group = criterion.benchmark_group("http_headers_negotiation_semantics/select"); + for (name, setup) in cases { + group.bench_function(criterion::BenchmarkId::new(TYPED.benchmark_name(), name), |b| { + b.iter_batched(setup, typed_selection, criterion::BatchSize::SmallInput); + }); + group.bench_function(criterion::BenchmarkId::new(RAW.benchmark_name(), name), |b| { + b.iter_batched(setup, raw_selection_workload, criterion::BatchSize::SmallInput); + }); + } + group.finish(); + let mut group = criterion.benchmark_group("http_headers_negotiation_semantics/retained"); + for (name, setup) in [cases[2], cases[5], cases[8]] { + group.bench_function(criterion::BenchmarkId::new(RETAINED.benchmark_name(), name), |b| { + b.iter_batched(setup, retained_reads, criterion::BatchSize::SmallInput); + }); + } + group.finish(); + let mut group = criterion.benchmark_group("http_headers_negotiation_semantics/forward"); + for (name, setup) in [cases[0], cases[3], cases[6]] { + group.bench_function(criterion::BenchmarkId::new(FORWARD.benchmark_name(), name), |b| { + b.iter_batched(setup, selection_forward, criterion::BatchSize::SmallInput); + }); + group.bench_function(criterion::BenchmarkId::new(RAW_FORWARD.benchmark_name(), name), |b| { + b.iter_batched(setup, raw_selection_forward, criterion::BatchSize::SmallInput); + }); + } + group.finish(); + let mut group = criterion.benchmark_group("http_headers_negotiation_semantics/decode"); + for (name, setup) in [cases[0], cases[3], cases[6]] { + group.bench_function(criterion::BenchmarkId::new(DECODE.benchmark_name(), name), |b| { + b.iter_batched(setup, decode_only, criterion::BatchSize::SmallInput); + }); + } + group.finish(); +} + +metabench::main!( + criterion = { + factory = criterion::Criterion::default, + benchmarks = criterion_benchmarks, + unit = "ns", + }, + benchmarks = [TYPED, RAW, RETAINED, FORWARD, RAW_FORWARD, DECODE], +); diff --git a/crates/http_headers/benches/http_headers_negotiation_shapes.rs b/crates/http_headers/benches/http_headers_negotiation_shapes.rs new file mode 100644 index 000000000..e52caa4c4 --- /dev/null +++ b/crates/http_headers/benches/http_headers_negotiation_shapes.rs @@ -0,0 +1,250 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! HTTP-backed and raw-source decode shapes for negotiation and media-type headers. + +use http_headers::DecodeErrorKind; +use http_headers::headers::{Accept, AcceptEncoding, AcceptLanguage, Allow, ContentType, Host, Server, Vary}; + +#[path = "http_headers_shapes_common.rs"] +mod shapes; + +use shapes::Expected; + +shapes::define_shapes!( + "http_headers_negotiation_shapes/parse"; + ( + accept_canonical_browser, + Accept, + &["text/html,application/xhtml+xml,application/xml;q=0.9,image/avif,image/webp,image/apng,*/*;q=0.8,application/signed-exchange;v=b3;q=0.7"], + Strict, + Expected::Valid + ), + (accept_short_html, Accept, &["text/html"], Strict, Expected::Valid), + (accept_common_api, Accept, &["application/json, text/plain;q=0.9, */*;q=0.8"], Strict, Expected::Valid), + (accept_repeated, Accept, &["text/html", "application/json;q=0.9", "*/*;q=0.1"], Strict, Expected::Valid), + ( + accept_quoted_comma, + Accept, + &["text/html;level=1, application/json;q=0.9;profile=\"a,b\"", "*/*;q=0.1"], + Strict, + Expected::Valid + ), + (accept_extension_flag, Accept, &["text/html;q=0.8;preview"], Strict, Expected::Valid), + (accept_relaxed_quality, Accept, &["text/html; q = .12345"], Relaxed, Expected::Valid), + ( + accept_strict_quality_whitespace, + Accept, + &["text/html; q = .12345"], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + (accept_bad_wildcard, Accept, &["*/json"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + ( + accept_late_unterminated_quote, + Accept, + &["text/html", "application/json;profile=\"unfinished"], + Strict, + Expected::Error(DecodeErrorKind::UnterminatedQuote) + ), + (accept_empty, Accept, &[""], Strict, Expected::Valid), + (accept_absent, Accept, &[], Strict, Expected::Absent), + (accept_encoding_canonical, AcceptEncoding, &["gzip, deflate, br, zstd"], Strict, Expected::Valid), + (accept_encoding_short, AcceptEncoding, &["gzip"], Strict, Expected::Valid), + ( + accept_encoding_weighted, + AcceptEncoding, + &["gzip;q=1.0, identity;q=0.5, *;q=0"], + Strict, + Expected::Valid + ), + ( + accept_encoding_long_preferences, + AcceptEncoding, + &["br;q=1.0, zstd;q=0.9, gzip;q=0.8, deflate;q=0.7, identity;q=0.5, *;q=0"], + Strict, + Expected::Valid + ), + (accept_encoding_repeated, AcceptEncoding, &["gzip", "br;q=0.8", "identity;q=0.5"], Strict, Expected::Valid), + (accept_encoding_relaxed_quality, AcceptEncoding, &["br; q = .1234"], Relaxed, Expected::Valid), + ( + accept_encoding_strict_quality_whitespace, + AcceptEncoding, + &["br; q = .1234"], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + accept_encoding_late_bad_token, + AcceptEncoding, + &["gzip", "bad encoding"], + Strict, + Expected::Error(DecodeErrorKind::InvalidToken) + ), + (accept_encoding_empty, AcceptEncoding, &[""], Strict, Expected::Valid), + (accept_encoding_absent, AcceptEncoding, &[], Strict, Expected::Absent), + ( + accept_language_canonical, + AcceptLanguage, + &["en-US,en;q=0.9,fr-FR;q=0.8,fr;q=0.7"], + Strict, + Expected::Valid + ), + (accept_language_short, AcceptLanguage, &["en"], Strict, Expected::Valid), + ( + accept_language_long_preferences, + AcceptLanguage, + &["zh-Hans-CN, zh-Hans;q=0.9, en-US;q=0.8, en;q=0.7, fr-CH;q=0.6, de;q=0.5, *;q=0.1"], + Strict, + Expected::Valid + ), + (accept_language_repeated, AcceptLanguage, &["en-US", "fr-CH;q=0.8", "de;q=0.5"], Strict, Expected::Valid), + (accept_language_wildcard, AcceptLanguage, &["*"], Strict, Expected::Valid), + (accept_language_relaxed_quality, AcceptLanguage, &["en-US; Q = 1.0000"], Relaxed, Expected::Valid), + ( + accept_language_strict_quality_whitespace, + AcceptLanguage, + &["en-US; Q = 1.0000"], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + accept_language_late_bad_subtag, + AcceptLanguage, + &["en-US", "en-123456789"], + Strict, + Expected::Error(DecodeErrorKind::InvalidToken) + ), + (accept_language_empty, AcceptLanguage, &[""], Strict, Expected::Valid), + (accept_language_absent, AcceptLanguage, &[], Strict, Expected::Absent), + (allow_canonical, Allow, &["GET, POST"], Strict, Expected::Valid), + (allow_short, Allow, &["GET"], Strict, Expected::Valid), + ( + allow_long_webdav, + Allow, + &["OPTIONS, GET, HEAD, POST, PUT, DELETE, TRACE, PROPFIND, PROPPATCH, MKCOL, COPY, MOVE, LOCK, UNLOCK"], + Strict, + Expected::Valid + ), + (allow_repeated, Allow, &["GET, HEAD", "POST", "PATCH, DELETE"], Strict, Expected::Valid), + (allow_extension, Allow, &["PROPFIND"], Strict, Expected::Valid), + (allow_ows_empty_members, Allow, &[" , GET,,\tPOST, "], Strict, Expected::Valid), + ( + allow_late_bad_token, + Allow, + &["GET", "BAD METHOD"], + Strict, + Expected::Error(DecodeErrorKind::InvalidToken) + ), + (allow_empty, Allow, &[""], Strict, Expected::Valid), + (allow_absent, Allow, &[], Strict, Expected::Absent), + (vary_canonical, Vary, &["accept-encoding, origin"], Strict, Expected::Valid), + (vary_short, Vary, &["accept-encoding"], Strict, Expected::Valid), + (vary_title_case, Vary, &["Accept-Encoding, Origin"], Strict, Expected::Valid), + ( + vary_long_selection, + Vary, + &["Accept-Encoding, Accept-Language, Origin, User-Agent, X-Requested-With, X-API-Version"], + Strict, + Expected::Valid + ), + (vary_repeated, Vary, &["accept-encoding", "Origin", "X-API-Version"], Strict, Expected::Valid), + (vary_wildcard, Vary, &["*"], Strict, Expected::Valid), + (vary_ows_empty_members, Vary, &[" , Accept-Encoding,,\tOrigin, "], Strict, Expected::Valid), + ( + vary_late_bad_token, + Vary, + &["accept-encoding", "bad field"], + Strict, + Expected::Error(DecodeErrorKind::InvalidToken) + ), + (vary_empty, Vary, &[""], Strict, Expected::Valid), + (vary_absent, Vary, &[], Strict, Expected::Absent), + (host_canonical, Host, &["example.com:8443"], Strict, Expected::Valid), + (host_common_domain, Host, &["api.example.com"], Strict, Expected::Valid), + ( + host_long_domain, + Host, + &["api.customer-a.region-west.service.production.widgets.example.com:8443"], + Strict, + Expected::Valid + ), + (host_ipv6, Host, &["[2001:db8::1]:443"], Strict, Expected::Valid), + (host_ipvfuture, Host, &["[v1.fe80::a]:443"], Strict, Expected::Valid), + (host_percent_encoded, Host, &["api%2Eexample.com:443"], Strict, Expected::Valid), + (host_empty_port, Host, &["example.com:"], Strict, Expected::Valid), + (host_relaxed_ascii, Host, &["www.example.com:443"], Relaxed, Expected::Valid), + (host_bad_port, Host, &["example.com:http"], Strict, Expected::Error(DecodeErrorKind::InvalidNumber)), + ( + host_repeated, + Host, + &["example.com", "example.org"], + Strict, + Expected::Error(DecodeErrorKind::UnexpectedMultipleValues) + ), + (host_empty, Host, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (host_absent, Host, &[], Strict, Expected::Absent), + (server_canonical, Server, &["example/1.0"], Strict, Expected::Valid), + (server_comment, Server, &["Apache/2.4.58 (Unix)"], Strict, Expected::Valid), + ( + server_long_products, + Server, + &["Apache/2.4.58 (Unix) OpenSSL/3.0.13 example-proxy/2.0 (public synthetic fixture)"], + Strict, + Expected::Valid + ), + (server_leading_ows, Server, &[" \tnginx/1.25.3 \t"], Strict, Expected::Valid), + ( + server_repeated, + Server, + &["example/1.0", "example-proxy/2.0"], + Strict, + Expected::Error(DecodeErrorKind::UnexpectedMultipleValues) + ), + (server_empty, Server, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (server_ows_only, Server, &[" \t \t"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (server_absent, Server, &[], Strict, Expected::Absent), + (content_type_canonical, ContentType, &["application/json; charset=utf-8"], Strict, Expected::Valid), + (content_type_common_json, ContentType, &["application/json"], Strict, Expected::Valid), + (content_type_case_variant, ContentType, &["Application/JSON; Charset=UTF-8"], Strict, Expected::Valid), + (content_type_two_parameters, ContentType, &["text/html; charset=utf-8; level=1"], Strict, Expected::Valid), + ( + content_type_many_parameters, + ContentType, + &["multipart/form-data; boundary=------------------------1a2b3c; charset=utf-8; name=\"upload\""], + Strict, + Expected::Valid + ), + ( + content_type_quoted_boundary, + ContentType, + &["multipart/form-data; boundary=\"example boundary, version=1\""], + Strict, + Expected::Valid + ), + (content_type_empty_slots, ContentType, &["text/plain; ; charset=utf-8;"], Strict, Expected::Valid), + (content_type_relaxed_slash, ContentType, &["text / html; charset=utf-8"], Relaxed, Expected::Valid), + ( + content_type_missing_slash, + ContentType, + &["application"], + Strict, + Expected::Error(DecodeErrorKind::InvalidSyntax) + ), + ( + content_type_unterminated_quote, + ContentType, + &["text/plain; charset=\"unfinished"], + Strict, + Expected::Error(DecodeErrorKind::UnterminatedQuote) + ), + ( + content_type_repeated, + ContentType, + &["text/plain", "application/json"], + Strict, + Expected::Error(DecodeErrorKind::UnexpectedMultipleValues) + ), + (content_type_empty, ContentType, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (content_type_absent, ContentType, &[], Strict, Expected::Absent), +); diff --git a/crates/http_headers/benches/http_headers_operations.rs b/crates/http_headers/benches/http_headers_operations.rs new file mode 100644 index 000000000..7e843948a --- /dev/null +++ b/crates/http_headers/benches/http_headers_operations.rs @@ -0,0 +1,171 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Shared `http_headers` operations compiled into both timing and instruction harnesses. + +use std::hint::black_box; + +use http::HeaderMap; +use http_headers::Field; +use http_headers::headers::{AcceptRanges, AccessControlAllowMethods, Range, Vary}; + +#[expect( + clippy::inline_always, + reason = "destruction is part of the operation while the operation boundary remains outlined" +)] +#[inline(always)] +fn consume(value: T) { + drop(black_box(value)); +} + +#[expect( + clippy::inline_always, + reason = "iteration is part of the operation while the operation boundary remains outlined" +)] +#[inline(always)] +fn consume_items(items: impl IntoIterator) -> usize { + let mut count = 0; + for item in items { + black_box(item); + count += 1; + } + black_box(count) +} + +#[inline(never)] +pub(crate) fn headers_owned(map: &HeaderMap) { + let header = headers::HeaderMapExt::typed_try_get::(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + consume(header); +} + +#[inline(never)] +pub(crate) fn http_headers_owned(map: &HeaderMap) { + let header = H::owned(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + consume(header); +} + +#[inline(never)] +pub(crate) fn http_headers_borrowed(map: &HeaderMap) { + let header = H::view(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + consume(header); +} + +#[inline(never)] +pub(crate) fn headers_range_owned(map: &HeaderMap) { + let header = headers::HeaderMapExt::typed_try_get::(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + consume(header); +} + +#[inline(never)] +pub(crate) fn http_headers_range_owned(map: &HeaderMap) { + let header = Range::owned(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + consume(header); +} + +#[inline(never)] +pub(crate) fn http_headers_range_borrowed(map: &HeaderMap) { + let header = Range::view(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + consume(header); +} + +#[inline(never)] +pub(crate) fn headers_accept_ranges(map: &HeaderMap) -> usize { + let header = headers::HeaderMapExt::typed_try_get::(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let result = usize::from(header.is_bytes()) | (usize::from(header.is_none()) << 1); + consume(header); + black_box(result) +} + +#[inline(never)] +pub(crate) fn http_headers_accept_ranges_owned(map: &HeaderMap) -> usize { + let header = AcceptRanges::owned(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let result = usize::from(header.units().eq(["bytes"])) | (usize::from(header.is_none()) << 1); + consume(header); + black_box(result) +} + +#[inline(never)] +pub(crate) fn http_headers_accept_ranges_borrowed(map: &HeaderMap) -> usize { + let header = AcceptRanges::view(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let result = usize::from(header.units().eq(["bytes"])) | (usize::from(header.is_none()) << 1); + consume(header); + black_box(result) +} + +#[inline(never)] +pub(crate) fn headers_allow_methods(map: &HeaderMap) -> usize { + let header = headers::HeaderMapExt::typed_try_get::(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let count = consume_items(header.iter()); + consume(header); + count +} + +#[inline(never)] +pub(crate) fn http_headers_allow_methods_owned(map: &HeaderMap) -> usize { + let header = AccessControlAllowMethods::owned(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let count = consume_items(header.iter()); + consume(header); + count +} + +#[inline(never)] +pub(crate) fn http_headers_allow_methods_borrowed(map: &HeaderMap) -> usize { + let header = AccessControlAllowMethods::view(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let count = consume_items(header.iter()); + consume(header); + count +} + +#[inline(never)] +pub(crate) fn headers_vary(map: &HeaderMap) -> usize { + let header = headers::HeaderMapExt::typed_try_get::(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let count = consume_items(header.iter_strs()); + consume(header); + count +} + +#[inline(never)] +pub(crate) fn http_headers_vary_owned(map: &HeaderMap) -> usize { + let header = Vary::owned(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let count = consume_items(header.items()); + consume(header); + count +} + +#[inline(never)] +pub(crate) fn http_headers_vary_borrowed(map: &HeaderMap) -> usize { + let header = Vary::view(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let count = consume_items(header.items()); + consume(header); + count +} diff --git a/crates/http_headers/benches/http_headers_per_header.rs b/crates/http_headers/benches/http_headers_per_header.rs new file mode 100644 index 000000000..5ed768656 --- /dev/null +++ b/crates/http_headers/benches/http_headers_per_header.rs @@ -0,0 +1,685 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Unified per-header comparisons against `headers` 0.4.1. + +use std::hint::black_box; +use std::sync::OnceLock; +use std::time::Duration; + +use criterion::{BatchSize, BenchmarkId, Criterion}; +use headers::HeaderMapExt as TheirMapExt; +use http::{HeaderMap, HeaderName, HeaderValue}; +use http_headers::Field; +use http_headers::headers::{ + Accept, AcceptEncoding, AcceptLanguage, AccessControlAllowCredentials, AccessControlAllowHeaders, AccessControlAllowOrigin, + AccessControlExposeHeaders, AccessControlMaxAge, AccessControlRequestHeaders, AccessControlRequestMethod, Allow, Authorization, Basic, + Bearer, CacheControl, ContentLength, ContentRange, ContentSecurityPolicy, ContentType, ETag, Host, IfMatch, IfModifiedSince, + IfNoneMatch, IfRange, IfUnmodifiedSince, LastModified, Location, ReferrerPolicy, SecWebSocketAccept, SecWebSocketExtensions, + SecWebSocketKey, SecWebSocketProtocol, SecWebSocketVersion, Server, SetCookie, StrictTransportSecurity, UserAgent, XContentTypeOptions, +}; +use paste::paste; + +#[path = "http_headers_common_values.rs"] +mod common_values; +#[expect( + dead_code, + reason = "the shared operations module also supports benchmark targets with different inventories" +)] +#[path = "http_headers_operations.rs"] +mod operations; + +fn map(name: &'static str, values: &'static [&'static str]) -> HeaderMap { + let mut map = HeaderMap::with_capacity(1); + let name = HeaderName::from_static(name); + for value in values { + let mut value = HeaderValue::from_bytes(value.as_bytes()).expect("valid benchmark fixture"); + value.set_sensitive(matches!(name.as_str(), "authorization" | "set-cookie")); + map.append(&name, value); + } + map +} + +fn consume(value: T) { + drop(black_box(value)); +} + +fn tuned_criterion() -> Criterion { + Criterion::default() + .warm_up_time(Duration::from_secs(1)) + .measurement_time(Duration::from_secs(3)) + .sample_size(60) +} + +macro_rules! compare_header { + ($id:ident, $name:literal, $values:expr, $ours:ty, $theirs:ty) => { + paste! { + fn [<$id _map>]() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map($name, $values)) + } + + fn [<$id _headers>](map: &'static HeaderMap) { + operations::headers_owned::<$theirs>(map); + } + + fn [<$id _owned>](map: &'static HeaderMap) { + operations::http_headers_owned::<$ours>(map); + } + + fn [<$id _borrowed>](map: &'static HeaderMap) { + operations::http_headers_borrowed::<$ours>(map); + } + } + }; +} + +macro_rules! compare_tag_header { + ( + $id:ident, + $name:literal, + $values:expr, + $ours:ty, + $theirs:ty, + ours = |$header:ident| $ours_read:expr + ) => { + paste! { + fn [<$id _map>]() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map($name, $values)) + } + + fn [<$id _headers>](map: &'static HeaderMap) -> bool { + static COMPARISON: OnceLock = OnceLock::new(); + let comparison = COMPARISON.get_or_init(|| { + "\"benchmark-never-matches\"".parse().expect("valid comparison ETag") + }); + let header = TheirMapExt::typed_try_get::<$theirs>(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let result = header.precondition_passes(black_box(comparison)); + consume(header); + black_box(result) + } + + fn [<$id _owned>](map: &'static HeaderMap) -> bool { + let $header = <$ours as Field>::owned(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let result = $ours_read; + consume($header); + black_box(result) + } + + fn [<$id _borrowed>](map: &'static HeaderMap) -> bool { + let $header = <$ours as Field>::view(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let result = $ours_read; + consume($header); + black_box(result) + } + } + }; +} + +macro_rules! compare_iter_header { + ( + $id:ident, + $name:literal, + $values:expr, + $ours:ty, + $theirs:ty, + ours = $our_iter:ident, + theirs = $their_iter:ident + ) => { + paste! { + fn [<$id _map>]() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map($name, $values)) + } + + fn [<$id _headers>](map: &'static HeaderMap) -> usize { + let header = TheirMapExt::typed_try_get::<$theirs>(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let count = header.$their_iter().map(black_box).count(); + consume(header); + black_box(count) + } + + fn [<$id _owned>](map: &'static HeaderMap) -> usize { + let header = <$ours as Field>::owned(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let count = header.$our_iter().map(black_box).count(); + consume(header); + black_box(count) + } + + fn [<$id _borrowed>](map: &'static HeaderMap) -> usize { + let header = <$ours as Field>::view(black_box(map)) + .expect("fixture must decode") + .expect("fixture must be present"); + let count = header.$our_iter().map(black_box).count(); + consume(header); + black_box(count) + } + } + }; +} + +macro_rules! compare_shared_iter_header { + ( + $id:ident, + $name:literal, + $values:expr, + headers = $headers:path, + owned = $owned:path, + borrowed = $borrowed:path + ) => { + paste! { + fn [<$id _map>]() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map($name, $values)) + } + + fn [<$id _headers>](map: &'static HeaderMap) -> usize { + $headers(map) + } + + fn [<$id _owned>](map: &'static HeaderMap) -> usize { + $owned(map) + } + + fn [<$id _borrowed>](map: &'static HeaderMap) -> usize { + $borrowed(map) + } + } + }; +} + +macro_rules! measure_header { + ($id:ident, $name:literal, $values:expr, $ours:ty) => { + paste! { + fn [<$id _map>]() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map($name, $values)) + } + + fn [<$id _owned>](map: &'static HeaderMap) { + operations::http_headers_owned::<$ours>(map); + } + + fn [<$id _borrowed>](map: &'static HeaderMap) { + operations::http_headers_borrowed::<$ours>(map); + } + } + }; +} + +measure_header!( + accept, + "accept", + &[ + "text/html,application/xhtml+xml,application/xml;q=0.9,image/avif,image/webp,image/apng,*/*;q=0.8,application/signed-exchange;v=b3;q=0.7" + ], + Accept +); +measure_header!(accept_encoding, "accept-encoding", &["gzip, deflate, br, zstd"], AcceptEncoding); +measure_header!( + accept_language, + "accept-language", + &["en-US,en;q=0.9,fr-FR;q=0.8,fr;q=0.7"], + AcceptLanguage +); +compare_shared_iter_header!( + accept_ranges, + "accept-ranges", + &["bytes"], + headers = operations::headers_accept_ranges, + owned = operations::http_headers_accept_ranges_owned, + borrowed = operations::http_headers_accept_ranges_borrowed +); +compare_header!( + access_control_allow_credentials, + "access-control-allow-credentials", + &["true"], + AccessControlAllowCredentials, + headers::AccessControlAllowCredentials +); +compare_iter_header!( + access_control_allow_headers, + "access-control-allow-headers", + &["content-type, x-request-id"], + AccessControlAllowHeaders, + headers::AccessControlAllowHeaders, + ours = iter, + theirs = iter +); +compare_shared_iter_header!( + access_control_allow_methods, + "access-control-allow-methods", + &["GET, POST"], + headers = operations::headers_allow_methods, + owned = operations::http_headers_allow_methods_owned, + borrowed = operations::http_headers_allow_methods_borrowed +); +compare_header!( + access_control_allow_origin, + "access-control-allow-origin", + &["https://example.com"], + AccessControlAllowOrigin, + headers::AccessControlAllowOrigin +); +compare_iter_header!( + access_control_expose_headers, + "access-control-expose-headers", + &["etag, x-request-id"], + AccessControlExposeHeaders, + headers::AccessControlExposeHeaders, + ours = iter, + theirs = iter +); +compare_header!( + access_control_max_age, + "access-control-max-age", + &["600"], + AccessControlMaxAge, + headers::AccessControlMaxAge +); +compare_iter_header!( + access_control_request_headers, + "access-control-request-headers", + &["content-type, x-request-id"], + AccessControlRequestHeaders, + headers::AccessControlRequestHeaders, + ours = iter, + theirs = iter +); +compare_header!( + access_control_request_method, + "access-control-request-method", + &["POST"], + AccessControlRequestMethod, + headers::AccessControlRequestMethod +); +compare_iter_header!(allow, "allow", &["GET, POST"], Allow, headers::Allow, ours = items, theirs = iter); +compare_header!( + authorization_basic, + "authorization", + &[common_values::BASIC_AUTHORIZATION], + Authorization, + headers::Authorization +); +compare_header!( + authorization_bearer, + "authorization", + &[ + "Bearer eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJzdWIiOiIxMjM0NTY3ODkwIiwibmFtZSI6IkpvaG4gRG9lIiwiaWF0IjoxNTE2MjM5MDIyLCJleHAiOjE1MTYyNDI2MjIsImF1ZCI6Imh0dHBzOi8vYXBpLmV4YW1wbGUuY29tIiwiaXNzIjoiaHR0cHM6Ly9hdXRoLmV4YW1wbGUuY29tIiwic2NvcGUiOiJyZWFkOnByb2ZpbGUgd3JpdGU6cHJvZmlsZSJ9.dBjftJeZ4CVPmB92K27uhbUJU1p1r_wW1gFWFOEjXkw" + ], + Authorization, + headers::Authorization +); +compare_header!( + cache_control, + "cache-control", + &["max-age=3600, private"], + CacheControl, + headers::CacheControl +); +compare_header!(content_length, "content-length", &["348"], ContentLength, headers::ContentLength); +compare_header!( + content_range, + "content-range", + &["bytes 0-499/1234"], + ContentRange, + headers::ContentRange +); +measure_header!( + content_security_policy, + "content-security-policy", + &[ + "default-src 'self'; script-src 'self' 'unsafe-inline' https://cdn.example.com https://analytics.example.com; style-src 'self' 'unsafe-inline' https://fonts.googleapis.com; img-src 'self' data: https:; font-src 'self' https://fonts.gstatic.com; connect-src 'self' https://api.example.com; frame-ancestors 'none'; base-uri 'self'; form-action 'self'" + ], + ContentSecurityPolicy +); +compare_header!( + content_type, + "content-type", + &["application/json; charset=utf-8"], + ContentType, + headers::ContentType +); +compare_header!(etag, "etag", &["\"revision-42\""], ETag, headers::ETag); +compare_header!(host, "host", &["example.com:8443"], Host, headers::Host); +compare_tag_header!( + if_match, + "if-match", + &["\"a\", W/\"b\""], + IfMatch, + headers::IfMatch, + ours = |header| header.is_wildcard() + || header + .tags() + .any(|tag| { !tag.is_weak() && tag.opaque_tag() == black_box(b"benchmark-never-matches") }) +); +compare_header!( + if_modified_since, + "if-modified-since", + &["Sun, 06 Nov 1994 08:49:37 GMT"], + IfModifiedSince, + headers::IfModifiedSince +); +compare_tag_header!( + if_none_match, + "if-none-match", + &["W/\"a\", \"b\""], + IfNoneMatch, + headers::IfNoneMatch, + ours = |header| !(header.is_wildcard() || header.tags().any(|tag| tag.opaque_tag() == black_box(b"benchmark-never-matches"))) +); +compare_header!(if_range, "if-range", &["\"revision-42\""], IfRange, headers::IfRange); +compare_header!( + if_unmodified_since, + "if-unmodified-since", + &["Sun, 06 Nov 1994 08:49:37 GMT"], + IfUnmodifiedSince, + headers::IfUnmodifiedSince +); +compare_header!( + last_modified, + "last-modified", + &["Sun, 06 Nov 1994 08:49:37 GMT"], + LastModified, + headers::LastModified +); +compare_header!( + location, + "location", + &["https://example.com/en-us/docs/reference/index.html?utm_source=newsletter&utm_campaign=spring&page=3"], + Location, + headers::Location +); +compare_header!( + range, + "range", + &["bytes=0-499, 1000-"], + http_headers::headers::Range, + headers::Range +); +compare_header!( + referrer_policy, + "referrer-policy", + &["strict-origin-when-cross-origin"], + ReferrerPolicy, + headers::ReferrerPolicy +); +compare_header!( + sec_websocket_accept, + "sec-websocket-accept", + &["s3pPLMBiTxaQ9kYGzzhZRbK+xOo="], + SecWebSocketAccept, + headers::SecWebsocketAccept +); +measure_header!( + sec_websocket_extensions, + "sec-websocket-extensions", + &["permessage-deflate; client_max_window_bits"], + SecWebSocketExtensions +); +compare_header!( + sec_websocket_key, + "sec-websocket-key", + &["dGhlIHNhbXBsZSBub25jZQ=="], + SecWebSocketKey, + headers::SecWebsocketKey +); +measure_header!( + sec_websocket_protocol, + "sec-websocket-protocol", + &["chat, superchat"], + SecWebSocketProtocol +); +compare_header!( + sec_websocket_version, + "sec-websocket-version", + &["13"], + SecWebSocketVersion, + headers::SecWebsocketVersion +); +measure_header!( + sec_websocket_version_advertisement, + "sec-websocket-version", + &["7, 8, 13"], + SecWebSocketVersion +); +compare_header!(server, "server", &["example/1.0"], Server, headers::Server); +compare_header!( + set_cookie, + "set-cookie", + &["session=eyJhbGciOiJIUzI1NiJ9.eyJzdWIiOiIxMjM0NTY3ODkwIn0; Path=/; Domain=example.com; Max-Age=3600; Secure; HttpOnly; SameSite=Lax"], + SetCookie, + headers::SetCookie +); +compare_header!( + strict_transport_security, + "strict-transport-security", + &["max-age=31536000; includeSubDomains"], + StrictTransportSecurity, + headers::StrictTransportSecurity +); +compare_header!( + user_agent, + "user-agent", + &["Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"], + UserAgent, + headers::UserAgent +); +compare_shared_iter_header!( + vary, + "vary", + &["accept-encoding, origin"], + headers = operations::headers_vary, + owned = operations::http_headers_vary_owned, + borrowed = operations::http_headers_vary_borrowed +); +measure_header!(x_content_type_options, "x-content-type-options", &["nosniff"], XContentTypeOptions); + +macro_rules! finish { + ( + supported = [$($supported:ident),+ $(,)?], + ours = [$($ours:ident),+ $(,)?], + ) => { + paste! { + #[derive(Clone, Copy)] + struct BenchmarkCase { + map: &'static HeaderMap, + operation: fn(&'static HeaderMap), + } + + impl BenchmarkCase { + fn run(self) { + (self.operation)(self.map); + } + } + + $( + fn [<$supported _headers_case>]() -> BenchmarkCase { + BenchmarkCase { + map: [<$supported _map>](), + operation: |map| consume([<$supported _headers>](map)), + } + } + + fn [<$supported _owned_case>]() -> BenchmarkCase { + BenchmarkCase { + map: [<$supported _map>](), + operation: |map| consume([<$supported _owned>](map)), + } + } + + fn [<$supported _borrowed_case>]() -> BenchmarkCase { + BenchmarkCase { + map: [<$supported _map>](), + operation: |map| consume([<$supported _borrowed>](map)), + } + } + )+ + + $( + fn [<$ours _owned_case>]() -> BenchmarkCase { + BenchmarkCase { + map: [<$ours _map>](), + operation: |map| consume([<$ours _owned>](map)), + } + } + + fn [<$ours _borrowed_case>]() -> BenchmarkCase { + BenchmarkCase { + map: [<$ours _map>](), + operation: |map| consume([<$ours _borrowed>](map)), + } + } + )+ + + #[metabench::benchmark(COMPETITOR, "http_headers_per_header/per_header", "headers")] + $(#[bench::$supported(setup = [<$supported _headers_case>])])+ + fn competitor(case: BenchmarkCase) { + case.run(); + } + + #[metabench::benchmark(OWNED, "http_headers_per_header/per_header", "http_headers_owned")] + $(#[bench::$supported(setup = [<$supported _owned_case>])])+ + $(#[bench::$ours(setup = [<$ours _owned_case>])])+ + fn owned(case: BenchmarkCase) { + case.run(); + } + + #[metabench::benchmark(BORROWED, "http_headers_per_header/per_header", "http_headers_borrowed")] + $(#[bench::$supported(setup = [<$supported _borrowed_case>])])+ + $(#[bench::$ours(setup = [<$ours _borrowed_case>])])+ + fn borrowed(case: BenchmarkCase) { + case.run(); + } + + fn criterion_benchmarks(criterion: &mut Criterion) { + let mut group = criterion.benchmark_group("http_headers_per_header/per_header"); + $( + group.bench_function( + BenchmarkId::new(COMPETITOR.benchmark_name(), stringify!($supported)), + |bencher| { + bencher.iter_batched( + [<$supported _headers_case>], + competitor, + BatchSize::SmallInput, + ); + }, + ); + group.bench_function( + BenchmarkId::new(OWNED.benchmark_name(), stringify!($supported)), + |bencher| { + bencher.iter_batched( + [<$supported _owned_case>], + owned, + BatchSize::SmallInput, + ); + }, + ); + group.bench_function( + BenchmarkId::new(BORROWED.benchmark_name(), stringify!($supported)), + |bencher| { + bencher.iter_batched( + [<$supported _borrowed_case>], + borrowed, + BatchSize::SmallInput, + ); + }, + ); + )+ + $( + group.bench_function( + BenchmarkId::new(OWNED.benchmark_name(), stringify!($ours)), + |bencher| { + bencher.iter_batched( + [<$ours _owned_case>], + owned, + BatchSize::SmallInput, + ); + }, + ); + group.bench_function( + BenchmarkId::new(BORROWED.benchmark_name(), stringify!($ours)), + |bencher| { + bencher.iter_batched( + [<$ours _borrowed_case>], + borrowed, + BatchSize::SmallInput, + ); + }, + ); + )+ + group.finish(); + } + + metabench::main!( + criterion = { + factory = tuned_criterion, + benchmarks = criterion_benchmarks, + unit = "ns", + }, + benchmarks = [COMPETITOR, OWNED, BORROWED], + ); + } + }; +} + +finish!( + supported = [ + accept_ranges, + access_control_allow_credentials, + access_control_allow_headers, + access_control_allow_methods, + access_control_allow_origin, + access_control_expose_headers, + access_control_max_age, + access_control_request_headers, + access_control_request_method, + allow, + authorization_basic, + authorization_bearer, + cache_control, + content_length, + content_range, + content_type, + etag, + host, + if_match, + if_modified_since, + if_none_match, + if_range, + if_unmodified_since, + last_modified, + location, + range, + referrer_policy, + sec_websocket_accept, + sec_websocket_key, + sec_websocket_version, + server, + set_cookie, + strict_transport_security, + user_agent, + vary, + ], + ours = [ + accept, + accept_encoding, + accept_language, + content_security_policy, + sec_websocket_extensions, + sec_websocket_protocol, + sec_websocket_version_advertisement, + x_content_type_options, + ], +); diff --git a/crates/http_headers/benches/http_headers_policy.rs b/crates/http_headers/benches/http_headers_policy.rs new file mode 100644 index 000000000..983b0f3f7 --- /dev/null +++ b/crates/http_headers/benches/http_headers_policy.rs @@ -0,0 +1,257 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Focused policy-header instruction-count experiments. + +#![expect( + clippy::unwrap_used, + reason = "fixed benchmark fixtures are asserted valid outside the measured operations" +)] + +use std::hint::black_box; +use std::sync::OnceLock; + +use criterion::{BatchSize, Criterion}; +use http::{HeaderMap, HeaderValue}; +use http_headers::headers::{ + Authorization, AuthorizationOwned, Basic, BasicCredentials, CacheControl, ContentSecurityPolicy, Location, ReferrerPolicy, SetCookie, + StrictTransportSecurity, +}; +use http_headers::{DecodeMode, Field}; + +fn map(name: &'static str, values: &[&'static str]) -> HeaderMap { + let mut map = HeaderMap::with_capacity(1); + for value in values { + map.append(name, HeaderValue::from_static(value)); + } + map +} + +fn referrer_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| { + map( + "referrer-policy", + &["future-policy, no-referrer", "origin, strict-origin-when-cross-origin"], + ) + }) +} + +fn cache_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| { + map( + "cache-control", + &["public, max-age=31536000, stale-while-revalidate=60, x-build=123456789012345678901"], + ) + }) +} + +fn hsts_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map("strict-transport-security", &["MAX-AGE=31536000; includeSubDomains; preload"])) +} + +fn location_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map("location", &["https://example.com/a%20path?query=1#fragment"])) +} + +fn repeated_csp_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| { + map( + "content-security-policy", + &[ + "default-src 'self'", + "script-src 'none'", + "img-src https:", + "style-src 'self'", + "font-src https:", + "connect-src 'self'", + "frame-src 'none'", + "object-src 'none'", + "base-uri 'self'", + "form-action 'self'", + "frame-ancestors 'none'", + "upgrade-insecure-requests", + "block-all-mixed-content", + "worker-src 'self'", + "manifest-src 'self'", + "media-src 'none'", + ], + ) + }) +} + +fn repeated_cookie_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| { + map( + "set-cookie", + &[ + "a=1", "b=2", "c=3", "d=4", "e=5", "f=6", "g=7", "h=8", "i=9", "j=10", "k=11", "l=12", "m=13", "n=14", "o=15", "p=16", + ], + ) + }) +} + +fn basic() -> AuthorizationOwned { + AuthorizationOwned::::basic(b"benchmark-user", b"benchmark-password").unwrap() +} + +fn authorization_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map("authorization", &["Basic YmVuY2htYXJrLXVzZXI6YmVuY2htYXJrLXBhc3N3b3Jk"])) +} + +fn basic_and_credentials() -> (AuthorizationOwned, BasicCredentials) { + (basic(), BasicCredentials::new()) +} + +fn drop_it(value: T) { + drop(value); +} + +#[metabench::benchmark(REFERRER_DECODE, "policy", "referrer_decode", gungraun_setup = referrer_map)] +fn referrer_decode(map: &'static HeaderMap) -> usize { + let value = ReferrerPolicy::view(black_box(map)).unwrap().unwrap(); + black_box(value.policies().count()) +} + +#[metabench::benchmark(REFERRER_PREFERRED_8, "policy", "referrer_preferred_8", gungraun_setup = referrer_map)] +fn referrer_preferred_8(map: &'static HeaderMap) -> usize { + let value = ReferrerPolicy::view(black_box(map)).unwrap().unwrap(); + let mut result = 0; + for _ in 0..8 { + result ^= value.preferred().unwrap() as usize; + } + black_box(result) +} + +#[metabench::benchmark(CACHE_DIRECTIVES, "policy", "cache_directives", gungraun_setup = cache_map)] +fn cache_directives(map: &'static HeaderMap) -> usize { + let value = CacheControl::view(black_box(map)).unwrap().unwrap(); + black_box(value.directives().map(|directive| directive.as_bytes().len()).sum()) +} + +#[metabench::benchmark(CACHE_MAX_AGE_8, "policy", "cache_max_age_8", gungraun_setup = cache_map)] +fn cache_max_age_8(map: &'static HeaderMap) -> u64 { + let value = CacheControl::view(black_box(map)).unwrap().unwrap(); + let mut result = 0; + for _ in 0..8 { + result ^= value.max_age().unwrap().as_secs(); + } + black_box(result) +} + +#[metabench::benchmark(HSTS_DECODE, "policy", "hsts_decode", gungraun_setup = hsts_map)] +fn hsts_decode(map: &'static HeaderMap) -> u64 { + let value = StrictTransportSecurity::view(black_box(map)).unwrap().unwrap(); + black_box(value.max_age().as_secs()) +} + +#[metabench::benchmark(HSTS_DIRECTIVES, "policy", "hsts_directives", gungraun_setup = hsts_map)] +fn hsts_directives(map: &'static HeaderMap) -> usize { + let value = StrictTransportSecurity::view(black_box(map)).unwrap().unwrap(); + black_box(value.directives().map(|directive| directive.unwrap().as_bytes().len()).sum()) +} + +#[metabench::benchmark(AUTH_ACCESS_8, "policy", "auth_access_8", gungraun_setup = basic, gungraun_teardown = drop_it)] +fn auth_access_8(value: AuthorizationOwned) -> (usize, AuthorizationOwned) { + let mut result = 0; + for _ in 0..8 { + result ^= value.encoded_credentials().unwrap().len(); + } + (black_box(result), value) +} + +#[metabench::benchmark(AUTH_DECODE, "policy", "auth_decode", gungraun_setup = authorization_map)] +fn auth_decode(map: &'static HeaderMap) -> usize { + let value = Authorization::::owned(black_box(map)).unwrap().unwrap(); + black_box(value.encoded_credentials().unwrap().len()) +} + +#[metabench::benchmark( + AUTH_EXTRACT_WARM, + "policy", + "auth_extract_warm", + gungraun_setup = basic_and_credentials, + gungraun_teardown = drop_it, +)] +fn auth_extract_warm(state: (AuthorizationOwned, BasicCredentials)) -> (usize, (AuthorizationOwned, BasicCredentials)) { + let (value, mut credentials) = state; + value.extract(&mut credentials).unwrap(); + let result = value.extract(&mut credentials).unwrap().username().len(); + (black_box(result), (value, credentials)) +} + +#[metabench::benchmark(CSP_FROM_BYTES, "policy", "csp_from_bytes")] +fn csp_from_bytes() -> usize { + let value = + http_headers::headers::ContentSecurityPolicyOwned::from_bytes(black_box(b"default-src 'self'; script-src 'nonce-abcdefghijklmno'")) + .unwrap(); + black_box(value.policies().next().unwrap().len()) +} + +#[metabench::benchmark(CSP_OWNED_16, "policy", "csp_owned_16", gungraun_setup = repeated_csp_map)] +fn csp_owned_16(map: &'static HeaderMap) -> usize { + let value = ContentSecurityPolicy::owned(black_box(map)).unwrap().unwrap(); + black_box(value.policies().count()) +} + +#[metabench::benchmark(SET_COOKIE_OWNED_16, "policy", "set_cookie_owned_16", gungraun_setup = repeated_cookie_map)] +fn set_cookie_owned_16(map: &'static HeaderMap) -> usize { + let value = SetCookie::owned(black_box(map)).unwrap().unwrap(); + black_box(value.len()) +} + +#[metabench::benchmark(LOCATION_RELAXED, "policy", "location_relaxed", gungraun_setup = location_map)] +fn location_relaxed(map: &'static HeaderMap) -> usize { + let value = Location::view_with(black_box(map), DecodeMode::Relaxed).unwrap().unwrap(); + black_box(value.as_bytes().len()) +} + +fn criterion_benchmarks(criterion: &mut Criterion) { + let mut group = criterion.benchmark_group("http_headers_policy/policy"); + macro_rules! with_setup { + ($id:ident, $setup:ident, $function:ident) => { + group.bench_function($id.benchmark_name(), |bencher| { + bencher.iter_batched($setup, $function, BatchSize::SmallInput); + }); + }; + } + with_setup!(REFERRER_DECODE, referrer_map, referrer_decode); + with_setup!(REFERRER_PREFERRED_8, referrer_map, referrer_preferred_8); + with_setup!(CACHE_DIRECTIVES, cache_map, cache_directives); + with_setup!(CACHE_MAX_AGE_8, cache_map, cache_max_age_8); + with_setup!(HSTS_DECODE, hsts_map, hsts_decode); + with_setup!(HSTS_DIRECTIVES, hsts_map, hsts_directives); + with_setup!(AUTH_ACCESS_8, basic, auth_access_8); + with_setup!(AUTH_DECODE, authorization_map, auth_decode); + with_setup!(AUTH_EXTRACT_WARM, basic_and_credentials, auth_extract_warm); + group.bench_function(CSP_FROM_BYTES.benchmark_name(), |b| b.iter(csp_from_bytes)); + with_setup!(CSP_OWNED_16, repeated_csp_map, csp_owned_16); + with_setup!(SET_COOKIE_OWNED_16, repeated_cookie_map, set_cookie_owned_16); + with_setup!(LOCATION_RELAXED, location_map, location_relaxed); + group.finish(); +} + +metabench::main!( + criterion = criterion_benchmarks, + benchmarks = [ + REFERRER_DECODE, + REFERRER_PREFERRED_8, + CACHE_DIRECTIVES, + CACHE_MAX_AGE_8, + HSTS_DECODE, + HSTS_DIRECTIVES, + AUTH_ACCESS_8, + AUTH_DECODE, + AUTH_EXTRACT_WARM, + CSP_FROM_BYTES, + CSP_OWNED_16, + SET_COOKIE_OWNED_16, + LOCATION_RELAXED, + ], +); diff --git a/crates/http_headers/benches/http_headers_policy_shapes.rs b/crates/http_headers/benches/http_headers_policy_shapes.rs new file mode 100644 index 000000000..607236368 --- /dev/null +++ b/crates/http_headers/benches/http_headers_policy_shapes.rs @@ -0,0 +1,127 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Policy, opaque-field and URI-reference parser shapes. + +use http_headers::DecodeErrorKind; +use http_headers::headers::{ + CacheControl, ContentSecurityPolicy, Location, ReferrerPolicy, SetCookie, StrictTransportSecurity, UserAgent, XContentTypeOptions, +}; + +#[path = "http_headers_shapes_common.rs"] +mod shapes; + +use shapes::Expected; + +shapes::define_shapes!( + "http_headers_policy_shapes/parse"; + (cache_control_canonical, CacheControl, &["max-age=3600, private"], Strict, Expected::Valid), + (cache_control_immutable, CacheControl, &["public, max-age=31536000, immutable"], Strict, Expected::Valid), + (cache_control_max_age_only, CacheControl, &["max-age=3600"], Strict, Expected::Valid), + (cache_control_overflow_token, CacheControl, &["max-age=18446744073709551616"], Strict, Expected::Valid), + (cache_control_first_valid_seconds, CacheControl, &["max-age=invalid, max-age=\"45\", max-age=90"], Strict, Expected::Valid), + (cache_control_quoted_extensions, CacheControl, &["no-cache, x-mode=\"fast, safe\", x-note=\"a\\\"b\", max-age=60"], Strict, Expected::Valid), + (cache_control_mixed_case_ows, CacheControl, &[" \tPUBLIC,\tMAX-AGE=60, NO-CACHE \t"], Strict, Expected::Valid), + (cache_control_repeated, CacheControl, &["max-age=3600", "public, must-revalidate", "no-transform, s-maxage=120"], Strict, Expected::Valid), + (cache_control_many_extensions, CacheControl, &["max-age=3600, ext0=value0, ext1=value1, ext2=value2, ext3=value3, ext4=value4, ext5=value5, ext6=value6, ext7=value7, ext8=value8, ext9=value9, ext10=value10, ext11=value11, ext12=value12, ext13=value13, ext14=value14, ext15=value15, ext16=value16, ext17=value17, ext18=value18, ext19=value19, ext20=value20, ext21=value21, ext22=value22, ext23=value23, ext24=value24, ext25=value25, ext26=value26, ext27=value27, ext28=value28, ext29=value29, ext30=value30, ext31=value31, public"], Strict, Expected::Valid), + (cache_control_empty, CacheControl, &[""], Strict, Expected::Valid), + (cache_control_empty_members, CacheControl, &[",, no-cache,", ",,,"], Strict, Expected::Valid), + (cache_control_bad_name, CacheControl, &["max-age =30"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (cache_control_bad_value, CacheControl, &["max-age= 30"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (cache_control_unterminated_quote, CacheControl, &["public", "x-note=\"unterminated"], Strict, Expected::Error(DecodeErrorKind::UnterminatedQuote)), + (cache_control_absent, CacheControl, &[], Strict, Expected::Absent), + + (content_security_policy_canonical, ContentSecurityPolicy, &["default-src 'self'"], Strict, Expected::Valid), + (content_security_policy_nonce, ContentSecurityPolicy, &["default-src 'self'; script-src 'nonce-abcdefghijklmno'; object-src 'none'; base-uri 'self'"], Strict, Expected::Valid), + (content_security_policy_large, ContentSecurityPolicy, &["default-src 'self'; script-src 'self' 'unsafe-inline' https://cdn.example.com https://analytics.example.com; style-src 'self' 'unsafe-inline' https://fonts.googleapis.com; img-src 'self' data: https:; font-src 'self' https://fonts.gstatic.com; connect-src 'self' https://api.example.com; frame-ancestors 'none'; base-uri 'self'; form-action 'self'"], Strict, Expected::Valid), + (content_security_policy_repeated_two, ContentSecurityPolicy, &["default-src 'self'", "frame-ancestors 'none'"], Strict, Expected::Valid), + (content_security_policy_repeated_sixteen, ContentSecurityPolicy, &["default-src 'self'", "script-src 'none'", "img-src https:", "style-src 'self'", "font-src https:", "connect-src 'self'", "frame-src 'none'", "object-src 'none'", "base-uri 'self'", "form-action 'self'", "frame-ancestors 'none'", "upgrade-insecure-requests", "block-all-mixed-content", "worker-src 'self'", "manifest-src 'self'", "media-src 'none'"], Strict, Expected::Valid), + (content_security_policy_empty, ContentSecurityPolicy, &[""], Strict, Expected::Valid), + (content_security_policy_opaque_fallback, ContentSecurityPolicy, &["\tunknown-directive value, other; report-to=\"unterminated"], Strict, Expected::Valid), + (content_security_policy_absent, ContentSecurityPolicy, &[], Strict, Expected::Absent), + + (referrer_policy_canonical, ReferrerPolicy, &["strict-origin-when-cross-origin"], Strict, Expected::Valid), + (referrer_policy_no_referrer, ReferrerPolicy, &["no-referrer"], Strict, Expected::Valid), + (referrer_policy_unknown_only, ReferrerPolicy, &["future-policy"], Strict, Expected::Valid), + (referrer_policy_ows, ReferrerPolicy, &[" \tstrict-origin-when-cross-origin \t"], Strict, Expected::Valid), + (referrer_policy_repeated_recognized, ReferrerPolicy, &["no-referrer", "strict-origin"], Strict, Expected::Valid), + (referrer_policy_repeated_mixed, ReferrerPolicy, &["future-policy, no-referrer", "origin, strict-origin-when-cross-origin"], Strict, Expected::Valid), + (referrer_policy_repeated_eight, ReferrerPolicy, &["no-referrer", "no-referrer-when-downgrade", "origin", "origin-when-cross-origin", "same-origin", "strict-origin", "strict-origin-when-cross-origin", "unsafe-url"], Strict, Expected::Valid), + (referrer_policy_many_extensions, ReferrerPolicy, &["future-origin-policy, vendor-origin-policy, private-origin-policy, site-origin-policy, strict-site-policy, private-network-policy, source-origin-policy, target-origin-policy, no-referrer, origin, same-origin, strict-origin, future-policy, strict-origin-when-cross-origin"], Strict, Expected::Valid), + (referrer_policy_empty_members, ReferrerPolicy, &[" , no-referrer,\t, origin , "], Strict, Expected::Valid), + (referrer_policy_empty, ReferrerPolicy, &[""], Strict, Expected::Error(DecodeErrorKind::MissingValue)), + (referrer_policy_bad_token, ReferrerPolicy, &["not a token"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (referrer_policy_unterminated_quote, ReferrerPolicy, &["origin, \"strict-origin"], Strict, Expected::Error(DecodeErrorKind::UnterminatedQuote)), + (referrer_policy_absent, ReferrerPolicy, &[], Strict, Expected::Absent), + + (strict_transport_security_canonical, StrictTransportSecurity, &["max-age=31536000; includeSubDomains"], Strict, Expected::Valid), + (strict_transport_security_preload, StrictTransportSecurity, &["max-age=63072000; includeSubDomains; preload"], Strict, Expected::Valid), + (strict_transport_security_max_age_only, StrictTransportSecurity, &["max-age=60"], Strict, Expected::Valid), + (strict_transport_security_maximum_seconds, StrictTransportSecurity, &["max-age=18446744073709551615"], Strict, Expected::Valid), + (strict_transport_security_mixed_case, StrictTransportSecurity, &["MAX-AGE=31536000; includeSubDomains; preload"], Strict, Expected::Valid), + (strict_transport_security_reordered, StrictTransportSecurity, &["preload; includeSubDomains; max-age=31536000"], Strict, Expected::Valid), + (strict_transport_security_quoted_seconds, StrictTransportSecurity, &["max-age=\"31536000\"; includeSubDomains"], Strict, Expected::Valid), + (strict_transport_security_escaped_seconds, StrictTransportSecurity, &["MAX-AGE=\"6\\0\"; includeSubDomains; preload; future=\"a,b\""], Strict, Expected::Valid), + (strict_transport_security_quoted_extension, StrictTransportSecurity, &["MAX-AGE=60 ; includeSubDomains ; x-vendor=\"a;b\"; report-to=\"a\\\";b\""], Strict, Expected::Valid), + (strict_transport_security_many_extensions, StrictTransportSecurity, &["max-age=31536000; includeSubDomains; preload; rollout=stable; report-to=audit; policy-version=20260918; x-region=global; x-service=frontend; x-owner=platform; x-mode=enforce; x-audit=enabled; x-source=edge; x-route=public; x-build=123456789012345678901"], Strict, Expected::Valid), + (strict_transport_security_repeated_field, StrictTransportSecurity, &["max-age=60", "max-age=120"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + (strict_transport_security_empty, StrictTransportSecurity, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (strict_transport_security_overflow, StrictTransportSecurity, &["max-age=18446744073709551616"], Strict, Expected::Error(DecodeErrorKind::InvalidNumber)), + (strict_transport_security_duplicate_subdomains, StrictTransportSecurity, &["max-age=60; includeSubDomains; includeSubDomains"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (strict_transport_security_absent, StrictTransportSecurity, &[], Strict, Expected::Absent), + + (x_content_type_options_canonical, XContentTypeOptions, &["nosniff"], Strict, Expected::Valid), + (x_content_type_options_case, XContentTypeOptions, &["NoSniff"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (x_content_type_options_ows, XContentTypeOptions, &["nosniff "], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (x_content_type_options_empty, XContentTypeOptions, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (x_content_type_options_repeated, XContentTypeOptions, &["nosniff", "nosniff"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + (x_content_type_options_absent, XContentTypeOptions, &[], Strict, Expected::Absent), + + (location_relative, Location, &["/next"], Strict, Expected::Valid), + (location_absolute, Location, &["https://example.com/en-us/docs/reference/index.html?utm_source=newsletter&utm_campaign=spring&page=3"], Strict, Expected::Valid), + (location_absolute_port, Location, &["https://example.com:8443/docs"], Strict, Expected::Valid), + (location_parent_relative, Location, &["../people?tab=1#profile"], Strict, Expected::Valid), + (location_query_only, Location, &["?page=2"], Strict, Expected::Valid), + (location_fragment_only, Location, &["#profile"], Strict, Expected::Valid), + (location_empty, Location, &[""], Strict, Expected::Valid), + (location_network_relative, Location, &["//cdn.example.com/assets/main.css"], Strict, Expected::Valid), + (location_userinfo, Location, &["https://user:pass@example.com/docs"], Strict, Expected::Valid), + (location_ipv6, Location, &["https://[2001:db8::1]:8443/docs?x=1#top"], Strict, Expected::Valid), + (location_ipvfuture, Location, &["https://[v1.address]/"], Strict, Expected::Valid), + (location_mailto, Location, &["mailto:user@example.com"], Strict, Expected::Valid), + (location_urn, Location, &["urn:example:animal:ferret:nose"], Strict, Expected::Valid), + (location_file_empty_authority, Location, &["file:///var/docs/index.html"], Strict, Expected::Valid), + (location_percent_encoded_long, Location, &["/download/reports/2026%2F09%2F18/quarterly%20summary.csv?filename=quarterly%20summary.csv&response-content-disposition=attachment%3B%20filename%3Dsummary.csv&redirect=https%3A%2F%2Fexample.com%2Faccount%3Ftab%3Dreports#download"], Strict, Expected::Valid), + (location_invalid_percent, Location, &["/bad%2"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (location_invalid_scheme, Location, &["1abc:def"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (location_strict_backslash, Location, &["/a\\b\\c"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (location_relaxed_backslash, Location, &["/a\\b\\c"], Relaxed, Expected::Valid), + (location_relaxed_common, Location, &["https://example.com/docs?x=1#top"], Relaxed, Expected::Valid), + (location_relaxed_percent_encoded, Location, &["https://example.com/a%20path?query=1#fragment"], Relaxed, Expected::Valid), + (location_relaxed_invalid_neighbor, Location, &["bad\\%zz"], Relaxed, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (location_relaxed_invalid_without_backslash, Location, &["%zz"], Relaxed, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (location_repeated, Location, &["/first", "/second"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + (location_absent, Location, &[], Strict, Expected::Absent), + + (set_cookie_canonical, SetCookie, &["session=benchmark-session-value; Path=/; HttpOnly; Secure; SameSite=Lax"], Strict, Expected::Valid), + (set_cookie_small, SetCookie, &["a=1"], Strict, Expected::Valid), + (set_cookie_expires, SetCookie, &["theme=dark; Expires=Wed, 21 Oct 2037 07:28:00 GMT; Path=/; SameSite=Lax"], Strict, Expected::Valid), + (set_cookie_repeated_four, SetCookie, &["session=benchmark-session-value; Path=/; HttpOnly; Secure; SameSite=Lax", "csrf=benchmark-csrf-value; Path=/; Secure; SameSite=Strict", "theme=dark; Path=/; Max-Age=31536000", "locale=en-US; Path=/; Max-Age=31536000"], Strict, Expected::Valid), + (set_cookie_repeated_sixteen, SetCookie, &["a=1", "b=2", "c=3", "d=4", "e=5", "f=6", "g=7", "h=8", "i=9", "j=10", "k=11", "l=12", "m=13", "n=14", "o=15", "p=16"], Strict, Expected::Valid), + (set_cookie_long, SetCookie, &["preferences=locale%3Den-US%26theme%3Ddark%26timezone%3DAmerica%2FLos_Angeles%26layout%3Dcompact%26notifications%3Denabled%26dashboard%3Doverview%26accessibility%3Dhigh-contrast%26table-columns%3Dname%2Cowner%2Cstatus%2Cupdated%26navigation%3Dexpanded; Path=/; Domain=example.com; Max-Age=31536000; Secure; SameSite=Lax"], Strict, Expected::Valid), + (set_cookie_opaque, SetCookie, &["opaque; not-a-cookie-pair; future=\"unterminated, value"], Strict, Expected::Valid), + (set_cookie_empty, SetCookie, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (set_cookie_late_empty, SetCookie, &["a=1; Path=/", "b=2; HttpOnly", ""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (set_cookie_absent, SetCookie, &[], Strict, Expected::Absent), + + (user_agent_canonical, UserAgent, &["curl/8.5.0"], Strict, Expected::Valid), + (user_agent_browser, UserAgent, &["Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/120.0.0.0 Safari/537.36"], Strict, Expected::Valid), + (user_agent_client_comment, UserAgent, &["example-client/1.0 (integration test)"], Strict, Expected::Valid), + (user_agent_long_components, UserAgent, &["ExampleApplication/4.2 Runtime/9.0 (Linux; x86_64; production) HttpClient/3.1 Telemetry/2.4 Identity/5.0 RetryPolicy/1.2 TraceContext/1.0 ProxySupport/2.1 Component/1.0 (build 20260918; feature-set-extended) Deployment/2026.09 Region/global Compatibility/legacy"], Strict, Expected::Valid), + (user_agent_padded, UserAgent, &[" client/1"], Strict, Expected::Valid), + (user_agent_opaque, UserAgent, &["not-a-product (unclosed; \"opaque, value"], Strict, Expected::Valid), + (user_agent_empty, UserAgent, &[""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (user_agent_blank, UserAgent, &[" \t \t "], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (user_agent_repeated, UserAgent, &["client/1", "proxy/2"], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + (user_agent_absent, UserAgent, &[], Strict, Expected::Absent), +); diff --git a/crates/http_headers/benches/http_headers_protocol.rs b/crates/http_headers/benches/http_headers_protocol.rs new file mode 100644 index 000000000..fa11fa452 --- /dev/null +++ b/crates/http_headers/benches/http_headers_protocol.rs @@ -0,0 +1,260 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Focused instruction-count and wall-clock cases for protocol header parsing. + +use std::hint::black_box; +use std::sync::OnceLock; + +use criterion::{BatchSize, BenchmarkId, Criterion}; +use http::{HeaderMap, HeaderName, HeaderValue}; +use http_headers::headers::{ + AcceptRanges, AccessControlAllowHeadersOwned, AccessControlAllowOriginOwned, ContentLength, ETagOwned, Host, IfNoneMatchOwned, + RangeOwned, SecWebSocketExtensionsOwned, SecWebSocketProtocolOwned, SecWebSocketVersion, +}; +use http_headers::{DecodeErrorKind, DecodeMode, FieldValue, SingleValueField}; + +const GROUP: &str = "http_headers_protocol/protocol"; + +fn map(name: &'static str, values: &'static [&'static str]) -> HeaderMap { + let mut map = HeaderMap::with_capacity(values.len()); + let name = HeaderName::from_static(name); + for value in values { + map.append(&name, HeaderValue::from_static(value)); + } + map +} + +fn accept_ranges_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map("accept-ranges", &["bytes", "items", "records"])) +} + +fn content_length_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map("content-length", &["18446744073709551615"])) +} + +fn cors_many_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| { + map( + "access-control-allow-headers", + &[ + "x-00", "x-01", "x-02", "x-03", "x-04", "x-05", "x-06", "x-07", "x-08", "x-09", "x-10", "x-11", "x-12", "x-13", "x-14", + "x-15", + ], + ) + }) +} + +fn websocket_version_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map("sec-websocket-version", &["13", "8", "\"7\""])) +} + +fn host_ipv_future_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| map("host", &["[v1.fe80::a]:443"])) +} + +fn host_domain() -> FieldValue { + FieldValue::from_static("www.example.com:443") +} + +fn host_idna() -> FieldValue { + FieldValue::from_static("münich.example:443") +} + +fn if_none_match() -> IfNoneMatchOwned { + IfNoneMatchOwned::try_from(FieldValue::from_static( + r#""short", W/"0123456789abcdefghijklmnopqrstuvwxyz", "obs\x80text""#, + )) + .expect("valid conditional tags") +} + +fn byte_range() -> RangeOwned { + RangeOwned::try_from("bytes=0-49, 100-, -500").expect("valid byte range") +} + +fn extension_range() -> RangeOwned { + RangeOwned::extension("items", "1-5").expect("valid extension range") +} + +fn websocket_protocol() -> SecWebSocketProtocolOwned { + SecWebSocketProtocolOwned::try_from("graphql-transport-ws, graphql-ws").expect("valid protocols") +} + +fn websocket_extensions() -> SecWebSocketExtensionsOwned { + SecWebSocketExtensionsOwned::try_from(r#"permessage-deflate; client_max_window_bits; mode="fast", x-test; token=value"#) + .expect("valid extensions") +} + +fn etag() -> ETagOwned { + ETagOwned::weak("0123456789abcdefghijklmnopqrstuvwxyz").expect("valid entity tag") +} + +fn sixteen_tags() -> impl Iterator { + (0..16).map(|index| ETagOwned::strong(format!("revision-{index}")).expect("valid entity tag")) +} + +#[metabench::benchmark(ACCEPT_RANGES_OWNED_MULTILINE, GROUP, "accept_ranges_owned_multiline", gungraun_setup = accept_ranges_map)] +fn accept_ranges_owned_multiline(map: &'static HeaderMap) -> usize { + AcceptRanges::owned(black_box(map)) + .expect("valid header") + .expect("present header") + .units() + .count() +} + +#[metabench::benchmark(CONTENT_LENGTH_20_DIGITS, GROUP, "content_length_20_digits", gungraun_setup = content_length_map)] +fn content_length_20_digits(map: &'static HeaderMap) -> u64 { + ContentLength::view(black_box(map)) + .expect("valid header") + .expect("present header") + .get() +} + +#[metabench::benchmark(CONDITIONAL_TAGS_LONG, GROUP, "conditional_tags_long", gungraun_setup = if_none_match)] +fn conditional_tags_long(value: IfNoneMatchOwned) -> (usize, IfNoneMatchOwned) { + let length = value.tags().map(|tag| tag.as_bytes().len()).sum(); + (black_box(length), value) +} + +#[metabench::benchmark(RANGE_BYTE_MEMBERS, GROUP, "range_byte_members", gungraun_setup = byte_range)] +fn range_byte_members(value: RangeOwned) -> (usize, RangeOwned) { + let count = value.byte_ranges().expect("byte unit").count(); + (black_box(count), value) +} + +#[metabench::benchmark(RANGE_EXTENSION_PROJECTION, GROUP, "range_extension_projection", gungraun_setup = extension_range)] +fn range_extension_projection(value: RangeOwned) -> (usize, RangeOwned) { + let length = value.extension_range_set().expect("extension unit").len(); + (black_box(length), value) +} + +#[metabench::benchmark(CORS_SINGLETON_CONVERSION, GROUP, "cors_singleton_conversion")] +fn cors_singleton_conversion() -> AccessControlAllowHeadersOwned { + AccessControlAllowHeadersOwned::try_from(black_box(FieldValue::from_static("x-custom-header"))).expect("valid CORS list") +} + +#[metabench::benchmark(CORS_OWNED_16_LINES, GROUP, "cors_owned_16_lines", gungraun_setup = cors_many_map)] +fn cors_owned_16_lines(map: &'static HeaderMap) -> usize { + AccessControlAllowHeadersOwned::from_field_values(map.get_all("access-control-allow-headers").iter().map(FieldValue::from).collect()) + .expect("valid CORS list") + .len() +} + +#[metabench::benchmark(CONDITIONAL_OWNED_16_TAGS, GROUP, "conditional_owned_16_tags")] +fn conditional_owned_16_tags() -> usize { + IfNoneMatchOwned::from_tags(sixteen_tags()).expect("valid tags").tags().count() +} + +#[metabench::benchmark(CORS_IPV6_ORIGIN, GROUP, "cors_ipv6_origin")] +fn cors_ipv6_origin() -> AccessControlAllowOriginOwned { + AccessControlAllowOriginOwned::try_from(black_box(FieldValue::from_static("https://[2001:db8::1]:8443"))) + .expect("canonical IPv6 origin") +} + +#[metabench::benchmark(WEBSOCKET_PROTOCOL_READ, GROUP, "websocket_protocol_read", gungraun_setup = websocket_protocol)] +fn websocket_protocol_read(value: SecWebSocketProtocolOwned) -> (usize, SecWebSocketProtocolOwned) { + let length = value.protocols().map(|item| item.expect("validated protocol").len()).sum(); + (black_box(length), value) +} + +#[metabench::benchmark(WEBSOCKET_EXTENSION_READ, GROUP, "websocket_extension_read", gungraun_setup = websocket_extensions)] +fn websocket_extension_read(value: SecWebSocketExtensionsOwned) -> (usize, SecWebSocketExtensionsOwned) { + let count = value + .extensions() + .map(|extension| extension.expect("validated extension").parameters().count()) + .sum(); + (black_box(count), value) +} + +#[metabench::benchmark(WEBSOCKET_VERSION_LATE_QUOTE, GROUP, "websocket_version_late_quote", gungraun_setup = websocket_version_map)] +fn websocket_version_late_quote(map: &'static HeaderMap) -> DecodeErrorKind { + SecWebSocketVersion::view(black_box(map)) + .expect_err("quoted version is invalid") + .kind() +} + +#[metabench::benchmark(HOST_IPV_FUTURE, GROUP, "host_ipv_future", gungraun_setup = host_ipv_future_map)] +fn host_ipv_future(map: &'static HeaderMap) -> usize { + Host::view(black_box(map)).expect("valid host").expect("present host").host().len() +} + +#[metabench::benchmark(HOST_RELAXED_DOMAIN, GROUP, "host_relaxed_domain", gungraun_setup = host_domain)] +fn host_relaxed_domain(value: FieldValue) -> (usize, FieldValue) { + let view = ::decode_view_with(value.as_field_value_ref(), DecodeMode::Relaxed).expect("valid relaxed host"); + let length = black_box(view.host().len() + view.port().map_or(0, str::len)); + drop(view); + (length, value) +} + +#[metabench::benchmark(HOST_RELAXED_IDNA, GROUP, "host_relaxed_idna", gungraun_setup = host_idna)] +fn host_relaxed_idna(value: FieldValue) -> (usize, FieldValue) { + let view = ::decode_view_with(value.as_field_value_ref(), DecodeMode::Relaxed).expect("valid relaxed host"); + let length = black_box(view.host().len() + view.port().map_or(0, str::len)); + drop(view); + (length, value) +} + +#[metabench::benchmark(ETAG_OPAQUE_READ, GROUP, "etag_opaque_read", gungraun_setup = etag)] +fn etag_opaque_read(value: ETagOwned) -> (usize, ETagOwned) { + let length = value.opaque_tag().expect("valid stored tag").len(); + (black_box(length), value) +} + +fn criterion_benchmarks(criterion: &mut Criterion) { + let mut group = criterion.benchmark_group(GROUP); + macro_rules! borrowed { + ($id:ident, $setup:ident, $function:ident) => { + group.bench_function(BenchmarkId::new($id.benchmark_name(), stringify!($function)), |bencher| { + bencher.iter_batched($setup, $function, BatchSize::SmallInput); + }); + }; + } + borrowed!(ACCEPT_RANGES_OWNED_MULTILINE, accept_ranges_map, accept_ranges_owned_multiline); + borrowed!(CONTENT_LENGTH_20_DIGITS, content_length_map, content_length_20_digits); + borrowed!(CONDITIONAL_TAGS_LONG, if_none_match, conditional_tags_long); + borrowed!(RANGE_BYTE_MEMBERS, byte_range, range_byte_members); + borrowed!(RANGE_EXTENSION_PROJECTION, extension_range, range_extension_projection); + group.bench_function(CORS_SINGLETON_CONVERSION.benchmark_name(), |bencher| { + bencher.iter(cors_singleton_conversion); + }); + borrowed!(CORS_OWNED_16_LINES, cors_many_map, cors_owned_16_lines); + group.bench_function(CONDITIONAL_OWNED_16_TAGS.benchmark_name(), |bencher| { + bencher.iter(conditional_owned_16_tags); + }); + group.bench_function(CORS_IPV6_ORIGIN.benchmark_name(), |bencher| bencher.iter(cors_ipv6_origin)); + borrowed!(WEBSOCKET_PROTOCOL_READ, websocket_protocol, websocket_protocol_read); + borrowed!(WEBSOCKET_EXTENSION_READ, websocket_extensions, websocket_extension_read); + borrowed!(WEBSOCKET_VERSION_LATE_QUOTE, websocket_version_map, websocket_version_late_quote); + borrowed!(HOST_IPV_FUTURE, host_ipv_future_map, host_ipv_future); + borrowed!(HOST_RELAXED_DOMAIN, host_domain, host_relaxed_domain); + borrowed!(HOST_RELAXED_IDNA, host_idna, host_relaxed_idna); + borrowed!(ETAG_OPAQUE_READ, etag, etag_opaque_read); + group.finish(); +} + +metabench::main!( + criterion = criterion_benchmarks, + benchmarks = [ + ACCEPT_RANGES_OWNED_MULTILINE, + CONTENT_LENGTH_20_DIGITS, + CONDITIONAL_TAGS_LONG, + RANGE_BYTE_MEMBERS, + RANGE_EXTENSION_PROJECTION, + CORS_SINGLETON_CONVERSION, + CORS_OWNED_16_LINES, + CONDITIONAL_OWNED_16_TAGS, + CORS_IPV6_ORIGIN, + WEBSOCKET_PROTOCOL_READ, + WEBSOCKET_EXTENSION_READ, + WEBSOCKET_VERSION_LATE_QUOTE, + HOST_IPV_FUTURE, + HOST_RELAXED_DOMAIN, + HOST_RELAXED_IDNA, + ETAG_OPAQUE_READ, + ] +); diff --git a/crates/http_headers/benches/http_headers_shapes_common.rs b/crates/http_headers/benches/http_headers_shapes_common.rs new file mode 100644 index 000000000..410168ca8 --- /dev/null +++ b/crates/http_headers/benches/http_headers_shapes_common.rs @@ -0,0 +1,268 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Shared HTTP and custom-source decode operations for parser-shape benchmarks. + +use std::hint::black_box; + +use http::{HeaderMap, HeaderName, HeaderValue}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldName, FieldValueRef}; + +#[derive(Clone, Copy, Debug)] +pub(crate) enum Expected { + Valid, + Absent, + Error(DecodeErrorKind), +} + +#[derive(Clone, Copy)] +pub(crate) struct Expectations { + pub(crate) http_owned: Expected, + pub(crate) http_borrowed: Expected, + pub(crate) raw_owned: Expected, + pub(crate) raw_borrowed: Expected, +} + +impl From for Expectations { + fn from(expected: Expected) -> Self { + Self { + http_owned: expected, + http_borrowed: expected, + raw_owned: expected, + raw_borrowed: expected, + } + } +} + +impl Expectations { + fn for_operation(self, raw: bool, owned: bool) -> Expected { + match (raw, owned) { + (false, false) => self.http_borrowed, + (false, true) => self.http_owned, + (true, false) => self.raw_borrowed, + (true, true) => self.raw_owned, + } + } +} + +pub(crate) struct Fixture { + http: HeaderMap, + raw: RawSource, +} + +struct RawSource { + name: &'static FieldName, + values: Vec>, +} + +impl FieldSource for RawSource { + #[inline] + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::from_borrowed(name, &self.values)).flatten() + } +} + +impl Fixture { + pub(crate) fn new(values: &'static [&'static str], mode: DecodeMode, expected: impl Into) -> Self { + let expected = expected.into(); + let name = F::name(); + let http_name = HeaderName::from_static(name.as_str()); + let mut http = HeaderMap::with_capacity(1); + for value in values { + let mut value = HeaderValue::from_bytes(value.as_bytes()).expect("fixture is a valid HTTP field value"); + value.set_sensitive(matches!(name.as_str(), "authorization" | "set-cookie")); + http.append(&http_name, value); + } + let fixture = Self { + http, + raw: RawSource { + name, + values: values.iter().map(|value| FieldValueRef::new(value.as_bytes())).collect(), + }, + }; + check(F::view_with(&fixture.http, mode), expected.http_borrowed, "HTTP borrowed"); + check(F::owned_with(&fixture.http, mode), expected.http_owned, "HTTP owned"); + check(F::view_with(&fixture.raw, mode), expected.raw_borrowed, "raw borrowed"); + check(F::owned_with(&fixture.raw, mode), expected.raw_owned, "raw owned"); + fixture + } +} + +#[expect(clippy::panic, reason = "invalid benchmark fixtures must fail setup")] +fn check(result: Result, DecodeError>, expected: Expected, operation: &str) { + match (result, expected) { + (Ok(Some(_)), Expected::Valid) | (Ok(None), Expected::Absent) => {} + (Err(error), Expected::Error(kind)) => assert_eq!(error.kind(), kind, "{operation}: {error:?}"), + (Err(error), _) => panic!("{operation}: fixture unexpectedly failed: {error:?}"), + (Ok(_), _) => panic!("{operation}: fixture did not produce {expected:?}"), + } +} + +#[derive(Clone, Copy)] +pub(crate) struct Case { + fixture: &'static Fixture, + operation: fn(&Fixture), +} + +impl Case { + #[inline] + pub(crate) fn run(self) { + (self.operation)(black_box(self.fixture)); + } +} + +const fn mode() -> DecodeMode { + if RELAXED { DecodeMode::Relaxed } else { DecodeMode::Strict } +} + +#[expect(clippy::inline_always, reason = "result consumption belongs inside each measured decode operation")] +#[expect(clippy::panic, reason = "a benchmark fixture must retain its expected outcome")] +#[inline(always)] +fn consume(result: Result, DecodeError>) { + match OUTCOME { + 0 => drop(black_box(result.expect("fixture must decode").expect("fixture must be present"))), + 1 => assert!(black_box(result.expect("absent fixture must decode")).is_none()), + _ => match result { + Err(error) => { + let _error = black_box(error); + } + Ok(_) => panic!("fixture must fail"), + }, + } +} + +#[inline(never)] +fn http_borrowed(fixture: &Fixture) { + consume::<_, OUTCOME>(F::view_with(black_box(&fixture.http), mode::())); +} + +#[inline(never)] +fn http_owned(fixture: &Fixture) { + consume::<_, OUTCOME>(F::owned_with(black_box(&fixture.http), mode::())); +} + +#[inline(never)] +fn raw_borrowed(fixture: &Fixture) { + consume::<_, OUTCOME>(F::view_with(black_box(&fixture.raw), mode::())); +} + +#[inline(never)] +fn raw_owned(fixture: &Fixture) { + consume::<_, OUTCOME>(F::owned_with(black_box(&fixture.raw), mode::())); +} + +fn operation(raw: bool, owned: bool) -> fn(&Fixture) { + match (raw, owned) { + (false, false) => http_borrowed::, + (false, true) => http_owned::, + (true, false) => raw_borrowed::, + (true, true) => raw_owned::, + } +} + +pub(crate) fn case( + fixture: &'static Fixture, + mode: DecodeMode, + expected: impl Into, + raw: bool, + owned: bool, +) -> Case { + let operation = match (mode, expected.into().for_operation(raw, owned)) { + (DecodeMode::Strict, Expected::Valid) => operation::(raw, owned), + (DecodeMode::Strict, Expected::Absent) => operation::(raw, owned), + (DecodeMode::Strict, Expected::Error(_)) => operation::(raw, owned), + (DecodeMode::Relaxed, Expected::Valid) => operation::(raw, owned), + (DecodeMode::Relaxed, Expected::Absent) => operation::(raw, owned), + (DecodeMode::Relaxed, Expected::Error(_)) => operation::(raw, owned), + }; + Case { fixture, operation } +} + +macro_rules! define_shapes { + ($group:literal; $(($id:ident, $header:ty, $values:expr, $mode:ident, $expected:expr)),+ $(,)?) => { + paste::paste! { + $( + fn [<$id _fixture>]() -> &'static $crate::shapes::Fixture { + static FIXTURE: std::sync::OnceLock<$crate::shapes::Fixture> = std::sync::OnceLock::new(); + FIXTURE.get_or_init(|| $crate::shapes::Fixture::new::<$header>( + $values, + http_headers::DecodeMode::$mode, + $expected, + )) + } + + fn [<$id _http_owned>]() -> $crate::shapes::Case { + $crate::shapes::case::<$header>( + [<$id _fixture>](), http_headers::DecodeMode::$mode, $expected, false, true, + ) + } + fn [<$id _http_borrowed>]() -> $crate::shapes::Case { + $crate::shapes::case::<$header>( + [<$id _fixture>](), http_headers::DecodeMode::$mode, $expected, false, false, + ) + } + fn [<$id _raw_owned>]() -> $crate::shapes::Case { + $crate::shapes::case::<$header>( + [<$id _fixture>](), http_headers::DecodeMode::$mode, $expected, true, true, + ) + } + fn [<$id _raw_borrowed>]() -> $crate::shapes::Case { + $crate::shapes::case::<$header>( + [<$id _fixture>](), http_headers::DecodeMode::$mode, $expected, true, false, + ) + } + )+ + + #[metabench::benchmark(HTTP_OWNED, $group, "http_owned")] + $(#[bench::$id(setup = [<$id _http_owned>])])+ + fn http_owned(case: $crate::shapes::Case) { case.run(); } + + #[metabench::benchmark(HTTP_BORROWED, $group, "http_borrowed")] + $(#[bench::$id(setup = [<$id _http_borrowed>])])+ + fn http_borrowed(case: $crate::shapes::Case) { case.run(); } + + #[metabench::benchmark(RAW_OWNED, $group, "raw_owned")] + $(#[bench::$id(setup = [<$id _raw_owned>])])+ + fn raw_owned(case: $crate::shapes::Case) { case.run(); } + + #[metabench::benchmark(RAW_BORROWED, $group, "raw_borrowed")] + $(#[bench::$id(setup = [<$id _raw_borrowed>])])+ + fn raw_borrowed(case: $crate::shapes::Case) { case.run(); } + + fn criterion_benchmarks(criterion: &mut criterion::Criterion) { + let mut group = criterion.benchmark_group($group); + $( + group.bench_function( + criterion::BenchmarkId::new(HTTP_OWNED.benchmark_name(), stringify!($id)), + |b| b.iter_batched([<$id _http_owned>], http_owned, criterion::BatchSize::SmallInput), + ); + group.bench_function( + criterion::BenchmarkId::new(HTTP_BORROWED.benchmark_name(), stringify!($id)), + |b| b.iter_batched([<$id _http_borrowed>], http_borrowed, criterion::BatchSize::SmallInput), + ); + group.bench_function( + criterion::BenchmarkId::new(RAW_OWNED.benchmark_name(), stringify!($id)), + |b| b.iter_batched([<$id _raw_owned>], raw_owned, criterion::BatchSize::SmallInput), + ); + group.bench_function( + criterion::BenchmarkId::new(RAW_BORROWED.benchmark_name(), stringify!($id)), + |b| b.iter_batched([<$id _raw_borrowed>], raw_borrowed, criterion::BatchSize::SmallInput), + ); + )+ + group.finish(); + } + + metabench::main!( + criterion = { + factory = criterion::Criterion::default, + benchmarks = criterion_benchmarks, + unit = "ns", + }, + benchmarks = [HTTP_OWNED, HTTP_BORROWED, RAW_OWNED, RAW_BORROWED], + ); + } + }; +} + +pub(crate) use define_shapes; diff --git a/crates/http_headers/benches/http_headers_storage.rs b/crates/http_headers/benches/http_headers_storage.rs new file mode 100644 index 000000000..abf9645ee --- /dev/null +++ b/crates/http_headers/benches/http_headers_storage.rs @@ -0,0 +1,680 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Unified `http_headers` benchmarks for storage and custom-sink encoding decisions. + +use std::hint::black_box; +use std::iter; +use std::sync::{LazyLock, OnceLock}; + +use compact_str::CompactString; +use criterion::{BatchSize, BenchmarkId, Criterion}; +use http::HeaderMap; +use http_headers::sink::{EncodedValues, FieldSink, InsertError, ValueRefsEncoder}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{FieldName, FieldValue, FieldValueRef}; +use smallvec::SmallVec; + +use self::storage_operations::{AppendInput, MapOutput}; + +#[path = "../tests/common/http_headers_storage_operations.rs"] +mod storage_operations; + +const STORAGE_TEXT: &str = "http_headers_storage/storage_text"; +const TEXT: &str = "http_headers_storage/text"; +const MAP_WRITER: &str = "http_headers_storage/map_writer"; +const MAP_APPEND: &str = "http_headers_storage/map_append"; +const DEFERRED_INSERTION: &str = "http_headers_storage/deferred_insertion"; +const STORAGE_REPEATED_VALUES: &str = "http_headers_storage/storage_repeated_values"; +const REPEATED_VALUES: &str = "http_headers_storage/repeated_values"; +const STORAGE_INSERTION: &str = "http_headers_storage/storage_insertion"; +const INSERTION: &str = "http_headers_storage/insertion"; +const CUSTOM_SINK_ENCODING: &str = "http_headers_storage/custom_sink_encoding"; + +fn short_text() -> &'static str { + "stale-if-error=30" +} + +fn long_text() -> &'static str { + "extension-directive-with-a-value-that-exceeds-inline-storage=enabled" +} + +fn one_value() -> FieldValue { + FieldValue::from_static("value") +} + +fn four_values() -> FieldValue { + FieldValue::from_static("value") +} + +#[metabench::benchmark(SHORT_STRING, TEXT, "short_string")] +#[bench::short_text(setup = short_text)] +fn short_string(value: &str) -> usize { + let value = String::from(value); + black_box(value).len() +} + +#[metabench::benchmark(SHORT_COMPACT_STRING, TEXT, "short_compact_string")] +#[bench::short_text(setup = short_text)] +fn short_compact_string(value: &str) -> usize { + let value = CompactString::from(value); + black_box(value).len() +} + +#[metabench::benchmark(LONG_STRING, TEXT, "long_string")] +#[bench::long_text(setup = long_text)] +fn long_string(value: &str) -> usize { + let value = String::from(value); + black_box(value).len() +} + +#[metabench::benchmark(LONG_COMPACT_STRING, TEXT, "long_compact_string")] +#[bench::long_text(setup = long_text)] +fn long_compact_string(value: &str) -> usize { + let value = CompactString::from(value); + black_box(value).len() +} + +fn empty_http_map() -> HeaderMap { + HeaderMap::with_capacity(4) +} + +const SHORT_VALUE: &str = "application/json; charset=utf-8"; +const LONG_VALUE: &str = "multipart/form-data; boundary=----WebKitFormBoundary7MA4YWxkTrZu0gW; charset=utf-8"; +const DEFAULT_SINK_SHORT_VALUE: &[u8] = &[b's'; 16]; +const DEFAULT_SINK_SPILLED_VALUE: &[u8] = &[b'l'; 65]; +const STREAMED_0: &[u8] = &[]; +const STREAMED_16: &[u8] = &[b's'; 16]; +const STREAMED_32: &[u8] = &[b'm'; 32]; +const STREAMED_65: &[u8] = &[b'l'; 65]; +const STREAMED_4096: &[u8] = &[b'x'; 4096]; +const STREAMED_65536: &[u8] = &[b'y'; 65_536]; +static CUSTOM_APPEND_NAME: LazyLock = LazyLock::new(|| FieldName::from_static("x-cookie-batch")); + +#[derive(Default)] +struct MinimalSink { + values: Vec, +} + +impl FieldSource for MinimalSink { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_slice(name, &self.values) + } +} + +impl FieldSink for MinimalSink { + fn set_values(&mut self, _name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.values = values.into_iter().collect(); + Ok(()) + } + + fn append_values(&mut self, _name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.values.extend(values); + Ok(()) + } + + fn remove_values(&mut self, _name: &'static FieldName) { + self.values.clear(); + } +} + +type CustomSinkInput = (FieldValueRef<'static>, usize); +type CustomSinkCase = (&'static str, fn() -> CustomSinkInput); + +fn custom_sink_input(value: &'static [u8], line_count: usize) -> CustomSinkInput { + (FieldValueRef::new(value), line_count) +} + +macro_rules! custom_sink_case { + ($name:ident, $value:ident, $line_count:literal) => { + fn $name() -> CustomSinkInput { + custom_sink_input($value, $line_count) + } + }; +} + +custom_sink_case!(short_1, DEFAULT_SINK_SHORT_VALUE, 1); +custom_sink_case!(short_4, DEFAULT_SINK_SHORT_VALUE, 4); +custom_sink_case!(short_16, DEFAULT_SINK_SHORT_VALUE, 16); +custom_sink_case!(short_64, DEFAULT_SINK_SHORT_VALUE, 64); +custom_sink_case!(spilled_1, DEFAULT_SINK_SPILLED_VALUE, 1); +custom_sink_case!(spilled_4, DEFAULT_SINK_SPILLED_VALUE, 4); +custom_sink_case!(spilled_16, DEFAULT_SINK_SPILLED_VALUE, 16); +custom_sink_case!(spilled_64, DEFAULT_SINK_SPILLED_VALUE, 64); + +#[metabench::benchmark(MINIMAL_SINK_SET_ENCODED, CUSTOM_SINK_ENCODING, "minimal_sink_set_encoded")] +#[bench::short_1(setup = short_1)] +#[bench::short_4(setup = short_4)] +#[bench::short_16(setup = short_16)] +#[bench::short_64(setup = short_64)] +#[bench::spilled_1(setup = spilled_1)] +#[bench::spilled_4(setup = spilled_4)] +#[bench::spilled_16(setup = spilled_16)] +#[bench::spilled_64(setup = spilled_64)] +fn minimal_sink_set_encoded(input: CustomSinkInput) -> usize { + let (value, line_count) = input; + let mut sink = MinimalSink::default(); + sink.set_encoded( + &FieldName::SetCookie, + ValueRefsEncoder::new(iter::repeat_n(black_box(value), line_count)), + ) + .expect("minimal sink accepts encoded values"); + black_box(sink.values.len()) +} + +#[metabench::benchmark(MINIMAL_SINK_APPEND_ENCODED, CUSTOM_SINK_ENCODING, "minimal_sink_append_encoded")] +#[bench::short_1(setup = short_1)] +#[bench::short_4(setup = short_4)] +#[bench::short_16(setup = short_16)] +#[bench::short_64(setup = short_64)] +#[bench::spilled_1(setup = spilled_1)] +#[bench::spilled_4(setup = spilled_4)] +#[bench::spilled_16(setup = spilled_16)] +#[bench::spilled_64(setup = spilled_64)] +fn minimal_sink_append_encoded(input: CustomSinkInput) -> usize { + let (value, line_count) = input; + let mut sink = MinimalSink::default(); + for _ in 0..line_count { + sink.append_encoded(&FieldName::SetCookie, black_box(value)) + .expect("minimal sink accepts encoded values"); + } + black_box(sink.values.len()) +} + +fn short_value_map() -> (HeaderMap, &'static str) { + (HeaderMap::with_capacity(4), SHORT_VALUE) +} + +fn long_value_map() -> (HeaderMap, &'static str) { + (HeaderMap::with_capacity(4), LONG_VALUE) +} + +fn long_owned_map() -> (HeaderMap, FieldValue) { + (HeaderMap::with_capacity(4), FieldValue::from_static(LONG_VALUE)) +} + +fn streamed_map(bytes: &'static [u8]) -> (HeaderMap, &'static [u8]) { + (HeaderMap::with_capacity(4), bytes) +} + +macro_rules! streamed_setup { + ($name:ident, $bytes:ident) => { + fn $name() -> (HeaderMap, &'static [u8]) { + streamed_map($bytes) + } + }; +} + +streamed_setup!(streamed_0_map, STREAMED_0); +streamed_setup!(streamed_16_map, STREAMED_16); +streamed_setup!(streamed_32_map, STREAMED_32); +streamed_setup!(streamed_65_map, STREAMED_65); +streamed_setup!(streamed_4096_map, STREAMED_4096); +streamed_setup!(streamed_65536_map, STREAMED_65536); + +#[metabench::benchmark(HTTP_WRITER_BORROWED, MAP_WRITER, "http_writer_borrowed")] +#[bench::short(setup = short_value_map)] +#[bench::long(setup = long_value_map)] +fn http_writer_borrowed(state: (HeaderMap, &'static str)) -> MapOutput { + storage_operations::http_writer_borrowed(state) +} + +#[metabench::benchmark(HTTP_WRITER_STREAMED, MAP_WRITER, "http_writer_streamed")] +#[bench::short(setup = short_value_map)] +#[bench::long(setup = long_value_map)] +fn http_writer_streamed(state: (HeaderMap, &'static str)) -> MapOutput { + storage_operations::http_writer_streamed(state) +} + +#[metabench::benchmark(HTTP_WRITER_STREAMED_SIZED, MAP_WRITER, "http_writer_streamed_sized")] +#[bench::bytes_0(setup = streamed_0_map)] +#[bench::bytes_16(setup = streamed_16_map)] +#[bench::bytes_32(setup = streamed_32_map)] +#[bench::bytes_65(setup = streamed_65_map)] +#[bench::bytes_4096(setup = streamed_4096_map)] +#[bench::bytes_65536(setup = streamed_65536_map)] +fn http_writer_streamed_sized(state: (HeaderMap, &'static [u8])) -> MapOutput { + storage_operations::http_writer_streamed_sized(state) +} + +#[metabench::benchmark(HTTP_WRITER_MATERIALIZED, MAP_WRITER, "http_writer_materialized")] +#[bench::long(setup = long_owned_map)] +fn http_writer_materialized(state: (HeaderMap, FieldValue)) -> MapOutput { + storage_operations::http_writer_materialized(state) +} + +type AppendCase = (&'static str, fn() -> AppendInput); + +fn append_input(name: &'static FieldName, occupied: bool, count: usize) -> AppendInput { + let mut map = HeaderMap::with_capacity(128); + if occupied { + let http_name = if name == &FieldName::SetCookie { + http::header::SET_COOKIE + } else { + http::HeaderName::from_static("x-cookie-batch") + }; + map.insert(http_name, http::HeaderValue::from_static("existing")); + } + let values = iter::repeat_n(FieldValue::from_static("value"), count).collect(); + (map, name, values) +} + +macro_rules! append_setup { + ($name:ident, $field_name:expr, $occupied:literal, $count:literal) => { + fn $name() -> AppendInput { + append_input($field_name, $occupied, $count) + } + }; +} + +append_setup!(known_vacant_1, &FieldName::SetCookie, false, 1); +append_setup!(known_vacant_4, &FieldName::SetCookie, false, 4); +append_setup!(known_vacant_16, &FieldName::SetCookie, false, 16); +append_setup!(known_vacant_64, &FieldName::SetCookie, false, 64); +append_setup!(known_occupied_1, &FieldName::SetCookie, true, 1); +append_setup!(known_occupied_4, &FieldName::SetCookie, true, 4); +append_setup!(known_occupied_16, &FieldName::SetCookie, true, 16); +append_setup!(known_occupied_64, &FieldName::SetCookie, true, 64); +append_setup!(custom_vacant_1, &CUSTOM_APPEND_NAME, false, 1); +append_setup!(custom_vacant_4, &CUSTOM_APPEND_NAME, false, 4); +append_setup!(custom_vacant_16, &CUSTOM_APPEND_NAME, false, 16); +append_setup!(custom_vacant_64, &CUSTOM_APPEND_NAME, false, 64); +append_setup!(custom_occupied_1, &CUSTOM_APPEND_NAME, true, 1); +append_setup!(custom_occupied_4, &CUSTOM_APPEND_NAME, true, 4); +append_setup!(custom_occupied_16, &CUSTOM_APPEND_NAME, true, 16); +append_setup!(custom_occupied_64, &CUSTOM_APPEND_NAME, true, 64); + +#[metabench::benchmark(HTTP_APPEND_VALUES, MAP_APPEND, "http_append_values")] +#[bench::known_vacant_1(setup = known_vacant_1)] +#[bench::known_vacant_4(setup = known_vacant_4)] +#[bench::known_vacant_16(setup = known_vacant_16)] +#[bench::known_vacant_64(setup = known_vacant_64)] +#[bench::known_occupied_1(setup = known_occupied_1)] +#[bench::known_occupied_4(setup = known_occupied_4)] +#[bench::known_occupied_16(setup = known_occupied_16)] +#[bench::known_occupied_64(setup = known_occupied_64)] +#[bench::custom_vacant_1(setup = custom_vacant_1)] +#[bench::custom_vacant_4(setup = custom_vacant_4)] +#[bench::custom_vacant_16(setup = custom_vacant_16)] +#[bench::custom_vacant_64(setup = custom_vacant_64)] +#[bench::custom_occupied_1(setup = custom_occupied_1)] +#[bench::custom_occupied_4(setup = custom_occupied_4)] +#[bench::custom_occupied_16(setup = custom_occupied_16)] +#[bench::custom_occupied_64(setup = custom_occupied_64)] +fn http_append_values(input: AppendInput) -> MapOutput { + storage_operations::http_append_values(input) +} + +#[metabench::benchmark(CONTENT_LENGTH_MATERIALIZED, DEFERRED_INSERTION, "content_length_materialized")] +#[bench::materialized(setup = empty_http_map)] +fn content_length_materialized(map: HeaderMap) -> MapOutput { + storage_operations::content_length_materialized(map) +} + +#[metabench::benchmark(CONTENT_LENGTH_DEFERRED_HTTP, DEFERRED_INSERTION, "content_length_deferred_http")] +#[bench::deferred_http(setup = empty_http_map)] +fn content_length_deferred_http(map: HeaderMap) -> MapOutput { + storage_operations::content_length_deferred_http(map) +} + +#[metabench::benchmark(ONE_VEC, REPEATED_VALUES, "one_vec")] +#[bench::one_value(setup = one_value)] +fn one_vec(value: FieldValue) -> usize { + let values = vec![value]; + black_box(values).len() +} + +#[metabench::benchmark(ONE_SMALLVEC, REPEATED_VALUES, "one_smallvec")] +#[bench::one_value(setup = one_value)] +fn one_smallvec(value: FieldValue) -> usize { + let values = SmallVec::<[FieldValue; 1]>::from_buf([value]); + black_box(values).len() +} + +#[metabench::benchmark(FOUR_VEC, REPEATED_VALUES, "four_vec")] +#[bench::four_values(setup = four_values)] +fn four_vec(value: FieldValue) -> usize { + let values = vec![value.clone(), value.clone(), value.clone(), value]; + black_box(values).len() +} + +#[metabench::benchmark(FOUR_SMALLVEC, REPEATED_VALUES, "four_smallvec")] +#[bench::four_values(setup = four_values)] +fn four_smallvec(value: FieldValue) -> usize { + let values = SmallVec::<[FieldValue; 1]>::from_vec(vec![value.clone(), value.clone(), value.clone(), value]); + black_box(values).len() +} + +#[metabench::benchmark(ONE_VEC_ENCODE, INSERTION, "one_vec_encode")] +#[bench::one_encode(setup = one_value)] +fn one_vec_encode(value: FieldValue) -> usize { + let values = EncodedValues::from_vec(vec![value]); + black_box(values).len() +} + +#[metabench::benchmark(ONE_SMALLVEC_ENCODE, INSERTION, "one_smallvec_encode")] +#[bench::one_encode(setup = one_value)] +fn one_smallvec_encode(value: FieldValue) -> usize { + let values = SmallVec::<[FieldValue; 1]>::from_buf([value]) + .into_iter() + .collect::(); + black_box(values).len() +} + +#[metabench::benchmark(SHORT_STRING_TIME, STORAGE_TEXT, "short_string")] +fn short_string_time() -> String { + black_box(String::from(black_box("stale-if-error=30"))) +} + +#[metabench::benchmark(SHORT_COMPACT_STRING_TIME, STORAGE_TEXT, "short_compact_string")] +fn short_compact_string_time() -> CompactString { + black_box(CompactString::from(black_box("stale-if-error=30"))) +} + +#[metabench::benchmark(LONG_STRING_TIME, STORAGE_TEXT, "long_string")] +fn long_string_time() -> String { + black_box(String::from(black_box( + "extension-directive-with-a-value-that-exceeds-inline-storage=enabled", + ))) +} + +#[metabench::benchmark(LONG_COMPACT_STRING_TIME, STORAGE_TEXT, "long_compact_string")] +fn long_compact_string_time() -> CompactString { + black_box(CompactString::from(black_box( + "extension-directive-with-a-value-that-exceeds-inline-storage=enabled", + ))) +} + +fn reused_value() -> &'static FieldValue { + static VALUE: OnceLock = OnceLock::new(); + VALUE.get_or_init(|| FieldValue::from_static("value")) +} + +#[metabench::benchmark(ONE_VEC_TIME, STORAGE_REPEATED_VALUES, "one_vec")] +#[bench::one_value(setup = reused_value)] +fn one_vec_time(value: &FieldValue) -> Vec { + black_box(vec![black_box(value.clone())]) +} + +#[metabench::benchmark(ONE_SMALLVEC_TIME, STORAGE_REPEATED_VALUES, "one_smallvec")] +#[bench::one_value(setup = reused_value)] +fn one_smallvec_time(value: &FieldValue) -> SmallVec<[FieldValue; 1]> { + black_box(SmallVec::<[FieldValue; 1]>::from_buf([black_box(value.clone())])) +} + +#[metabench::benchmark(FOUR_VEC_TIME, STORAGE_REPEATED_VALUES, "four_vec")] +#[bench::four_values(setup = reused_value)] +fn four_vec_time(value: &FieldValue) -> Vec { + black_box(vec![value.clone(), value.clone(), value.clone(), black_box(value.clone())]) +} + +#[metabench::benchmark(FOUR_SMALLVEC_TIME, STORAGE_REPEATED_VALUES, "four_smallvec")] +#[bench::four_values(setup = reused_value)] +fn four_smallvec_time(value: &FieldValue) -> SmallVec<[FieldValue; 1]> { + black_box(SmallVec::<[FieldValue; 1]>::from_vec(vec![ + value.clone(), + value.clone(), + value.clone(), + black_box(value.clone()), + ])) +} + +#[metabench::benchmark(ONE_VEC_ENCODE_TIME, STORAGE_INSERTION, "one_vec_encode")] +#[bench::one_encode(setup = reused_value)] +fn one_vec_encode_time(value: &FieldValue) -> EncodedValues { + black_box(EncodedValues::from_vec(vec![black_box(value.clone())])) +} + +#[metabench::benchmark(ONE_SMALLVEC_ENCODE_TIME, STORAGE_INSERTION, "one_smallvec_encode")] +#[bench::one_encode(setup = reused_value)] +fn one_smallvec_encode_time(value: &FieldValue) -> EncodedValues { + black_box( + SmallVec::<[FieldValue; 1]>::from_buf([black_box(value.clone())]) + .into_iter() + .collect::(), + ) +} + +fn register_storage_text(criterion: &mut Criterion) { + let mut storage_text = criterion.benchmark_group(STORAGE_TEXT); + storage_text.bench_function(SHORT_STRING_TIME.benchmark_name(), |bencher| { + bencher.iter(short_string_time); + }); + storage_text.bench_function(SHORT_COMPACT_STRING_TIME.benchmark_name(), |bencher| { + bencher.iter(short_compact_string_time); + }); + storage_text.bench_function(LONG_STRING_TIME.benchmark_name(), |bencher| { + bencher.iter(long_string_time); + }); + storage_text.bench_function(LONG_COMPACT_STRING_TIME.benchmark_name(), |bencher| { + bencher.iter(long_compact_string_time); + }); + storage_text.finish(); +} + +fn register_text(criterion: &mut Criterion) { + let mut text = criterion.benchmark_group(TEXT); + text.bench_function(BenchmarkId::new(SHORT_STRING.benchmark_name(), "short_text"), |bencher| { + bencher.iter_batched(short_text, short_string, BatchSize::SmallInput); + }); + text.bench_function(BenchmarkId::new(SHORT_COMPACT_STRING.benchmark_name(), "short_text"), |bencher| { + bencher.iter_batched(short_text, short_compact_string, BatchSize::SmallInput); + }); + text.bench_function(BenchmarkId::new(LONG_STRING.benchmark_name(), "long_text"), |bencher| { + bencher.iter_batched(long_text, long_string, BatchSize::SmallInput); + }); + text.bench_function(BenchmarkId::new(LONG_COMPACT_STRING.benchmark_name(), "long_text"), |bencher| { + bencher.iter_batched(long_text, long_compact_string, BatchSize::SmallInput); + }); + text.finish(); +} + +fn register_map_writer(criterion: &mut Criterion) { + let mut map_writer = criterion.benchmark_group(MAP_WRITER); + map_writer.bench_function(BenchmarkId::new(HTTP_WRITER_BORROWED.benchmark_name(), "short"), |bencher| { + bencher.iter_batched(short_value_map, http_writer_borrowed, BatchSize::SmallInput); + }); + map_writer.bench_function(BenchmarkId::new(HTTP_WRITER_BORROWED.benchmark_name(), "long"), |bencher| { + bencher.iter_batched(long_value_map, http_writer_borrowed, BatchSize::SmallInput); + }); + map_writer.bench_function(BenchmarkId::new(HTTP_WRITER_STREAMED.benchmark_name(), "short"), |bencher| { + bencher.iter_batched(short_value_map, http_writer_streamed, BatchSize::SmallInput); + }); + map_writer.bench_function(BenchmarkId::new(HTTP_WRITER_STREAMED.benchmark_name(), "long"), |bencher| { + bencher.iter_batched(long_value_map, http_writer_streamed, BatchSize::SmallInput); + }); + let streamed_cases = [ + ("bytes_0", streamed_0_map as fn() -> (HeaderMap, &'static [u8])), + ("bytes_16", streamed_16_map), + ("bytes_32", streamed_32_map), + ("bytes_65", streamed_65_map), + ("bytes_4096", streamed_4096_map), + ("bytes_65536", streamed_65536_map), + ]; + for &(case, setup) in &streamed_cases { + map_writer.bench_function(BenchmarkId::new(HTTP_WRITER_STREAMED_SIZED.benchmark_name(), case), |bencher| { + bencher.iter_batched(setup, http_writer_streamed_sized, BatchSize::SmallInput); + }); + } + map_writer.bench_function(BenchmarkId::new(HTTP_WRITER_MATERIALIZED.benchmark_name(), "long"), |bencher| { + bencher.iter_batched(long_owned_map, http_writer_materialized, BatchSize::SmallInput); + }); + map_writer.finish(); +} + +fn register_map_append(criterion: &mut Criterion) { + let cases: [AppendCase; 16] = [ + ("known_vacant_1", known_vacant_1), + ("known_vacant_4", known_vacant_4), + ("known_vacant_16", known_vacant_16), + ("known_vacant_64", known_vacant_64), + ("known_occupied_1", known_occupied_1), + ("known_occupied_4", known_occupied_4), + ("known_occupied_16", known_occupied_16), + ("known_occupied_64", known_occupied_64), + ("custom_vacant_1", custom_vacant_1), + ("custom_vacant_4", custom_vacant_4), + ("custom_vacant_16", custom_vacant_16), + ("custom_vacant_64", custom_vacant_64), + ("custom_occupied_1", custom_occupied_1), + ("custom_occupied_4", custom_occupied_4), + ("custom_occupied_16", custom_occupied_16), + ("custom_occupied_64", custom_occupied_64), + ]; + let mut group = criterion.benchmark_group(MAP_APPEND); + for &(case, setup) in &cases { + group.bench_function(BenchmarkId::new(HTTP_APPEND_VALUES.benchmark_name(), case), |bencher| { + bencher.iter_batched(setup, http_append_values, BatchSize::SmallInput); + }); + } + group.finish(); +} + +fn register_deferred_insertion(criterion: &mut Criterion) { + let mut deferred = criterion.benchmark_group(DEFERRED_INSERTION); + deferred.bench_function( + BenchmarkId::new(CONTENT_LENGTH_MATERIALIZED.benchmark_name(), "materialized"), + |bencher| { + bencher.iter_batched(empty_http_map, content_length_materialized, BatchSize::SmallInput); + }, + ); + deferred.bench_function( + BenchmarkId::new(CONTENT_LENGTH_DEFERRED_HTTP.benchmark_name(), "deferred_http"), + |bencher| { + bencher.iter_batched(empty_http_map, content_length_deferred_http, BatchSize::SmallInput); + }, + ); + deferred.finish(); +} + +fn register_storage_repeated_values(criterion: &mut Criterion, value: &FieldValue) { + let mut storage_repeated = criterion.benchmark_group(STORAGE_REPEATED_VALUES); + storage_repeated.bench_function(BenchmarkId::new(ONE_VEC_TIME.benchmark_name(), "one_value"), |bencher| { + bencher.iter(|| one_vec_time(value)); + }); + storage_repeated.bench_function(BenchmarkId::new(ONE_SMALLVEC_TIME.benchmark_name(), "one_value"), |bencher| { + bencher.iter(|| one_smallvec_time(value)); + }); + storage_repeated.bench_function(BenchmarkId::new(FOUR_VEC_TIME.benchmark_name(), "four_values"), |bencher| { + bencher.iter(|| four_vec_time(value)); + }); + storage_repeated.bench_function(BenchmarkId::new(FOUR_SMALLVEC_TIME.benchmark_name(), "four_values"), |bencher| { + bencher.iter(|| four_smallvec_time(value)); + }); + storage_repeated.finish(); +} + +fn register_repeated_values(criterion: &mut Criterion) { + let mut repeated = criterion.benchmark_group(REPEATED_VALUES); + repeated.bench_function(BenchmarkId::new(ONE_VEC.benchmark_name(), "one_value"), |bencher| { + bencher.iter_batched(one_value, one_vec, BatchSize::SmallInput); + }); + repeated.bench_function(BenchmarkId::new(ONE_SMALLVEC.benchmark_name(), "one_value"), |bencher| { + bencher.iter_batched(one_value, one_smallvec, BatchSize::SmallInput); + }); + repeated.bench_function(BenchmarkId::new(FOUR_VEC.benchmark_name(), "four_values"), |bencher| { + bencher.iter_batched(four_values, four_vec, BatchSize::SmallInput); + }); + repeated.bench_function(BenchmarkId::new(FOUR_SMALLVEC.benchmark_name(), "four_values"), |bencher| { + bencher.iter_batched(four_values, four_smallvec, BatchSize::SmallInput); + }); + repeated.finish(); +} + +fn register_storage_insertion(criterion: &mut Criterion, value: &FieldValue) { + let mut storage_insertion = criterion.benchmark_group(STORAGE_INSERTION); + storage_insertion.bench_function(BenchmarkId::new(ONE_VEC_ENCODE_TIME.benchmark_name(), "one_encode"), |bencher| { + bencher.iter(|| one_vec_encode_time(value)); + }); + storage_insertion.bench_function( + BenchmarkId::new(ONE_SMALLVEC_ENCODE_TIME.benchmark_name(), "one_encode"), + |bencher| bencher.iter(|| one_smallvec_encode_time(value)), + ); + storage_insertion.finish(); +} + +fn register_insertion(criterion: &mut Criterion) { + let mut insertion = criterion.benchmark_group(INSERTION); + insertion.bench_function(BenchmarkId::new(ONE_VEC_ENCODE.benchmark_name(), "one_encode"), |bencher| { + bencher.iter_batched(one_value, one_vec_encode, BatchSize::SmallInput); + }); + insertion.bench_function(BenchmarkId::new(ONE_SMALLVEC_ENCODE.benchmark_name(), "one_encode"), |bencher| { + bencher.iter_batched(one_value, one_smallvec_encode, BatchSize::SmallInput); + }); + insertion.finish(); +} + +fn criterion_benchmarks(criterion: &mut Criterion) { + register_storage_text(criterion); + register_text(criterion); + register_map_writer(criterion); + register_map_append(criterion); + register_deferred_insertion(criterion); + + let value = FieldValue::from_static("value"); + register_storage_repeated_values(criterion, &value); + register_repeated_values(criterion); + register_storage_insertion(criterion, &value); + register_insertion(criterion); + register_custom_sink_encoding(criterion); +} + +fn register_custom_sink_encoding(criterion: &mut Criterion) { + let cases: [CustomSinkCase; 8] = [ + ("short_1", short_1), + ("short_4", short_4), + ("short_16", short_16), + ("short_64", short_64), + ("spilled_1", spilled_1), + ("spilled_4", spilled_4), + ("spilled_16", spilled_16), + ("spilled_64", spilled_64), + ]; + let mut group = criterion.benchmark_group(CUSTOM_SINK_ENCODING); + for &(case, setup) in &cases { + group.bench_function(BenchmarkId::new(MINIMAL_SINK_SET_ENCODED.benchmark_name(), case), |bencher| { + bencher.iter_batched(setup, minimal_sink_set_encoded, BatchSize::SmallInput); + }); + } + for &(case, setup) in &cases { + group.bench_function(BenchmarkId::new(MINIMAL_SINK_APPEND_ENCODED.benchmark_name(), case), |bencher| { + bencher.iter_batched(setup, minimal_sink_append_encoded, BatchSize::SmallInput); + }); + } + group.finish(); +} + +metabench::main!( + criterion = criterion_benchmarks, + benchmarks = [ + SHORT_STRING, + SHORT_COMPACT_STRING, + LONG_STRING, + LONG_COMPACT_STRING, + HTTP_WRITER_BORROWED, + HTTP_WRITER_STREAMED, + HTTP_WRITER_STREAMED_SIZED, + HTTP_WRITER_MATERIALIZED, + HTTP_APPEND_VALUES, + CONTENT_LENGTH_MATERIALIZED, + CONTENT_LENGTH_DEFERRED_HTTP, + ONE_VEC, + ONE_SMALLVEC, + FOUR_VEC, + FOUR_SMALLVEC, + ONE_VEC_ENCODE, + ONE_SMALLVEC_ENCODE, + SHORT_STRING_TIME, + SHORT_COMPACT_STRING_TIME, + LONG_STRING_TIME, + LONG_COMPACT_STRING_TIME, + ONE_VEC_TIME, + ONE_SMALLVEC_TIME, + FOUR_VEC_TIME, + FOUR_SMALLVEC_TIME, + ONE_VEC_ENCODE_TIME, + ONE_SMALLVEC_ENCODE_TIME, + MINIMAL_SINK_SET_ENCODED, + MINIMAL_SINK_APPEND_ENCODED, + ], +); diff --git a/crates/http_headers/benches/http_headers_typed_tokens.rs b/crates/http_headers/benches/http_headers_typed_tokens.rs new file mode 100644 index 000000000..bb28609de --- /dev/null +++ b/crates/http_headers/benches/http_headers_typed_tokens.rs @@ -0,0 +1,158 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Typed token reads versus explicit owned-token materialization. + +#![expect(clippy::unwrap_used, reason = "benchmark fixtures and validated token conversions must succeed")] + +use std::hint::black_box; +use std::sync::OnceLock; + +use criterion::Criterion; +use http::{HeaderMap, HeaderValue, Method}; +use http_headers::FieldName; +use http_headers::headers::{Allow, MethodView, Vary}; + +fn allow_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| { + let mut map = HeaderMap::new(); + map.insert("allow", HeaderValue::from_static("GET, HEAD, POST, CUSTOM, PURGE")); + map + }) +} + +fn vary_map() -> &'static HeaderMap { + static MAP: OnceLock = OnceLock::new(); + MAP.get_or_init(|| { + let mut map = HeaderMap::new(); + map.insert("vary", HeaderValue::from_static("Accept-Encoding, X-Tenant, X-Region")); + map + }) +} + +#[metabench::benchmark(ALLOW_TYPED, "http_headers_typed_tokens/reads", "allow_typed", gungraun_setup = allow_map)] +fn allow_typed(map: &'static HeaderMap) -> bool { + let value = Allow::view(black_box(map)).unwrap().unwrap(); + let query = MethodView::new(black_box("PURGE")).unwrap(); + black_box(value.methods().any(|method| method == query)) +} + +#[metabench::benchmark(ALLOW_MATERIALIZED, "http_headers_typed_tokens/reads", "allow_materialized", gungraun_setup = allow_map)] +fn allow_materialized(map: &'static HeaderMap) -> bool { + let value = Allow::view(black_box(map)).unwrap().unwrap(); + let query = Method::from_bytes(black_box(b"PURGE")).unwrap(); + black_box(value.items().any(|method| Method::from_bytes(method).unwrap() == query)) +} + +#[metabench::benchmark(ALLOW_TYPED_8, "http_headers_typed_tokens/reads", "allow_typed_8", gungraun_setup = allow_map)] +fn allow_typed_8(map: &'static HeaderMap) -> usize { + let value = Allow::view(black_box(map)).unwrap().unwrap(); + let query = MethodView::new("PURGE").unwrap(); + (0..8) + .map(|_| usize::from(black_box(value.methods().any(|method| method == black_box(query))))) + .sum() +} + +#[metabench::benchmark(ALLOW_MATERIALIZED_8, "http_headers_typed_tokens/reads", "allow_materialized_8", gungraun_setup = allow_map)] +fn allow_materialized_8(map: &'static HeaderMap) -> usize { + let value = Allow::view(black_box(map)).unwrap().unwrap(); + let query = Method::from_bytes(b"PURGE").unwrap(); + (0..8) + .map(|_| { + usize::from(black_box( + value + .items() + .any(|method| Method::from_bytes(method).unwrap() == *black_box(&query)), + )) + }) + .sum() +} + +#[metabench::benchmark(VARY_TYPED, "http_headers_typed_tokens/reads", "vary_typed", gungraun_setup = vary_map)] +fn vary_typed(map: &'static HeaderMap) -> bool { + let value = Vary::view(black_box(map)).unwrap().unwrap(); + black_box(value.entries().any(|entry| { + entry.is_wildcard() + || entry + .field_name() + .is_some_and(|name| name.eq_ignore_ascii_case(black_box("x-region"))) + })) +} + +#[metabench::benchmark(VARY_MATERIALIZED, "http_headers_typed_tokens/reads", "vary_materialized", gungraun_setup = vary_map)] +fn vary_materialized(map: &'static HeaderMap) -> bool { + let value = Vary::view(black_box(map)).unwrap().unwrap(); + black_box(value.items().any(|name| { + name == b"*" + || FieldName::try_from_bytes(name) + .unwrap() + .as_str() + .eq_ignore_ascii_case(black_box("x-region")) + })) +} + +#[metabench::benchmark(VARY_TYPED_8, "http_headers_typed_tokens/reads", "vary_typed_8", gungraun_setup = vary_map)] +fn vary_typed_8(map: &'static HeaderMap) -> usize { + let value = Vary::view(black_box(map)).unwrap().unwrap(); + (0..8) + .map(|_| { + usize::from(black_box(value.entries().any(|entry| { + entry.is_wildcard() + || entry + .field_name() + .is_some_and(|name| name.eq_ignore_ascii_case(black_box("x-region"))) + }))) + }) + .sum() +} + +#[metabench::benchmark(VARY_MATERIALIZED_8, "http_headers_typed_tokens/reads", "vary_materialized_8", gungraun_setup = vary_map)] +fn vary_materialized_8(map: &'static HeaderMap) -> usize { + let value = Vary::view(black_box(map)).unwrap().unwrap(); + (0..8) + .map(|_| { + usize::from(black_box(value.items().any(|name| { + name == b"*" + || FieldName::try_from_bytes(name) + .unwrap() + .as_str() + .eq_ignore_ascii_case(black_box("x-region")) + }))) + }) + .sum() +} + +fn criterion_benchmarks(criterion: &mut Criterion) { + let mut group = criterion.benchmark_group(ALLOW_TYPED.group_name()); + macro_rules! register { + ($id:ident, $function:ident, $setup:ident) => { + group.bench_function($id.benchmark_name(), |bencher| { + let input = $setup(); + bencher.iter(|| $function(input)); + }); + }; + } + register!(ALLOW_TYPED, allow_typed, allow_map); + register!(ALLOW_MATERIALIZED, allow_materialized, allow_map); + register!(ALLOW_TYPED_8, allow_typed_8, allow_map); + register!(ALLOW_MATERIALIZED_8, allow_materialized_8, allow_map); + register!(VARY_TYPED, vary_typed, vary_map); + register!(VARY_MATERIALIZED, vary_materialized, vary_map); + register!(VARY_TYPED_8, vary_typed_8, vary_map); + register!(VARY_MATERIALIZED_8, vary_materialized_8, vary_map); +} + +metabench::main!( + criterion = criterion_benchmarks, + benchmarks = [ + ALLOW_TYPED, + ALLOW_MATERIALIZED, + ALLOW_TYPED_8, + ALLOW_MATERIALIZED_8, + VARY_TYPED, + VARY_MATERIALIZED, + VARY_TYPED_8, + VARY_MATERIALIZED_8, + ], +); diff --git a/crates/http_headers/benches/http_headers_websocket_shapes.rs b/crates/http_headers/benches/http_headers_websocket_shapes.rs new file mode 100644 index 000000000..f6a20ccbc --- /dev/null +++ b/crates/http_headers/benches/http_headers_websocket_shapes.rs @@ -0,0 +1,158 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! WebSocket decode shapes across HTTP/custom sources and ownership modes. + +use http_headers::DecodeErrorKind; +use http_headers::headers::{SecWebSocketAccept, SecWebSocketExtensions, SecWebSocketKey, SecWebSocketProtocol, SecWebSocketVersion}; + +#[path = "http_headers_shapes_common.rs"] +mod shapes; + +use shapes::{Expectations, Expected}; + +shapes::define_shapes!( + "http_headers_websocket_shapes/parse"; + (websocket_accept_standard, SecWebSocketAccept, &["s3pPLMBiTxaQ9kYGzzhZRbK+xOo="], Strict, Expected::Valid), + (websocket_accept_zero, SecWebSocketAccept, &["AAAAAAAAAAAAAAAAAAAAAAAAAAA="], Strict, Expected::Valid), + (websocket_accept_padding, SecWebSocketAccept, &["s3pPLMBiTxaQ9kYGzzhZRbK+xOp="], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_accept_alphabet, SecWebSocketAccept, &["!3pPLMBiTxaQ9kYGzzhZRbK+xOo="], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_accept_alphabet_tail, SecWebSocketAccept, &["s3pPLMBiTxaQ9kYGzzhZRbK+xO!="], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_accept_missing_pad, SecWebSocketAccept, &["s3pPLMBiTxaQ9kYGzzhZRbK+xOo!"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_accept_short, SecWebSocketAccept, &["s3pPLMBiTxaQ9kYG"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_accept_repeated, SecWebSocketAccept, &["s3pPLMBiTxaQ9kYGzzhZRbK+xOo=", "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + (websocket_accept_absent, SecWebSocketAccept, &[], Strict, Expected::Absent), + (websocket_key_standard, SecWebSocketKey, &["dGhlIHNhbXBsZSBub25jZQ=="], Strict, Expected::Valid), + (websocket_key_zero, SecWebSocketKey, &["AAAAAAAAAAAAAAAAAAAAAA=="], Strict, Expected::Valid), + (websocket_key_padding, SecWebSocketKey, &["dGhlIHNhbXBsZSBub25jZR=="], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_key_alphabet, SecWebSocketKey, &["!GhlIHNhbXBsZSBub25jZQ=="], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_key_alphabet_tail, SecWebSocketKey, &["dGhlIHNhbXBsZSBub25jZ!=="], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_key_missing_first_pad, SecWebSocketKey, &["dGhlIHNhbXBsZSBub25jZQ!="], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_key_missing_last_pad, SecWebSocketKey, &["dGhlIHNhbXBsZSBub25jZQ=!"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_key_short, SecWebSocketKey, &["dGhlIHNhbXBsZQ=="], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_key_repeated, SecWebSocketKey, &["dGhlIHNhbXBsZSBub25jZQ==", "dGhlIHNhbXBsZSBub25jZQ=="], Strict, Expected::Error(DecodeErrorKind::UnexpectedMultipleValues)), + (websocket_key_absent, SecWebSocketKey, &[], Strict, Expected::Absent), + (websocket_protocol_standard, SecWebSocketProtocol, &["graphql-transport-ws"], Strict, Expected::Valid), + (websocket_protocol_graphql, SecWebSocketProtocol, &["graphql-ws"], Strict, Expected::Valid), + (websocket_protocol_graphql_advertisement, SecWebSocketProtocol, &["graphql-transport-ws, graphql-ws"], Strict, Expected::Valid), + (websocket_protocol_graphql_near_first, SecWebSocketProtocol, &["Xraphql-ws"], Strict, Expected::Valid), + (websocket_protocol_graphql_near_second, SecWebSocketProtocol, &["gRaphql-ws"], Strict, Expected::Valid), + (websocket_protocol_graphql_near_middle, SecWebSocketProtocol, &["grapHql-ws"], Strict, Expected::Valid), + (websocket_protocol_graphql_near_last, SecWebSocketProtocol, &["graphql-wS"], Strict, Expected::Valid), + (websocket_protocol_graphql_invalid_last, SecWebSocketProtocol, &["graphql-w/"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_protocol_list_near_first, SecWebSocketProtocol, &["Chat, superchat"], Strict, Expected::Valid), + (websocket_protocol_list_near_second, SecWebSocketProtocol, &["cHat, superchat"], Strict, Expected::Valid), + (websocket_protocol_list_near_middle, SecWebSocketProtocol, &["chat, suPerchat"], Strict, Expected::Valid), + (websocket_protocol_list_near_last, SecWebSocketProtocol, &["chat, superchaT"], Strict, Expected::Valid), + (websocket_protocol_list_invalid_first, SecWebSocketProtocol, &["/hat, superchat"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_protocol_list_invalid_last, SecWebSocketProtocol, &["chat, supercha/"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_protocol_transport_near_first, SecWebSocketProtocol, &["Graphql-transport-ws"], Strict, Expected::Valid), + (websocket_protocol_transport_near_second, SecWebSocketProtocol, &["gRaphql-transport-ws"], Strict, Expected::Valid), + (websocket_protocol_transport_near_middle, SecWebSocketProtocol, &["graphql-transPort-ws"], Strict, Expected::Valid), + (websocket_protocol_transport_near_last, SecWebSocketProtocol, &["graphql-transport-wS"], Strict, Expected::Valid), + (websocket_protocol_transport_invalid_first, SecWebSocketProtocol, &["/raphql-transport-ws"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_protocol_transport_invalid_last, SecWebSocketProtocol, &["graphql-transport-w/"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_protocol_advertisement_near_first, SecWebSocketProtocol, &["Graphql-transport-ws, graphql-ws"], Strict, Expected::Valid), + (websocket_protocol_advertisement_near_second, SecWebSocketProtocol, &["gRaphql-transport-ws, graphql-ws"], Strict, Expected::Valid), + (websocket_protocol_advertisement_near_middle, SecWebSocketProtocol, &["graphql-transport-wS, graphql-ws"], Strict, Expected::Valid), + (websocket_protocol_advertisement_near_last, SecWebSocketProtocol, &["graphql-transport-ws, graphql-wS"], Strict, Expected::Valid), + (websocket_protocol_advertisement_invalid_first, SecWebSocketProtocol, &["/raphql-transport-ws, graphql-ws"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_protocol_advertisement_invalid_last, SecWebSocketProtocol, &["graphql-transport-ws, graphql-w/"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_protocol_invalid_first, SecWebSocketProtocol, &["/raphql-ws"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_protocol_token, SecWebSocketProtocol, &["binary.protocol-v2"], Strict, Expected::Valid), + (websocket_protocol_list, SecWebSocketProtocol, &["chat, superchat"], Strict, Expected::Valid), + (websocket_protocol_ows, SecWebSocketProtocol, &[" \tchat , binary.protocol-v2\t"], Strict, Expected::Valid), + (websocket_protocol_repeated, SecWebSocketProtocol, &["chat", "superchat", "binary.protocol-v2"], Strict, Expected::Valid), + (websocket_protocol_quote, SecWebSocketProtocol, &["chat, \"quoted\""], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_protocol_empty, SecWebSocketProtocol, &[""], Strict, Expectations { + http_owned: Expected::Error(DecodeErrorKind::InvalidSyntax), + http_borrowed: Expected::Error(DecodeErrorKind::MissingValue), + raw_owned: Expected::Error(DecodeErrorKind::InvalidSyntax), + raw_borrowed: Expected::Error(DecodeErrorKind::MissingValue), + }), + (websocket_protocol_empty_members, SecWebSocketProtocol, &[" , , \t"], Strict, Expectations { + http_owned: Expected::Error(DecodeErrorKind::InvalidSyntax), + http_borrowed: Expected::Error(DecodeErrorKind::MissingValue), + raw_owned: Expected::Error(DecodeErrorKind::InvalidSyntax), + raw_borrowed: Expected::Error(DecodeErrorKind::MissingValue), + }), + (websocket_protocol_unterminated, SecWebSocketProtocol, &["chat", "\"quoted"], Strict, Expectations { + http_owned: Expected::Error(DecodeErrorKind::InvalidToken), + http_borrowed: Expected::Error(DecodeErrorKind::UnterminatedQuote), + raw_owned: Expected::Error(DecodeErrorKind::InvalidToken), + raw_borrowed: Expected::Error(DecodeErrorKind::UnterminatedQuote), + }), + (websocket_protocol_absent, SecWebSocketProtocol, &[], Strict, Expected::Absent), + (websocket_extensions_standard, SecWebSocketExtensions, &["permessage-deflate"], Strict, Expected::Valid), + (websocket_extensions_window, SecWebSocketExtensions, &["permessage-deflate; client_max_window_bits"], Strict, Expected::Valid), + (websocket_extensions_takeover, SecWebSocketExtensions, &["permessage-deflate; server_no_context_takeover; client_no_context_takeover"], Strict, Expected::Valid), + (websocket_extensions_near_first, SecWebSocketExtensions, &["Permessage-deflate"], Strict, Expected::Valid), + (websocket_extensions_near_second, SecWebSocketExtensions, &["pErmessage-deflate"], Strict, Expected::Valid), + (websocket_extensions_near_middle, SecWebSocketExtensions, &["permessAge-deflate"], Strict, Expected::Valid), + (websocket_extensions_near_last, SecWebSocketExtensions, &["permessage-deflatE"], Strict, Expected::Valid), + (websocket_extensions_window_near_first, SecWebSocketExtensions, &["Permessage-deflate; client_max_window_bits"], Strict, Expected::Valid), + (websocket_extensions_window_near_second, SecWebSocketExtensions, &["pErmessage-deflate; client_max_window_bits"], Strict, Expected::Valid), + (websocket_extensions_window_invalid_first, SecWebSocketExtensions, &["/ermessage-deflate; client_max_window_bits"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_extensions_window_near_middle, SecWebSocketExtensions, &["permessage-deflate; clieNt_max_window_bits"], Strict, Expected::Valid), + (websocket_extensions_window_near_last, SecWebSocketExtensions, &["permessage-deflate; client_max_window_bitS"], Strict, Expected::Valid), + (websocket_extensions_takeover_near_first, SecWebSocketExtensions, &["Permessage-deflate; server_no_context_takeover; client_no_context_takeover"], Strict, Expected::Valid), + (websocket_extensions_takeover_near_second, SecWebSocketExtensions, &["pErmessage-deflate; server_no_context_takeover; client_no_context_takeover"], Strict, Expected::Valid), + (websocket_extensions_takeover_invalid_first, SecWebSocketExtensions, &["/ermessage-deflate; server_no_context_takeover; client_no_context_takeover"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_extensions_takeover_near_middle, SecWebSocketExtensions, &["permessage-deflate; server_no_context_Takeover; client_no_context_takeover"], Strict, Expected::Valid), + (websocket_extensions_takeover_near_last, SecWebSocketExtensions, &["permessage-deflate; server_no_context_takeover; client_no_context_takeoveR"], Strict, Expected::Valid), + (websocket_extensions_invalid_first, SecWebSocketExtensions, &["/ermessage-deflate"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_extensions_invalid_last, SecWebSocketExtensions, &["permessage-deflat/"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_extensions_window_invalid_last, SecWebSocketExtensions, &["permessage-deflate; client_max_window_bit/"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_extensions_takeover_invalid_last, SecWebSocketExtensions, &["permessage-deflate; server_no_context_takeover; client_no_context_takeove/"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_extensions_numbers, SecWebSocketExtensions, &["permessage-deflate; server_max_window_bits=15; client_max_window_bits=15"], Strict, Expected::Valid), + (websocket_extensions_token_15, SecWebSocketExtensions, &["xxxxxxxxxxxxxxx"], Strict, Expected::Valid), + (websocket_extensions_token_16, SecWebSocketExtensions, &["xxxxxxxxxxxxxxxx"], Strict, Expected::Valid), + (websocket_extensions_token_17, SecWebSocketExtensions, &["xxxxxxxxxxxxxxxxx"], Strict, Expected::Valid), + (websocket_extensions_token_31, SecWebSocketExtensions, &["xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx"], Strict, Expected::Valid), + (websocket_extensions_token_32, SecWebSocketExtensions, &["xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx"], Strict, Expected::Valid), + (websocket_extensions_token_33, SecWebSocketExtensions, &["xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx"], Strict, Expected::Valid), + (websocket_extensions_quoted_name, SecWebSocketExtensions, &["\"x\""], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_extensions_unterminated_name, SecWebSocketExtensions, &["\"x"], Strict, Expectations { + http_owned: Expected::Error(DecodeErrorKind::InvalidToken), + http_borrowed: Expected::Error(DecodeErrorKind::UnterminatedQuote), + raw_owned: Expected::Error(DecodeErrorKind::InvalidToken), + raw_borrowed: Expected::Error(DecodeErrorKind::UnterminatedQuote), + }), + (websocket_extensions_backslash_value, SecWebSocketExtensions, &["x; mode=\\bad"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_extensions_quoted, SecWebSocketExtensions, &["permessage-deflate; client_max_window_bits=\"15\""], Strict, Expected::Valid), + (websocket_extensions_repeated, SecWebSocketExtensions, &["permessage-deflate; client_max_window_bits", "x-custom; mode=fast"], Strict, Expected::Valid), + (websocket_extensions_invalid, SecWebSocketExtensions, &["permessage-deflate; =15"], Strict, Expected::Error(DecodeErrorKind::InvalidToken)), + (websocket_extensions_empty, SecWebSocketExtensions, &[""], Strict, Expectations { + http_owned: Expected::Error(DecodeErrorKind::InvalidSyntax), + http_borrowed: Expected::Error(DecodeErrorKind::MissingValue), + raw_owned: Expected::Error(DecodeErrorKind::InvalidSyntax), + raw_borrowed: Expected::Error(DecodeErrorKind::MissingValue), + }), + (websocket_extensions_absent, SecWebSocketExtensions, &[], Strict, Expected::Absent), + (websocket_version_standard, SecWebSocketVersion, &["13"], Strict, Expected::Valid), + (websocket_version_draft, SecWebSocketVersion, &["8"], Strict, Expected::Valid), + (websocket_version_zero, SecWebSocketVersion, &["0"], Strict, Expected::Valid), + (websocket_version_two_digits, SecWebSocketVersion, &["17"], Strict, Expected::Valid), + (websocket_version_maximum, SecWebSocketVersion, &["255"], Strict, Expected::Valid), + (websocket_version_ows, SecWebSocketVersion, &[" \t13\t "], Strict, Expected::Valid), + (websocket_version_advertisement, SecWebSocketVersion, &["7, 8, 13"], Strict, Expected::Valid), + (websocket_version_repeated, SecWebSocketVersion, &["13", "8", "255"], Strict, Expected::Valid), + (websocket_version_duplicate, SecWebSocketVersion, &["13", "13"], Strict, Expected::Valid), + (websocket_version_late_invalid, SecWebSocketVersion, &["13", "256"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_version_quote, SecWebSocketVersion, &["13", "\"13"], Strict, Expected::Error(DecodeErrorKind::UnterminatedQuote)), + (websocket_version_quote_only, SecWebSocketVersion, &["\"13"], Strict, Expected::Error(DecodeErrorKind::UnterminatedQuote)), + (websocket_version_quoted_member, SecWebSocketVersion, &["\"13\""], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_version_invalid_before_quote, SecWebSocketVersion, &["256, \"13"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_version_quote_after_member, SecWebSocketVersion, &["13, \"13"], Strict, Expected::Error(DecodeErrorKind::UnterminatedQuote)), + (websocket_version_empty, SecWebSocketVersion, &[""], Strict, Expected::Error(DecodeErrorKind::MissingValue)), + (websocket_version_empty_members, SecWebSocketVersion, &[" , , \t"], Strict, Expected::Error(DecodeErrorKind::MissingValue)), + (websocket_version_leading_zero, SecWebSocketVersion, &["013"], Strict, Expected::Error(DecodeErrorKind::InvalidSyntax)), + (websocket_version_relaxed, SecWebSocketVersion, &["13, 8"], Relaxed, Expected::Valid), + (websocket_version_custom_line_limit, SecWebSocketVersion, &["13"; 129], Strict, Expectations { + http_owned: Expected::Valid, + http_borrowed: Expected::Valid, + raw_owned: Expected::Error(DecodeErrorKind::InvalidSyntax), + raw_borrowed: Expected::Error(DecodeErrorKind::InvalidSyntax), + }), + (websocket_version_absent, SecWebSocketVersion, &[], Strict, Expected::Absent), +); diff --git a/crates/http_headers/docs/COMPATIBILITY.md b/crates/http_headers/docs/COMPATIBILITY.md new file mode 100644 index 000000000..5c25ddd30 --- /dev/null +++ b/crates/http_headers/docs/COMPATIBILITY.md @@ -0,0 +1,201 @@ +# Compatibility, wire guarantees, and versioning + +This document states what a caller may rely on across releases: the semver +policy this pre-1.0 crate follows, the minimum supported Rust version, the +feature flags and how they interact, and which parts of the wire behavior +described in [`DESIGN.md`](DESIGN.md) are guaranteed rather than incidental. + +## Semantic versioning + +`http_headers` is `0.y.z`. Cargo therefore treats every `0.(y+1).0` release as +potentially breaking, and this crate uses exactly that latitude: a minor +version bump may add, change, or remove public API, tighten or loosen a +parser's accepted grammar, or change the concrete type returned by an +existing method. A patch release (`0.y.(z+1)`) never changes public API and +never changes what wire input is accepted or what bytes are produced from a +given typed value — it is reserved for bug fixes, performance improvements, +and documentation. + +Once this crate reaches `1.0.0` it will adopt standard semver: no breaking +change without a major version bump, and any acceptance-grammar change that +could reject a previously-accepted value or newly accept a previously-rejected +one will be treated as breaking regardless of whether the public API text +changed. + +There is no commitment yet about when `1.0.0` ships. Track breaking changes +between `0.y` releases in the crate's release notes rather than assuming any +particular stability beyond what is stated above. + +## Minimum supported Rust version (MSRV) + +The workspace declares `rust-version = "1.95"` with the 2024 edition +(`Cargo.toml`). Raising the MSRV is treated as a breaking change under the +policy above: it happens only in a `0.(y+1).0` release, and the new minimum is +stated in that release's notes. A patch release never raises the MSRV. + +## Feature flags + +Declared in `crates/http_headers/Cargo.toml`: + +| Feature | Default | Effect | +|---|---|---| +| No features (`default-features = false`) | — | Core field names, values, errors, `Field`/`SingleValueField`, and source/sink APIs. No built-in header families or optional dependencies are enabled by this crate. | +| `headers-all` | on | Enables all built-in header families listed below, including their optional dependencies. | +| `http` | off | Adds an optional `http::HeaderMap` adapter and conversions to and from the `http` crate's name, value, and method types. The adapter may retain `HeaderValue` storage internally; typed-header semantics do not change. | +| `serde` | off | Enables serialization and deserialization for `FieldName`, `FieldValue`, `EncodedValues`, and owned headers from enabled families (`dep:serde`). Does not enable header families itself. | +| `benchmarking` | off | Exposes private-backend instrumentation (`http_headers_simd/benchmarking`) needed by the workspace benchmarks. Not for downstream use; the surface it exposes is not covered by the semver policy above. | + +To select individual families, disable default features and enable the required +`headers-*` features. Each family below is enabled by `headers-all`. +`BasicCredentials` belongs to `headers-authorization`, not the feature-free core. + +| Family feature | Additional optional dependencies or family features | +|---|---| +| `headers-authorization` | `base64`, `zeroize` | +| `headers-cache-control` | `compact_str` | +| `headers-conditional` | `httpdate`, `headers-etag` | +| `headers-content-length` | None | +| `headers-content-type` | None | +| `headers-cors` | None | +| `headers-etag` | None | +| `headers-location` | `fluent-uri` | +| `headers-negotiation` | `idna` | +| `headers-range` | None | +| `headers-security` | `compact_str` | +| `headers-set-cookie` | None | +| `headers-user-agent` | None | +| `headers-websocket` | `base64`, `sha1` | + +For example, this selects only the two named families and the HTTP adapter: + +```toml +http_headers = { version = "0.1", default-features = false, features = ["headers-content-length", "headers-content-type", "http"] } +``` + +Cargo features are additive: another dependency enabling a family or +`headers-all` can enable it for the same package in the resolved build. + +Enabling `benchmarking` is not a supported way to depend on this crate; the +items it exposes can change or disappear in a patch release. Ordinary header +decoding, encoding, and map operations never require it. + +`http_headers_simd`, the sibling crate holding the SIMD scanning kernels, is +published only so Cargo can resolve `http_headers`. Its rustdoc API is hidden, +and it carries no independent compatibility guarantee. Downstream manifests +should depend on `http_headers`, never on `http_headers_simd` directly. + +## Extension points + +Every public trait in this crate except one is open to downstream +implementation, and none is sealed. Two extensions are supported: teaching the +crate about a new container, and teaching it about a new field. + +| Trait | Implement it to | Required | Provided | +|---|---|---|---| +| `source::FieldSource` | read fields from your container | `lines` | `contains` | +| `sink::FieldSink` | write fields to your container | `set_values`, `append_values`, `remove_values` | `set_encoded`, `append_encoded` | +| `sink::FieldEncodeOutput` | let encoders write straight into your container's storage | `Writer`, `begin_value`, `push_value` | `push_u64` | +| `sink::FieldValueWriter` | receive one value's bytes for the above | `write_bytes`, `finish` | — | +| `sink::FieldEncoder` | define how a value becomes field lines | `encode` | — | +| `SingleValueField` | define a field carried by exactly one field line | `View`, `Owned`, `name`, `decode_view`, `decode_owned`, `as_field_value`, `into_field_value` | `decode_view_with`, `decode_owned_with` | +| `Field` | define a field that may span repeated field lines | `View`, `Owned`, `name`, `view_with`, `owned_with`, `insert` | `view`, `owned`, `remove` | + +`sink::FieldSinkExt` is the exception. Its blanket implementation covers every +`FieldSink`, so no type in any crate can implement it; it exists to be called, +not implemented, and it is listed here only to state that implementing it is +not an available extension. + +A blanket implementation also relates the last two rows: implementing +`SingleValueField` supplies `Field` automatically. Implement `Field` directly +only for a field whose value spans repeated field lines, and never implement +both for one type — the implementations would overlap. + +Implementing `FieldEncodeOutput` and `FieldValueWriter` is optional. A +container that implements only `FieldSink` still accepts every typed field, +because the default `set_encoded` collects into `EncodedValues` first. The two +traits exist so a container can skip that intermediate collection; the bundled +`http::HeaderMap` adapter overrides `set_encoded` for exactly that reason, and +the traits are public so a third-party container can reach the same path. + +The split between required and provided methods above is the part that matters +across releases: adding a required method to any of these traits breaks every +downstream implementation, while adding a provided one does not. New methods +will therefore always arrive with defaults. Under the pre-1.0 policy stated +above this remains a `0.(y+1).0`-only concern either way, but it is the +commitment that will carry into `1.0.0`. + +Note that the reverse direction is not symmetric: because these traits are +unsealed, adding a *blanket* implementation of any of them in a later release +could conflict with a downstream implementation. No such blanket +implementation will be added except where one already exists. + +## Decode modes + +`DecodeMode::Strict` is the default used by `view`, `owned`, direct +`TryFrom` constructors, and Basic credential decoding. +For typed reads from a `FieldSource`, select `DecodeMode::Relaxed` through +`Field::view_with` or `Field::owned_with`. Direct single-value decoding accepts +the mode through `SingleValueField::decode_view_with` and +`SingleValueField::decode_owned_with`; standalone quality parsing accepts it +through `QualityView::parse`. Relaxed decoding is not a raw-value or +skip-validation mode. + +Relaxed decoding recognizes the following interoperability deviations: + +| Headers | Relaxation | Example | +|---|---|---| +| `Accept`, `Accept-Encoding`, `Accept-Language` | Spaces or tabs around the quality `=`, a fractional value without a leading zero, more than three fractional digits below one, or more than three zero digits after one | `gzip; q = .12345` | +| `ETag`, `If-Match`, `If-None-Match` | A lowercase weak-validator prefix | `w/"revision"` | +| `Content-Type` | Optional whitespace around the type/subtype slash | `text / html` | +| `Range`, `Content-Range` | Optional whitespace around range delimiters | `bytes = 0 - 499` | +| `Last-Modified`, `If-Modified-Since`, `If-Unmodified-Since`, `If-Range` | Outer whitespace, the `UTC` zone, or one-digit date/time components in an IMF-style date | ` Sun, 6 Nov 1994 8:49:37 UTC ` | +| `Host` | UTF-8 internationalized registered names that successfully convert through IDNA | `münich.example:443` | +| `Location` | Backslashes treated as slashes for URI-reference validation | `/a\b\c` | + +These are acceptance exceptions, not a general recovery mode. Numeric bounds, +token and list structure, entity-tag contents, range ordering, valid IDNA and +ports, parameter ordering and count, quoting and escaping, list boundaries, +cardinality, framing, authorization, CORS, security, and WebSocket invariants +remain enforced. Relaxed quality values still reject signs, exponents, +multiple decimal points, quoted values, nonzero fractional digits after one, +and empty leading-dot fractions. A decoder never drops an invalid list member +to salvage the rest of a field. Headers absent from the table process +`DecodeMode::Relaxed` exactly as `DecodeMode::Strict`. + +Accepted relaxed values retain their original bytes. Encoding therefore does +not silently canonicalize a relaxed value, and callers inspecting `items()` or +`values()` see the same nonconforming quality syntax that arrived on the wire. + +## Sensitivity inputs + +Sensitivity-setting APIs use `FieldSensitivity::{Sensitive, NonSensitive}` +rather than Boolean arguments. This applies to `FieldValue`, +`FieldValueRef`, `ValueRefsEncoder`, and `FieldEncodeOutput::begin_value`. +The enum's supported path is `http_headers::FieldSensitivity`, alongside the +crate-root field value types that consume it. Code written against the earlier +pre-1.0 draft should replace `set_sensitive(bool)` and `with_sensitive(bool)` +with `set_sensitivity(FieldSensitivity)` and +`with_sensitivity(FieldSensitivity)`; custom encode outputs receive the same +enum directly. The duplicate `http_headers::sink::FieldSensitivity` alias was +removed before 1.0 so the type has one documented public path. + +## Wire guarantees + +The behaviors [`DESIGN.md`](DESIGN.md#protocol-behavior) documents per header +family are part of the semver-covered contract above: a patch release will +not change which bytes a constructor accepts, which bytes a decoder accepts +or rejects, or which bytes an encoder produces for a value that already +round-trips today. Two guarantees apply across every header type: + +- **Byte-valued, not string-valued.** Field values that are not guaranteed + UTF-8 by their grammar are exposed as `&[u8]`; an explicit `*_str` accessor + performs the fallible UTF-8 conversion. A decoder never panics on bytes a + `FieldValue` can hold, including `obs-text` (`0x80..=0xff`). +- **Field-line boundaries are preserved.** A header delivered as several + repeated field lines is never silently flattened into one comma-joined + value before the caller sees it; `FieldLines::repeated()` and the borrowed + iterators reconstruct the same boundaries that were parsed. + +A change to either of the two guarantees above, or to any bullet in +`DESIGN.md`'s protocol-behavior list, is a breaking change under the semver +policy stated above regardless of which release channel introduces it. diff --git a/crates/http_headers/docs/DESIGN.md b/crates/http_headers/docs/DESIGN.md new file mode 100644 index 000000000..7bd4d0d16 --- /dev/null +++ b/crates/http_headers/docs/DESIGN.md @@ -0,0 +1,388 @@ +# Design + +`http_headers` is a typed header layer with its own core abstractions. Its +central design choice is to preserve HTTP wire structure while letting callers +choose between borrowed parsing, owned values, and reusable Basic credential +storage. + +## Core abstractions and container adapters + +The core owns these abstractions and depends on no external header crate: + +- `FieldValue` stores copied or adopted owned wire bytes inline through 64 + bytes and uses shared storage for longer values. `from_static` and + `from_shared` retain their backing storage even for short values. + `FieldValueRef<'a>` is the borrowed counterpart every decoder observes. +- `FieldName` is the single map-key, wire-name, and typed-field name type. It + has one variant for every recognized standard header and `Custom(Arc)` + for other names. Runtime parsing normalizes custom names to lowercase. + The source/sink boundary uses `&'static FieldName`, so known variants and + static custom descriptors are passed without cloning their owned representation. + Custom descriptors can use `static LazyLock`. Runtime-owned names + support validation, comparison, and conversion to container-specific names; + locally constructed names cannot be used for source/sink lookup, insertion, + append, removal, or `FieldLines` construction. Dynamic name operations use + the container's native API instead. +- `source::FieldSource` supplies the `source::FieldLines` stored for a + typed field, and `sink::FieldSink` stores them. Each generic operation + obtains its key from `Field::name`: a container that indexes well-known + headers reads `FieldName::index` and answers without hashing or comparing + anything, while any other container compares the name's bytes. +- `FieldSource` carries the typed read API, while `FieldSink` carries the + typed write API. Typed methods look a field up by `Field::name` alone, and + because that name is a compile-time constant the indexed path folds away + during monomorphization. Nothing on a decoding path is a trait object. + +Container integrations sit outside these abstractions. The optional `http` +feature implements both traits for `http::HeaderMap` and supplies conversions +at that boundary. Values retained from a `HeaderMap` may use an internal +adapter-specific representation to share its storage; this does not alter the +public `FieldValue` model or typed-header behavior. Other integrations can +implement the same source and sink contracts without depending on `http`; +disabling the feature removes that dependency without removing any typed +header. + +## Owned headers and borrowed views + +`MethodView` and `FieldNameView` are shared by the CORS and negotiation +families and are available when either family is enabled. Method identity is +case-sensitive; field-name equality and hashing are ASCII-case-insensitive. +Both views preserve the original wire spelling. + +Every `Field` definition associates a generic `View<'a>` and an `Owned` +representation. `view` returns the view, which can borrow +`FieldValueRef` storage directly from the header source. +`owned` invokes the implementation's owned path directly. + +This split avoids forcing a clone for read-only access while still giving +callers a type that can outlive the map. Some headers, such as +`ContentLength`, are small values and use the same copyable type for both +forms. + +`Field::insert` lets each implementation move its validated storage into the +destination without exposing an intermediate encoding operation. + +## FieldLines and encoded values + +`FieldLines<'a>` borrows the field lines a field source already stores — a +slice of `FieldValue`s, borrowed `FieldValueRef`s into arbitrary transport +storage, one parsed line, or the field lines of an adapted +`http::HeaderMap` — rather than collecting or joining them. The set of +representations is closed, so iteration dispatches statically. It therefore +retains each field-line boundary and insertion order. Parsers can require +exactly one value, reiterate repeated values, or scan comma- and +semicolon-delimited items across all lines. Delimited scanning trims optional +whitespace and respects quoted delimiters within each physical line. An item +cannot span lines because iterator items borrow one contiguous source line. + +`FieldSource::lines` returns `None` for absence and +`Some(FieldLines)` for presence. `FieldLines` therefore always contains +at least one raw field line; a zero-length line remains present and is decoded +according to the typed header's grammar rather than being confused with +absence. Slice, borrowed, and `http::HeaderMap` constructors return `None` when +given no lines. + +Custom sources are bounded at the shared decode boundary: one field name may +contribute at most 64 KiB across at most 128 field lines, and one delimited +decode may yield at most 1,024 list items. Owned decoding checks the byte and +line budgets before copying source data, while `DelimitedItems` applies the +item budget before yielding another item. Specialized list parsers invoke the +same `FieldLines` item-budget check before their optimized scans. Opaque +single-value parsers do not interpret delimiters as list members. The +`http::HeaderMap` representation is exempt because that adapter exposes +already validated values and is not a custom `FieldSource`. + +Validated `FieldValue` slices supplied by custom sources remain subject to the +same budgets. A budget rejection is `DecodeErrorKind::SourceLimitExceeded`, +not `InvalidSyntax`; valid HTTP grammar can exceed an admission limit. Typed +negotiation constructors use the same category for their byte and item bounds. +Typed Serde deserialization also applies byte, line, and applicable list-item +budgets during collection and maps admission failures into Serde errors with +`source limit exceeded` diagnostics. Raw `FieldValue` and `EncodedValues` +deserialization does not impose these typed-source budgets. + +For `FieldLines`, checks retain their scan order: field-line count is checked first, then +aggregate byte count and raw-byte validity in physical line order. Thus an +oversized line wins over invalid bytes in that same line, but invalid bytes +in an earlier line can win over a later byte overflow. List preflight may +reject an over-budget item before its grammar is parsed; streaming +`DelimitedItems` instead diagnoses an unterminated quote before admitting +that item. `InvalidSyntax` and the more specific grammar errors remain +reserved for malformed input. + +This distinction is required for `Set-Cookie`, where each cookie remains a +separate field line. List-valued headers can still treat multiple lines as one +logical list without first allocating a flattened string. + +`EncodedValues` mirrors that model on output. A singleton is stored directly +in an `Option` with no `Vec` allocation. Repeated values can be +appended, while `from_vec` adopts an owned vector without reallocating it. + +## Fallible, atomic map operations + +`FieldSource` provides raw container lookup primitives. `FieldSink` adds +`set_values`, `append_values`, and `remove_values` as the raw mutation +primitives. Requiring append separately lets the provided `append_encoded` +operation encode its complete delta before one atomic, delta-proportional +mutation. `Field` provides strict `view` and `owned` conveniences; +implementations supply `view_with`, `owned_with`, and `insert`. Its provided +`remove` operation deletes matching field lines without parsing them. + +`Field::insert` returns `Result<(), InsertError>`. It encodes before mutation +and reserves or obtains the destination entry before replacing an existing +value. If the container has reached its maximum capacity, insertion reports the +error and leaves the existing field lines unchanged. An empty encoding +removes the existing header. + +`Field::remove` performs no decoding and therefore cannot fail because a +stored value is malformed. + +## Errors and secret-safe diagnostics + +`DecodeError` records a copyable static reference to `FieldName`, a non-exhaustive +`DecodeErrorKind`, and an optional zero-based field-value index. It +deliberately does not retain or print raw field bytes because headers may +contain credentials, cookies, signed URLs, or other secrets. + +`InsertError` is also `Copy`. Its non-exhaustive `InsertErrorKind` distinguishes +invalid values, encoder length-contract violations, buffer reservation failures, +and value/container capacity limits. It retains neither field contents nor an +underlying error source. These categories describe failures, not retry policy: +the built-in sinks reject lengths above `isize::MAX` as `CapacityExceeded` +before reserving, while later buffer reservation failures are `AllocationFailed`. +Neither category promises that retrying without other changes will help. +No automatic recovery classification or `recoverable` dependency is imposed. + +Sensitive built-in types set `FieldValue::set_sensitive` where appropriate. +Their `Debug` implementations, along with those for `FieldLines`, +`BasicCredentials`, and encoded collections, report structure or counts rather +than contents. Error handling remains structured without making secret data +part of routine logs. + +## Reusable Basic credentials + +`BasicCredentials` owns decoded Basic authorization bytes and reuses its +allocation across requests. `AuthorizationOwned::extract` and +`AuthorizationView<'_, Basic>::extract` return a reference to that storage, +preventing reuse while the username and password are borrowed. Previous +credentials are zeroized before extraction, when cleared, and on drop. +`BasicCredentials::clear` retains capacity up to a caller-configurable limit +and drops an oversized allocation. The default retention cap is 64 KiB. + +## Downstream extension contract + +Downstream crates can implement `Field` directly when they need repeated +values, custom ownership, caching, or a nonstandard encoding. Implementations +must: + +1. return a `&'static FieldName`; custom headers can initialize one + `FieldName::Custom` lazily with `std::sync::LazyLock`; +2. validate grammar and cardinality when decoding; +3. return views that borrow only for their declared lifetime; +4. produce only valid `FieldValue` objects when encoding; and +5. preserve any security invariant, including sensitivity markings and + secret-safe `Debug`. + +`SingleValueField` is the preferred adapter for one-field-value types. It +supplies exact singleton cardinality and move-based encoding; the downstream +`view` function validates inbound values, while constructors must establish +the same invariant for owned values. + +Built-in headers validate through a private `validate` module that wraps safe, +shared primitives for token and field-value validation, token-byte checks, +ASCII case-insensitive comparison, and scanning for list/quote/whitespace +bytes. It is an internal detail built on the rustdoc-hidden +`http_headers_simd` companion crate. That crate is published only as an +implementation dependency and is not a supported downstream API. Custom +headers implement their own validation instead of depending on it. + +## Protocol behavior + +The parsers preserve byte-level RFC behavior instead of treating every field +as a Unicode string. + +- **Cache-Control:** recipients ignore empty comma-list members, including an + entirely empty logical list. Sender constructors and the builder require at + least one nonempty valid directive. Directive values are exposed as bytes; + `value_str` is an explicit fallible UTF-8 conversion. +- **Quoted strings and `obs-text`:** field-value validation accepts bytes + `0x80..=0xff`, and Cache-Control and Content-Type quoted strings retain them. + APIs return raw bytes where UTF-8 is not guaranteed. +- **Content-Type:** empty parameter slots such as trailing `;` or repeated + semicolons are tolerated. Parameter names and values must have no optional + whitespace around `=`. +- **Location:** values are RFC 3986 URI-references, including absolute, + relative, fragment-only, and empty references. Validation is delegated to + `fluent-uri` with IPvFuture support rather than to an HTTP URI model with + different restrictions. +- **Content-Length:** repeated field lines and comma-separated duplicates are + accepted only when every decimal value agrees. Conflicts, invalid digits, + and `u64` overflow are errors. +- **Set-Cookie:** repeated field lines are retained and encoded independently; + they are never parsed as a comma list. +- **Negotiation:** media ranges, codings, languages, methods, and field names + remain allocation-free list iterators. `Host` separates host and optional + port without normalizing the authority. +- **Conditional requests and dates:** entity-tag lists preserve wildcard and + weak/strong semantics. HTTP dates accept the three recipient date formats; + semantic constructors emit IMF-fixdate. `If-Range` accepts only a strong + entity tag or an HTTP date. +- **Ranges:** byte-range and content-range components are parsed without + flattening field lines. Unknown syntactically valid range units remain + visible rather than being rewritten as bytes. +- **CORS:** origins, methods, and field-name lists are syntax types, not + authorization policy. In particular, wildcard and credential combinations + remain the application's policy decision. +- **WebSocket:** key and accept values require canonical base64 lengths; + `SecWebSocketAcceptOwned::from_key` computes the RFC handshake digest. Extension + parameters preserve quoted wire data and reject whitespace around `=`. +- **Security fields:** HSTS is structured and requires exactly one `max-age`; + CSP deliberately remains an opaque safe field value; Referrer-Policy keeps + unknown future tokens while selecting the last recognized fallback. + +## Inline storage + +`FieldValue` copies byte slices and adopts owned strings or vectors into +inline storage when they contain at most 64 bytes; longer values use shared +`bytes::Bytes` storage. `from_static` and `from_shared` instead retain their +static or shared backing storage regardless of length. The optional `http` +adapter copies short values inline and can retain longer source +`HeaderValue`s without copying their private buffers. +Sensitivity is stored beside every representation and does not affect +equality, ordering, or hashing. + +`SmallVec` removes common temporary allocations in the generic encoding sink, +the `http::HeaderMap` adapter, and Cache-Control storage. `CompactString` +stores extension directives in the Cache-Control and HSTS builders. These +choices are covered by the `storage` metabench target and by unit tests at +their inline/spill boundaries; they are implementation details rather than +layout guarantees. + +## SIMD dispatch and unsafe isolation + +The token, token68, field-value, and interesting-byte scanners use scalar Rust +below 16 bytes on x86 and x86-64, and below 32 bytes on other architectures. +Token-list, Base64, and URI scanners use a 16-byte threshold on every +architecture; ASCII case-insensitive equality uses 32 bytes. At or above each +scanner's threshold, dispatch selects: + +- SSE2 on x86-64, where it is part of the architecture baseline, with an + SSE4.2 `pcmpestrm` fast path for token68 when runtime detection succeeds; +- SSE2 on x86 when runtime detection (or the compile-time target feature + without `std`) confirms support, again preferring SSE4.2 for token68; +- NEON on AArch64 when runtime detection (or the compile-time target feature + without `std`) confirms support; or +- the scalar implementation on all other targets and whenever the optimized + backend is unavailable. + +Vector loops process 16-byte lanes and pass tails to the scalar reference +implementation. Differential properties check backend equivalence, including +lengths around vector and dispatch boundaries. + +The user-facing `http_headers` crate declares `#![forbid(unsafe_code)]`. +Project-owned pointer loads, intrinsics, and `target_feature` functions live +only in the rustdoc-hidden `http_headers_simd` crate, which denies unsafe +operations inside unsafe functions unless they are explicit. Its outward API +is safe, so target-feature preconditions and pointer bounds cannot leak into +`http_headers`. + +## Deliberate API decisions + +The following surfaces are deliberate rather than oversights. + +**Duration-based max ages require whole seconds.** `AccessControlMaxAgeOwned` +offers fallible duration construction and `TryFrom`, not a lossy +`From`. `CacheControlBuilder` and `StrictTransportSecurityBuilder` +retain the supplied duration and reject fractional seconds at build time; +direct Cache-Control encoding and the CORS duration setter reject them before +changing the sink. Constructors report `DecodeErrorKind::InvalidNumber`; +insertion reports `InsertErrorKind::InvalidValue`. No path silently floors a +caller-supplied max age. + +**`SetCookieOwned` stays construction-fallible; no `FromIterator`/`Extend`.** +`SetCookieOwned::push` and `push_str` validate every field value +against `SET_COOKIE`'s nonempty-field-value grammar and reject one that +fails. `EncodedValues`' +`FromIterator`/`Extend` adopt trusted, already-encoded `FieldValue` storage, +not attacker-reachable cookie content; a `FromIterator`/ +`Extend` impl for `SetCookieOwned` would need to either silently drop invalid +values or +panic (turning untrusted request- or response-adjacent data into a crash +surface), and neither is acceptable. + +Immutable iteration preserves validity, but `iter_mut` and +`IntoIterator for &mut SetCookieOwned` can replace an entry with an empty +`FieldValue`: valid raw storage, but not a valid cookie field line. +`SetCookie::insert` and `FieldSinkExt::append_set_cookie` therefore revalidate +nonemptiness before changing the sink and restore cookie sensitivity. +An empty collection remains valid; it is an empty stored entry that fails. + +**Built-in descriptors forward essential `Field` operations.** Every built-in +descriptor exposes inherent `view`, `owned`, `insert`, and `remove` methods so +callers can discover and invoke its essential operations without importing the +`Field` trait. These methods are generated from the same built-in field list +and delegate directly to `Field`, keeping their behavior and signatures tied +to the trait implementation instead of maintaining hand-written duplicates. +`Field` remains the generic contract and the extension point for downstream +field descriptors. + +**Type names follow RFC 9110; the crate name follows common usage.** RFC 9110 +calls the genus a *field*: a field has a field name and a field value, and a +field that appears in the header section is a *header field* (Section 6.3), +which the specification itself notes is called a "header" only colloquially. +The types therefore say `FieldName`, `FieldValue`, `FieldSource`, and +`FieldSink`, because none of them is specific to the header section — the same +`FieldSource` is correct over a trailer section. The crate stays +`http_headers` because a crate name is a domain label in the reader's +vocabulary rather than a node in the specification's taxonomy, and because +`FieldName` no longer collides with `http::HeaderName` at any import site. +Headers of the header section keep the word "header" in prose and in the +`headers` module, where it is accurate. + +**Types are grouped by role, and each has exactly one path.** Concrete +headers, their views, builders, scheme markers, and related enums live under +`http_headers::headers`. The read-path extension points a custom container +implements to supply field values live under `http_headers::source`, and the +write-path ones it implements to store them live under `http_headers::sink`. +The crate root holds only the vocabulary every user names: field values, +field names, decode errors and modes, and the `Field` traits. No type is +re-exported from two paths, so there is one way to import each name and +`cargo doc` lists it once. + +**All header families by default, individually selectable when needed.** +The default `headers-all` feature enables every built-in header family. +Disabling default features leaves the core field names, values, errors, +traits, and source/sink APIs; consumers then enable only the `headers-*` +families they need. Family features activate their optional dependencies: +for example, `headers-location` enables `fluent-uri`, while +`headers-authorization` enables `base64` and `zeroize`. +`headers-conditional` also enables `headers-etag`. + +The `http` and `serde` features add integrations independently of family +selection. Serde supports owned headers from whichever families are enabled; +neither integration enables all families. A consumer selecting only +`Content-Type` and `Content-Length` therefore need not enable the optional +dependencies of unrelated families. Cargo features remain additive, so +another dependent can enable more families for the same resolved package. +[`COMPATIBILITY.md`](COMPATIBILITY.md#feature-flags) lists the complete feature +and dependency mapping. + +**Opaque owned headers retain `FieldValue`.** Their owned forms clone and hold +the validated `FieldValue` rather than introducing a second opaque-byte +representation. This preserves sensitivity metadata, inline/shared storage, +and round-trip behavior. Callers that only inspect a value should use `view` +to borrow source storage instead. + +**Accessors keep re-validating UTF-8; views keep returning `Result`.** +`LocationView::as_str` and the equivalent per-item CORS, +`AcceptRanges`, and security-header accessors keep re-running +`str::from_utf8` over bytes the decoder already proved to be ASCII, rather +than carrying a pre-validated `&str` in the view and dropping the `Result` +those accessors return. Storing `&str` instead of `&[u8]` in a view is an +accessor-signature change for every caller matching on today's +`Result<&str, Utf8Error>` return; per this item's own escape clause, the +`from_utf8` pass is a validation over an already-short, already-ASCII slice +next to work — grammar validation, quote/whitespace scanning — that has +already walked the same bytes once, so it is not the dominant cache-and-clone +cost. Kept as-is; accessor signatures and their `Result` returns are +preserved. diff --git a/crates/http_headers/docs/PERF.md b/crates/http_headers/docs/PERF.md new file mode 100644 index 000000000..36f356ac4 --- /dev/null +++ b/crates/http_headers/docs/PERF.md @@ -0,0 +1,60 @@ +# Performance + +This table is the comparative snapshot imported with `http_headers` 0.1.0. It +is not a promise of current timing on different hardware or dependency +versions; use the repository's metabench targets to measure the current +checkout. + +Each cell reports metabench's median wall-clock time, Callgrind instruction +count, allocation count, and total allocated bytes for one typed decode and +read. Benchmarks consume the decoded value and force comparable semantic work +when `headers 0.4.1` defers parsing. Setup uses a prebuilt `HeaderMap` outside +the measured operation. Criterion uses a 1 s warm-up, 3 s measurement period, +and 60 samples per arm. Each cell is formatted as time, instructions, then +allocation count / allocated bytes. `n/a` means `headers 0.4.1` does not +provide that typed header. + +| Header | `headers 0.4.1` | `http_headers` (owned) | `http_headers` (borrowed) | +|---|---:|---:|---:| +| Accept | n/a | 399.0 ns
1,387 instr
1 allocs / 24 B | 343.5 ns
1,055 instr
0 allocs / 0 B | +| Accept-Encoding | n/a | 119.1 ns
679 instr
0 allocs / 0 B | 74.0 ns
449 instr
0 allocs / 0 B | +| Accept-Language | n/a | 148.1 ns
731 instr
0 allocs / 0 B | 104.7 ns
502 instr
0 allocs / 0 B | +| Accept-Ranges | 37.5 ns
590 instr
1 allocs / 24 B | 70.9 ns
408 instr
0 allocs / 0 B | 57.2 ns
369 instr
0 allocs / 0 B | +| Access-Control-Allow-Credentials | 20.3 ns
224 instr
0 allocs / 0 B | 28.1 ns
227 instr
0 allocs / 0 B | 23.9 ns
227 instr
0 allocs / 0 B | +| Access-Control-Allow-Headers | 195.6 ns
2,021 instr
2 allocs / 36 B | 149.9 ns
914 instr
0 allocs / 0 B | 82.8 ns
791 instr
0 allocs / 0 B | +| Access-Control-Allow-Methods | 85.4 ns
1,185 instr
1 allocs / 24 B | 129.5 ns
725 instr
0 allocs / 0 B | 88.5 ns
529 instr
0 allocs / 0 B | +| Access-Control-Allow-Origin | 192.1 ns
1,335 instr
2 allocs / 43 B | 136.5 ns
694 instr
0 allocs / 0 B | 54.6 ns
583 instr
0 allocs / 0 B | +| Access-Control-Expose-Headers | 133.6 ns
1,816 instr
2 allocs / 36 B | 114.7 ns
827 instr
0 allocs / 0 B | 93.1 ns
704 instr
0 allocs / 0 B | +| Access-Control-Max-Age | 27.8 ns
325 instr
0 allocs / 0 B | 24.0 ns
291 instr
0 allocs / 0 B | 29.4 ns
291 instr
0 allocs / 0 B | +| Access-Control-Request-Headers | 159.0 ns
1,982 instr
2 allocs / 36 B | 122.1 ns
914 instr
0 allocs / 0 B | 68.8 ns
791 instr
0 allocs / 0 B | +| Access-Control-Request-Method | 22.6 ns
269 instr
0 allocs / 0 B | 28.3 ns
270 instr
0 allocs / 0 B | 21.1 ns
255 instr
0 allocs / 0 B | +| Allow | 76.1 ns
1,190 instr
1 allocs / 24 B | 110.5 ns
645 instr
0 allocs / 0 B | 107.0 ns
440 instr
0 allocs / 0 B | +| Authorization (Basic) | 148.6 ns
995 instr
1 allocs / 18 B | 160.5 ns
654 instr
0 allocs / 0 B | 112.1 ns
539 instr
0 allocs / 0 B | +| Authorization (Bearer) | 84.4 ns
546 instr
1 allocs / 24 B | 211.7 ns
702 instr
1 allocs / 24 B | 91.3 ns
489 instr
0 allocs / 0 B | +| Cache-Control | 261.2 ns
1,338 instr
0 allocs / 0 B | 214.8 ns
1,060 instr
0 allocs / 0 B | 131.5 ns
838 instr
0 allocs / 0 B | +| Content-Length | 56.5 ns
367 instr
0 allocs / 0 B | 47.1 ns
333 instr
0 allocs / 0 B | 50.0 ns
333 instr
0 allocs / 0 B | +| Content-Range | 142.7 ns
1,019 instr
0 allocs / 0 B | 140.1 ns
741 instr
0 allocs / 0 B | 99.1 ns
619 instr
0 allocs / 0 B | +| Content-Security-Policy | n/a | 96.1 ns
576 instr
1 allocs / 24 B | 21.8 ns
223 instr
0 allocs / 0 B | +| Content-Type | 199.6 ns
1,606 instr
1 allocs / 31 B | 118.2 ns
438 instr
0 allocs / 0 B | 42.9 ns
283 instr
0 allocs / 0 B | +| ETag | 58.9 ns
554 instr
1 allocs / 24 B | 116.0 ns
472 instr
0 allocs / 0 B | 52.2 ns
346 instr
0 allocs / 0 B | +| Host | 104.3 ns
1,021 instr
2 allocs / 40 B | 141.1 ns
708 instr
0 allocs / 0 B | 106.4 ns
683 instr
0 allocs / 0 B | +| If-Match | 148.6 ns
1,664 instr
2 allocs / 49 B | 166.2 ns
890 instr
0 allocs / 0 B | 115.3 ns
707 instr
0 allocs / 0 B | +| If-Modified-Since | 151.1 ns
1,007 instr
0 allocs / 0 B | 122.0 ns
620 instr
0 allocs / 0 B | 74.0 ns
523 instr
0 allocs / 0 B | +| If-None-Match | 155.0 ns
1,690 instr
2 allocs / 49 B | 119.1 ns
867 instr
0 allocs / 0 B | 87.0 ns
713 instr
0 allocs / 0 B | +| If-Range | 53.7 ns
567 instr
1 allocs / 24 B | 89.8 ns
452 instr
0 allocs / 0 B | 45.9 ns
354 instr
0 allocs / 0 B | +| If-Unmodified-Since | 113.6 ns
1,007 instr
0 allocs / 0 B | 112.4 ns
620 instr
0 allocs / 0 B | 66.3 ns
523 instr
0 allocs / 0 B | +| Last-Modified | 124.7 ns
1,007 instr
0 allocs / 0 B | 104.2 ns
620 instr
0 allocs / 0 B | 60.5 ns
523 instr
0 allocs / 0 B | +| Location | 42.2 ns
352 instr
1 allocs / 24 B | 172.6 ns
1,029 instr
1 allocs / 24 B | 97.0 ns
786 instr
0 allocs / 0 B | +| Range | 42.8 ns
553 instr
1 allocs / 24 B | 94.7 ns
492 instr
0 allocs / 0 B | 46.6 ns
419 instr
0 allocs / 0 B | +| Referrer-Policy | 68.5 ns
839 instr
0 allocs / 0 B | 110.5 ns
556 instr
0 allocs / 0 B | 35.5 ns
301 instr
0 allocs / 0 B | +| Sec-WebSocket-Accept | 35.3 ns
352 instr
1 allocs / 24 B | 103.6 ns
423 instr
0 allocs / 0 B | 33.5 ns
301 instr
0 allocs / 0 B | +| Sec-WebSocket-Extensions | n/a | 109.9 ns
575 instr
0 allocs / 0 B | 38.4 ns
311 instr
0 allocs / 0 B | +| Sec-WebSocket-Key | 36.3 ns
488 instr
1 allocs / 24 B | 99.8 ns
423 instr
0 allocs / 0 B | 36.1 ns
301 instr
0 allocs / 0 B | +| Sec-WebSocket-Protocol | n/a | 93.7 ns
537 instr
0 allocs / 0 B | 35.1 ns
300 instr
0 allocs / 0 B | +| Sec-WebSocket-Version | 19.6 ns
226 instr
0 allocs / 0 B | 23.5 ns
234 instr
0 allocs / 0 B | 27.0 ns
234 instr
0 allocs / 0 B | +| Server | 47.0 ns
627 instr
1 allocs / 24 B | 99.2 ns
418 instr
0 allocs / 0 B | 30.7 ns
291 instr
0 allocs / 0 B | +| Set-Cookie | 44.9 ns
666 instr
2 allocs / 184 B | 113.9 ns
657 instr
1 allocs / 24 B | 45.1 ns
308 instr
0 allocs / 0 B | +| Strict-Transport-Security | 99.1 ns
1,434 instr
0 allocs / 0 B | 91.9 ns
436 instr
0 allocs / 0 B | 46.5 ns
380 instr
0 allocs / 0 B | +| User-Agent | 49.1 ns
610 instr
1 allocs / 24 B | 105.2 ns
526 instr
1 allocs / 24 B | 29.3 ns
291 instr
0 allocs / 0 B | +| Vary | 76.1 ns
1,370 instr
1 allocs / 24 B | 106.2 ns
737 instr
0 allocs / 0 B | 43.5 ns
546 instr
0 allocs / 0 B | +| X-Content-Type-Options | n/a | 73.2 ns
403 instr
0 allocs / 0 B | 33.3 ns
295 instr
0 allocs / 0 B | diff --git a/crates/http_headers/docs/TODO.md b/crates/http_headers/docs/TODO.md new file mode 100644 index 000000000..8579331a6 --- /dev/null +++ b/crates/http_headers/docs/TODO.md @@ -0,0 +1,1417 @@ +# TODO + +This file tracks outstanding work for the `http_headers` family. Completed +items are deleted rather than retained as history. + +## Contents + +### Security + +- [SEC2](#sec2) — Pin code executed with the GitHub release token +- [SEC3](#sec3) — Preserve sensitivity across typed reconstruction and forwarding + +### Conformance + +- [CON1](#con1) — Make checked byte-to-text conversions explicit +- [CON3](#con3) — Establish current unsafe-code verification evidence + +### Performance + +- [P1](#p1) — Accelerate Cache-Control fallback token validation +- [P6](#p6) — Reduce HTTP-date digit-validation overhead +- [P9](#p9) — Reuse accelerated validation and searches for entity tags +- [P10](#p10) — Evaluate Content-Type common-literal dispatch +- [P11](#p11) — Resume Content-Type lookup after the cached parameter prefix +- [P12](#p12) — Reuse delimiter searches in negotiation validation +- [P13](#p13) — Reduce Host literal and international-fallback overhead +- [P14](#p14) — Investigate Authorization validation data flow +- [P15](#p15) — Reduce repeated CORS list traversal and projection +- [P18](#p18) — Tighten CORS scalar token decoding +- [P19](#p19) — Investigate Content-Length singleton and decimal decoding +- [P21](#p21) — Reduce shared source and result construction overhead +- [P24](#p24) — Reduce retained URI and authority metadata costs +- [P25](#p25) — Compare streaming negotiation with bounded retained member metadata +- [P26](#p26) — Evaluate inline general Content-Type metadata +- [P27](#p27) — Transfer completed Allow and Vary construction buffers + +### Features + +- [F1](#f1) — Classify semantically sensitive typed headers with `data_privacy_core` +- [F2](#f2) — Add byte-oriented redaction before redacting typed headers +- [F3](#f3) — Convert `templated_uri::Uri` into `LocationOwned` + +### Documentation + +- [D1](#d1) — Document typed headers in the workspace HTTP APIs + +### Benchmarks + +- [B1](#b1) — Cover semantic reads separately from decode-and-drop +- [B2](#b2) — Measure source representations and ownership boundaries +- [B3](#b3) — Cover nonliteral grammar and fallback distributions +- [B4](#b4) — Gate parser instruction and allocation regressions +- [B5](#b5) — Measure downstream release code generation and layouts +- [B6](#b6) — Establish workload shares and benchmark repeatability + +### Testing + +- [T1](#t1) — Make counter-lock rendezvous tolerate spurious failure and abort safely + +## Security + + +### SEC2 — Pin code executed with the GitHub release token + +**Area:** shared GitHub release build boundary used by the `http_headers` family · **Priority:** Medium · **Effort:** Medium + +**Severity:** Medium · **Confidence:** High · **Kind:** Hardening · **Threat:** compromise of an upstream dependency or tool release selected by a subsequent trusted release run · **Scope:** 1 credential-bearing release job with 2 mutable code-selection paths, exhaustive for this workflow + +Make the code receiving the release credential a reviewed, reproducible +selection. The contents-write workflow exports `GH_TOKEN` while invoking +`just`, which launches `cargo -Zscript`. The script declares `argh = "0.1"` +without a repository-controlled script lock at this invocation, so a fresh +runner can select dependency updates not reviewed with the release commit. +Their build scripts or procedural macros inherit the token before the release +program runs. The setup action also selects any `just >=1.46.0`; that selected +executable later runs directly in the credential-bearing step. + +- `.github/workflows/publish-gh-release.yml:11` — contents-write token scope +- `.github/workflows/publish-gh-release.yml:39` — token-bearing invocation +- `justfile:54` — recipe compiles and executes the Rust script +- `scripts/publish-gh-release.rs:12` — script dependency selection +- `.github/actions/anvil-setup/action.yml:115` — minimum-version rather than exact `just` selection + +An upstream compromise could consequently modify repository contents or +GitHub Releases using the job token. This is conditional supply-chain +hardening, not evidence of compromise, fork-PR escalation, or exposure of a +crates.io publishing credential. Keep compilation/bootstrap outside the +credential-bearing step where possible, and ensure any code that still +receives the token uses reviewed dependency/tool identities. Change generated +Anvil setup through its generator/configuration rather than hand-editing it. + +**Done when:** a clean release runner cannot select unreviewed tool or script +dependency versions, dependency compilation does not inherit the release +token, and an isolated workflow regression check with a dummy token fails for +the current mutable selections/inherited build environment without creating +a real release or accessing a real credential. + +--- + + +### SEC3 — Preserve sensitivity across typed reconstruction and forwarding + +**Area:** `http_headers` normalization, compact storage and borrowed forwarding · **Priority:** High · **Effort:** Small + +**Severity:** Medium · **Confidence:** Medium · **Kind:** Hardening · **Threat:** a reader of lower-trust diagnostics after an application reconstructs or forwards explicitly sensitive typed values · **Scope:** 5 typed header families, exhaustive for the normalization, compact-storage and borrowed-forwarding paths cited below + +Carry sensitivity through normalization rather than reconstructing only the +wire bytes. Cache-Control and Accept-Ranges feed bare member slices into +`normalized_comma_value`, which creates a fresh nonsensitive `FieldValue`. +Their `Field::insert` implementations then send that value to the sink, even +when retained input lines were marked sensitive. The HTTP adapter faithfully +copies the already-cleared flag, so normal downstream Debug output can reveal +marked extension values. Accept-Ranges additionally discards the marker when +its owned or borrowed decoder replaces canonical `bytes`/`none` storage with +the compact representation. + +- `crates/http_headers/src/headers/cache_control.rs:182` — normalization drops the retained lines' metadata; insertion uses it at line 450 +- `crates/http_headers/src/headers/range/accept_ranges.rs:138` — second normalization caller; insertion uses it at line 374 +- `crates/http_headers/src/headers/range/accept_ranges.rs:379` — borrowed and owned canonical shortcuts discard original storage; direct construction does likewise at line 435 +- `crates/http_headers/src/headers/shared.rs:364` — helper accepts only byte slices and constructs a fresh value +- `crates/http_headers/src/field_value.rs:624` — byte-vector conversion calls the nonsensitive constructor at line 329 +- `crates/http_headers/src/http_adapter.rs:189` — insertion forwards the value through the marker-preserving HTTP conversion + +Preserve the marker when any contributing line is sensitive, including mixed +repeated lines and compact canonical representations. Keep existing wire +normalization, grammar validation, and nonsensitive behavior; no new privacy +dependency is needed. + +**How to disprove:** show that the consuming application never marks these +headers sensitive or never exposes the resulting sink to lower-trust +diagnostics. No deployed disclosure is established by this static path. + +**Done when:** custom and HTTP sink round-trip tests fail on the current code +and demonstrate retained sensitivity for single and mixed repeated lines, +owned insertion, and borrowed/owned canonical Accept-Ranges forwarding; +downstream Debug omits marked extension content and nonsensitive cases retain +their existing behavior. + +**See also:** P18, P21, P24, P25 +(preserve this marker before evaluating compact storage or forwarding); +F1, F2 (future privacy features, not preservation of this marker). + +The same metadata-reconstruction mechanism also affects three additional +typed forwarding surfaces, exhaustively identified in these paths: + +- `crates/http_headers/src/headers/cors/access_control_request_method.rs:85` + — the registered-method shortcut retains only a static method name; + `into_field_value` at line 95 creates `FieldValue::from_static(method)` + without the original marker. Direct `TryFrom` repeats the + shortcut at line 397, whereas extension methods retain their input value. +- `crates/http_headers/src/sink/field_sink_ext.rs:455` — borrowed conditional + insertion reconstructs a wildcard with `FieldValueRef::new(b"*")`, or tags + with `FieldValueRef::new(tag.as_bytes())` at line 465. The macro instantiates + both `IfMatchView` and `IfNoneMatchView` at lines 477–479; neither path + transfers the contributing source lines' sensitivity. + +**Additional acceptance criteria:** the same marker-preservation regressions +cover sensitive registered request methods through source decoding, direct +construction, serde and insertion, plus both borrowed conditional views +through wildcard and tagged single/mixed repeated lines. Preserve the +existing method spelling and conditional-tag wire behavior. A sensitive +entity tag must remain redacted when the resulting sink is formatted; these +are extensions of this existing root, not new privacy-classification work. + +**Guideline obligation:** [M-STRONG-TYPES-GUARD](https://microsoft.github.io/rust-guidelines/guidelines/libs/resilience/#M-STRONG-TYPES-GUARD) +requires strong types to enforce their encoded invariants where applicable. +Here the retained sensitivity property must survive implicit reconstruction; +explicit caller-directed sensitivity changes remain allowed. Test the final +sink representation as well as the intermediate typed value. + +--- + +## Conformance + + +### CON1 — Make checked byte-to-text conversions explicit + +**Area:** borrowed byte-oriented text accessors · **Priority:** Low · **Effort:** Small + +**Guideline:** [C-CONV](https://rust-lang.github.io/api-guidelines/naming.html#c-conv) — should use `as_` for free borrowed projections and `to_` for expensive conversions, explicitly including borrowed UTF-8 validation · **Confidence:** High · **Scope:** 2 checked accessors, exhaustive for `UserAgentView` and `ServerView`; sampled across the wider conversion API + +Make `to_str` the canonical spelling for these checked conversions. They can +scan the complete byte sequence and reject invalid UTF-8; they are not free +projections of a retained string. The two opaque header views expose +the checked operation as `as_str` while internally calling `to_str`. + +- `crates/http_headers/src/headers/user_agent.rs:158` — `as_str` calls `.to_str()` at line 160 +- `crates/http_headers/src/headers/negotiation/server.rs:161` — the same checked delegation is at line 163 + +Retain compatibility aliases if needed and explicitly document their +validation cost; do not replace validation with unchecked conversion or +rename constant-time getters such as `LocationView::as_str`. This is a naming +and discoverability correction, not a claim of measured performance loss. + +**Done when:** both types offer and teach the checked `to_str` operation, +any retained `as_str` aliases clearly identify the same checked cost, and +package-scoped documentation/API checks confirm existing callers still work. +Public-API tests preserve valid ASCII and multibyte UTF-8 results and invalid +byte rejection, including unchanged header-specific error kinds. Regenerate +crate README examples from rustdoc rather than editing generated files. + +--- + + +### CON3 — Establish current unsafe-code verification evidence + +**Area:** `http_headers_simd` unsafe acceptance evidence · **Priority:** Medium · **Effort:** Small + +**Guideline:** [M-UNSAFE](https://microsoft.github.io/rust-guidelines/guidelines/correctness/#M-UNSAFE) — unsafe code must have a valid reason, safety reasoning and passing Miri verification; performance-motivated unsafe should follow benchmarking · **Confidence:** High that current acceptance is unassessed, not that the code fails · **Scope:** 1 unsafe implementation crate; sampled ASCII, vector-load and allocator boundaries plus the existing Miri recipe + +Establish revision-attributable safety verification using the existing +Anvil checks. Static inspection finds local proof comments and a configured +Miri gate, but neither proves that the current revision passes it. This +static-only audit did not run Miri or observe current CI artifacts. It found +no confirmed undefined behavior; do not describe this evidence gap as an +unsoundness finding or as missing Miri tooling. + +- `crates/http_headers_simd/src/api.rs:39` — `str::from_utf8_unchecked(bytes)` follows the ASCII proof +- `crates/http_headers_simd/src/x86.rs:33` — `_mm_loadu_si128(pointer)` has a local slice-bound proof +- `crates/http_headers_simd/src/tracking.rs:127` — `unsafe impl GlobalAlloc for TrackingAllocator` delegates allocator obligations to `System` +- `justfiles/anvil/checks/miri.just:35` — existing package-scoped `anvil-miri` recipe; its implementation selects `--all-features` test targets + +**Done when:** current-commit artifacts from +`just anvil-miri --package http_headers_simd` and +`just anvil-miri --package http_headers` establish the applicable checks' +outcomes, with exact toolchain, target, feature selection and test selection. +Account explicitly for any unsupported intrinsic, target, ignored test or +unselected path instead of implying full coverage. Investigate observed +failures if any; do not invent them from configuration alone. Ordinary +isolated no_std correctness checks are separate from this unsafe-code +verification, and the repository's no_std-only coverage/mutation exemptions +remain unchanged. + +**See also:** B3, B5, B6 own the fresh performance evidence for unsafe +optimization decisions; this item does not duplicate their measurement work. + +--- + +## Performance + +Evaluate these candidates against the typed semantic APIs. They are +hypotheses to re-evaluate, not patches to apply blindly or promised gains. +Establish fresh equivalent-work baselines for the resulting APIs, pin the +compiler and target flags, and preserve executable paths, fixture inventories, +and measurement boundaries. Judge CPU cost using instruction counts rather +than short wall-time samples, and independently reject increases in measured +allocation counts or allocated bytes: +every existing case must remain equal or improve, including owned/borrowed, +raw/HTTP, repeated-line, fallback, and error cases. Never offset a regression +with gains elsewhere or weaken validation, source limits, error precedence, +or safety checks. Inspect assembly and additive Callgrind attribution after +each improvement. Leave `PERF.md` regeneration to the maintainer. + +Prioritize the shared representation/proof boundary (P21), retained semantic +layout (P24), and negotiation traversal strategy (P25) before polishing their +inner loops. These are alternative experiments, not permission to redesign +the public facade API or introduce unsafe code into it. Every candidate below +still awaits measurement; no current workload share or end-to-end speedup is +established. Resolve the linked security prerequisites first. +Candidate redesigns with unmeasured benefit do not outrank confirmed contract +defects. B1's observable-reader repair precedes trusting affected baselines; +B4 then makes accepted performance reproducible and enforceable rather than a +one-off measurement. + + +### P1 — Accelerate Cache-Control fallback token validation + +**Area:** Cache-Control fallback validation · **Priority:** Medium · **Effort:** Small + +**Evidence:** Reasoned · **Expected impact:** fewer classification instructions per nonnumeric/overflow token, still linear in token bytes; size and overall workload share unknown · **Risk:** short-input regressions and changed error precedence; no new API or unsafe · **Scope:** 1 fallback validator, exhaustive + +Each directive routed through `validate_token_value` pays this fallback after +an unsuccessful seconds parse, or directly for a nonnumeric value. +`crates/http_headers/src/headers/cache_control.rs:1025` uses +`for byte in bytes.iter().copied()`; the proposed replacement changes that +scan, not the need to validate the entire token. The end-to-end ceiling is +limited to the unknown share spent in these directive scans. + +Evaluate the existing word-based token validator as a replacement for the +scalar fallback after unsuccessful decimal parsing and for nonnumeric +directive values. Keep complete validation without introducing another +numeric pass or changing the separate quoted-value grammar. + +- `crates/http_headers/src/headers/cache_control.rs:1016` — `validate_token_value` + +**Done when:** complete per-case instruction measurements support retaining +or rejecting the replacement, and +oracles cover `u64::MAX`, overflow, long leading zeroes, nondigit tokens, +invalid suffixes, and the separate quoted-value path. + +**Decision:** use B3's `http_headers_policy_shapes` overflow/token/quoted +strata and B1 directive reads with B4's per-case gate. Retain a measured win +or close the candidate as not worthwhile with the comparison recorded. + +--- + + +### P6 — Reduce HTTP-date digit-validation overhead + +**Area:** conditional-header shared date parser · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Reasoned · **Expected impact:** lower fixed digit-validation cost per date decode; no measured magnitude or known overall date share · **Risk:** accepting bad digits, calendar changes, or precision drift; no new API or unsafe · **Scope:** 2 shared digit helpers, exhaustive + +The fixed IMF parser calls `two_digits` six times; relaxed normalization +uses the short-decimal helper for day, year and three clock fields. +`crates/http_headers/src/headers/conditional/shared.rs:1107` checks +`value.as_bytes().iter().all(u8::is_ascii_digit)` before `value.parse()` at +line 1110. This is a bounded per-date cost, not an unbounded parsing +bottleneck; its unknown request share caps the possible overall gain. + +Evaluate packed two-digit validation and the separate digit-validation pass +before accumulation in `parse_short_decimal`. Preserve checked conversions, +calendar validation, and the allocation-free fixed-width normalization path. + +- `crates/http_headers/src/headers/conditional/shared.rs:1106` — short decimal validation and accumulation +- `crates/http_headers/src/headers/conditional/shared.rs:1230` — two-digit validation + +**Done when:** instruction measurements across every date header and If-Range +support retaining or rejecting each candidate; exhaustive byte-pair checks +and an independent date parser preserve normalization, leap/calendar validity, +supported years, and malformed forms. + +**Decision:** B3's conditional/range date strata and B1 cached date-read +controls decide acceptance or closure as not worthwhile, under B4. Preserve +If-Range metadata's agreement with wire precision when comparing date paths. + +--- + + +### P9 — Reuse accelerated validation and searches for entity tags + +**Area:** ETag construction and conditional tag lists · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Reasoned · **Expected impact:** fewer instructions in linear opaque-byte and delimiter scans per constructed/read tag; magnitude and overall share unknown · **Risk:** conflating opaque tag bytes with list grammar; no new API or unsafe · **Scope:** 4 named constructor/validator/iterator routines, sampled call paths + +`crates/http_headers/src/headers/etag.rs:372` performs +`opaque.iter().copied().all(valid_opaque_byte)` once per constructor call. +Conditional readers traverse each requested tag list anew. Long tags offer +more scan work to amortize a wide helper, while early invalid and short tags +may lose; the unknown construction/read frequency limits end-to-end impact. + +Evaluate using the existing word validator in ETag constructors instead of +the bytewise predicate. For If-Match and If-None-Match, investigate the +existing delimiter-search helpers for closing quotes and long-tag scans; +preserve the distinct roles of commas, semicolons, backslashes, and whitespace +inside opaque tags. + +- `crates/http_headers/src/headers/etag.rs:371` — `construct` +- `crates/http_headers/src/headers/etag.rs:444` — existing `opaque_is_valid` +- `crates/http_headers/src/headers/conditional/shared.rs:776` — `TagIter` +- `crates/http_headers/src/headers/conditional/shared.rs:952` — `validate_tag_line_with` + +**Done when:** constructor, decode, and semantic-reader measurements establish +each retained gain, with short/long/obs-text tags, weak markers, wildcard +mixing, repeated lines, and exact error precedence unchanged. + +**See also:** B1, B3. + +**Decision:** compare measured-body constructors and decode-plus-tag-reader +cases in B1/B3, not constructors hidden in fixture setup. Retain only B4 +nonregressing wins or close as not worthwhile with the measured rejection. SEC3's borrowed +conditional-tag forwarding marker fix remains a prerequisite for forwarding +comparisons. + +--- + + +### P10 — Evaluate Content-Type common-literal dispatch + +**Area:** Content-Type parse pipeline · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Speculative · **Expected impact:** possibly fewer literal-dispatch instructions per Content-Type decode; compiler lowering, magnitude and overall share unknown · **Risk:** redundant manual dispatch, near-miss regressions and code-size growth; no new API or unsafe · **Scope:** 1 recognizer with 8 literal spellings, exhaustive + +Every metadata parse probes `common_metadata` before the general parser. +`crates/http_headers/src/headers/content_type.rs:696` first compares +`bytes == b"application/json; charset=utf-8"` and line 699 starts +`match bytes`. The compiler may already group comparisons by length; source +comparison count is not emitted instruction count. The ceiling is the +unknown share spent recognizing these values, not all Content-Type work. + +Investigate length dispatch for common literal recognition, with original-code +hit and early/late near-miss controls before changing the recognizer. A faster +hit is not useful if equal-length unknown values pay for additional probes. + +- `crates/http_headers/src/headers/content_type.rs:695` — `common_metadata` + +**Done when:** representative hit, miss, casing, and fallback measurements +support retaining or rejecting the candidate under the complete instruction +gate. Preserve semantic metadata, exact wire bytes, error positions, and +allocation behavior. + +**See also:** B3. + +**Decision:** B3's hit/early-miss/late-miss strata and B5 ordinary +release versus fat-LTO disassembly decide whether explicit dispatch adds +anything. Accept a B4-gated win or close as not worthwhile. Keep wire-value +identity independent of cache representation; coordinate P26's layout study. + +--- + + +### P11 — Resume Content-Type lookup after the cached parameter prefix + +**Area:** parameter semantic accessors · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Structural · **Expected impact:** avoid rescanning up to the cached prefix on each later/missing lookup; saved bytes and overall lookup share unknown · **Risk:** noncontiguous/large offsets and quoted/duplicate semantics; no new API or unsafe · **Scope:** 1 lookup function and its 2-entry prefix cache, exhaustive + +`crates/http_headers/src/headers/content_type.rs:957` starts +`ParameterScanner::new(bytes, ...)` at `head.parameter_start`, then calls +`.skip(usize::from(head.inline_count))`. Thus a caller making repeated +uncached lookups pays again for the cached wire prefix. The saving is bounded +by that prefix per lookup, not the entire remaining parameter list, and +application lookup frequency is unknown. + +Lookup beyond the inline cache restarts a scanner at `parameter_start` and +skips already cached entries. Investigate resuming at the last cached value +end instead. The existing cached `charset` lookup is not evidence for a +third-parameter hit or a missing-name lookup. + +- `crates/http_headers/src/headers/content_type.rs:946` — `find_parameter` + +**Done when:** original-code baselines cover later/missing lookups in both +ownership modes; retained gains preserve checked offsets, cache-contiguity +invariants, first-match duplicate semantics, quoted values, empty slots, and +large-offset fallbacks. + +**See also:** B1, B3. + +**Decision:** use B1/B3's first/third/last/missing lookups with quoted long +prefixes and B5 layout controls. Record either a retained B4 nonregressing +win or closure as not worthwhile; coordinate P26 so cursor storage does not +silently enlarge every common result. + +**Guideline connection:** [C-INTERMEDIATE](https://rust-lang.github.io/api-guidelines/flexibility.html#c-intermediate) +recommends useful intermediate results that avoid duplicate work. Preserve +the existing retained metadata and public semantic APIs; this experiment is +about the uncached lookup's resume point, not reparsing every getter or +requiring new public intermediate types. + +--- + + +### P12 — Reuse delimiter searches in negotiation validation + +**Area:** Accept and Accept-Language scanners · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Reasoned · **Expected impact:** fewer delimiter-search instructions per media/language member; remains linear, with magnitude and overall share unknown · **Risk:** extra short-token dispatch and changed malformed-input precedence; no new API or unsafe · **Scope:** 2 member validators, exhaustive + +`crates/http_headers/src/headers/negotiation/accept.rs:225` uses +`bytes.splitn(2, |byte| *byte == b'/')`; language validation similarly discovers +subtag boundaries per member. A header decode reaches these routines for +members not settled by earlier recognition paths, so the traffic share of +the fallback, not the mere presence of Accept, bounds the benefit. + +Evaluate the existing safe `find_either` helper for Accept's +first-slash split and language-subtag delimiter discovery. Preserve the cheap +overlong-subtag rejection before character validation rather than assuming +a single pass is always better. + +- `crates/http_headers/src/headers/negotiation/accept.rs:225` — media-range split and repeated-slash error precedence +- `crates/http_headers/src/headers/negotiation/accept_language.rs:141` — language-range subtag validation + +**Done when:** each retained change passes canonical, raw, repeated, long, +wildcard, and malformed cases without losing the typed semantic output; +vector-boundary slash positions and language length boundaries have exact +result/error oracles and fresh instruction baselines. + +**See also:** B1. + +**Decision:** compare B3 negotiation shapes with B1 typed read/selection +controls after considering P25's broader traversal design. B4 measurements +must support acceptance; an unchanged or slower result closes this candidate +as not worthwhile without a forced implementation. + +--- + + +### P13 — Reduce Host literal and international-fallback overhead + +**Area:** Host parsing · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Reasoned · **Expected impact:** reduce literal delimiter scans or common-path code footprint per Host decode; magnitude and overall literal/IDNA share unknown · **Risk:** short-host regressions, normalization/framing drift and unhelpful outlining; no new API or unsafe · **Scope:** 2 fallback routines, exhaustive + +`crates/http_headers/src/headers/negotiation/host.rs:853` separately checks +`bytes.contains(&b'@')`; line 858 searches +`.position(|byte| *byte == b']')`. These execute on bracketed literal +decodes, while IDNA belongs only to relaxed international fallback. +Their frequencies are not established, and a cold-boundary gain requires +emitted-code evidence rather than assuming the compiler inlined either path. + +Examine a cold boundary around the international-name fallback and reuse safe +delimiter searches for `@` and closing brackets. Measure both malformed +framing and valid long literals rather than assuming a scanner improves +short hostnames. + +- `crates/http_headers/src/headers/negotiation/host.rs:747` — IDNA fallback +- `crates/http_headers/src/headers/negotiation/host.rs:759` — normalized-name delimiter rejection +- `crates/http_headers/src/headers/negotiation/host.rs:853` — literal framing and closing-bracket search + +**Done when:** owned and borrowed Unicode-IDNA controls are baselined alongside +strict/relaxed ASCII, IPv6/IPvFuture, delimiter boundaries, and malformed +ports; only per-case nonregressing changes remain, with original error +precedence and typed components intact. + +**See also:** P24, B3. + +**Decision:** B3 literal/IDNA shapes, the authority semantic target, and B5 +cross-profile attribution decide acceptance or closure as not worthwhile +under B4. Preserve sensitive Host diagnostic redaction and SEC3 sensitivity +when comparing retained metadata or forwarding. + +--- + + +### P14 — Investigate Authorization validation data flow + +**Area:** Basic and Bearer validation · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Speculative · **Expected impact:** possibly lower validation instructions per credential block or token; compiler behavior, magnitude and overall authorization share unknown · **Risk:** padding/colon/zeroization mistakes and code-size regressions; no new API or unsafe · **Scope:** 2 scheme-validation paths, exhaustive; codegen effects unmeasured + +`crates/http_headers/src/headers/authorization.rs:863` visits +`body.chunks_exact(4)` and line 873 accumulates `colons |= colon_marks(group)`. +This is once per Basic decode, not a fresh decoded allocation. Bearer's +local scalar/SIMD crossover and the companion's crossover need matched +length/backend controls. The optimizer may already remove the suspected +boundary overhead; unknown decode/extraction shares cap any overall gain. + +Inspect Basic's validator boundary and decoded-colon data flow, and Bearer's +short scalar tail versus long-token dispatch. Look for redundant work in +generated code rather than assuming extra inlining or outlining helps. + +- `crates/http_headers/src/headers/authorization.rs:857` — `validate_basic` +- `crates/http_headers/benches/http_headers_auth_cors_shapes.rs:18` — source/ownership and adverse-shape measurement target + +**Done when:** each concrete change has attributed instruction gains across +short/long tokens, early/late invalid bytes, padding, and missing-colon +inputs; extraction semantics, sensitive storage, and zeroization are +preserved. Unsupported hypotheses do not become production changes. + +**See also:** B3. + +**Decision:** use B3 auth shapes and separately bounded cold/warm extraction, +with B5 disassembly and B4 per-case measurements. Retain a justified change +or close as not worthwhile; a boundary annotation without an attributed win +does not complete the item. + +--- + + +### P15 — Reduce repeated CORS list traversal and projection + +**Area:** Allow-Headers, Expose-Headers, Request-Headers, and Allow-Methods · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Structural · **Expected impact:** reduce repeated linear line/member work per decode-plus-reader lifecycle; dispatch savings and overall CORS share unknown · **Risk:** code duplication, larger results and changed wildcard/error-index semantics; no new API or unsafe · **Scope:** 4 list families using the shared validator, exhaustive + +Each decode validates physical lines, and each requested semantic traversal +projects their members again. `crates/http_headers/src/headers/cors/shared.rs:926` +takes `values: &mut dyn Iterator>`; line 935 +iterates `(1_usize..).zip(values)`. The traversal exists structurally, but +whether indirect calls survive depends on the profile. Repeated readers and +line counts determine the possible saving; their application share is unknown. + +Investigate repeated list traversal, erased-validator call boundaries, and +member projection after validation. Preserve existing wildcard, method/name +comparison, custom-source preflight, and error-index semantics; use existing +typed item machinery rather than an alternate parser. + +- `crates/http_headers/src/headers/cors/shared.rs:910` — shared list validator +- `crates/http_headers/src/headers/cors/shared.rs:1043` — per-member validation + +**Done when:** decode and semantic-reader baselines independently demonstrate +the retained gains for singleton/repeated, wildcard, mixed-case, extension, +and early/late error cases, with no regression in other users of shared code. + +**See also:** B1. + +**Decision:** compare B1 decode-only/1/2/8 reads, B2 representation/line-count +strata and B5 erased-versus-specialized codegen. Accept only B4-supported +wins or close as not worthwhile with a measured rejection. Like P21, retain the exact validated +`FieldLines`; never reacquire a changing source to reuse its validation proof. + +--- + + +### P18 — Tighten CORS scalar token decoding + +**Area:** Allow-Credentials and Request-Method · **Priority:** Low · **Effort:** Medium + +**Evidence:** Speculative · **Expected impact:** possibly fewer fixed framing/comparison instructions per scalar decode; magnitude and overall share unknown · **Risk:** trading cheap failures for wider work or dropping sensitivity; no new API or unsafe · **Scope:** 2 scalar decoders, exhaustive + +`crates/http_headers/src/headers/cors/access_control_allow_credentials.rs:217` +calls `is_credentials_true(lines.exactly_one()?.as_bytes())`; +`crates/http_headers/src/headers/cors/access_control_request_method.rs:350` calls +`request_method_of(value.as_bytes())?`. Each is one scalar decode after +source checks. Their already-small operation and unknown request frequency +bound the opportunity; emitted comparison lowering may already be optimal. + +Inspect framed-whitespace/range data flow around the `true` literal and +known-method literal comparison lowering. These already-cheap paths need +attribution-backed candidates, not unconditional wider comparisons that +penalize early failures. + +- `crates/http_headers/src/headers/cors/access_control_allow_credentials.rs:209` — singleton borrowed decode +- `crates/http_headers/src/headers/cors/access_control_request_method.rs:341` — method borrowed decode + +**Done when:** any retained change improves instructions while preserving +owned/borrowed, padded, duplicate, case-sensitive, unknown-method, and +same-length near-miss behavior; stop without a code change if no such +candidate remains. + +**Decision:** resolve SEC3's registered-method sensitivity loss before +comparing compact states or forwarding. B3 scalar near-miss/framing strata +and B5 generated code decide acceptance under B4 or closure as not worthwhile. + +--- + + +### P19 — Investigate Content-Length singleton and decimal decoding + +**Area:** numeric header decoding · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Structural · **Expected impact:** avoid a separate singleton comma scan, with possible decimal-loop instruction savings; magnitude and overall Content-Length share unknown · **Risk:** overflow, full-consumption and duplicate/error-precedence changes; no new API or unsafe · **Scope:** 1 singleton decision and 2 decimal/list routines, exhaustive + +`crates/http_headers/src/headers/content_length.rs:125` probes +`!first.as_bytes().contains(&b',')` before `parse_decimal_ows` at line 126. +The decimal loop at line 179 uses +`value.checked_mul(10)?.checked_add(...)` for every digit. Fuse discovery +before considering grouped accumulation; both are conditional experiments, +not permission to omit checked arithmetic or the custom-source preflight. +Unknown source and numeric-length distributions limit end-to-end benefit. + +Inspect short-decimal accumulation, singleton scanning, and call boundaries +for unnecessary work before falling back to duplicate/list handling. Preserve +checked arithmetic and validation of every supplied value. +In particular, investigate recognizing a singleton number and its first comma +in one pass without adding a speculative numeric parse before the comma-list +fallback. Measure malformed and comma-joined inputs as well as bare numbers. + +- `crates/http_headers/src/headers/content_length.rs:158` — `parse_content_length_line` +- `crates/http_headers/benches/http_headers_auth_cors_shapes.rs:18` — existing Content-Length shape coverage + +**Done when:** retained gains satisfy canonical, 19/20-digit, leading-zero, +equal/conflicting duplicate, overflow, and malformed-list instruction gates +with identical values and exact errors. + +**Decision:** B3's Content-Length shape series, B2's raw/HTTP source matrix +and B5 attribution decide whether either transformation earns a B4-gated +win. Close as not worthwhile if it does not; count long leading zeroes as +validation work even though at most 20 significant digits fit in `u64`. + +--- + + +### P21 — Reduce shared source and result construction overhead + +**Area:** `Field`, `FieldLines`, and owned/view forwarding · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Structural · **Expected impact:** reduce repeated preflight passes, list growth or raw/result movement per typed operation; emitted copies, magnitude and overall share unknown · **Risk:** proof reuse across changing sources, larger layouts and lost markers/atomicity; no new API or unsafe · **Scope:** 4 source representations and 2 generic forwarding entry points, exhaustive; downstream consumers and serde seam sampled + +For example, `crates/http_headers/src/source/field_lines.rs:356` calls +`self.validate_custom_source()?` again when acquiring repeated owned values. +Consider a private validated-ownership state tied to the exact immutable +snapshot, rather than deleting checks or asking `FieldSource::lines()` again. +At line 390, the HTTP `len()` arm uses `values.iter().count()`: do not add a +pre-count traversal merely to reserve a collector. + +Compare singleton promotion and closed-enum dispatch across an entire +decode/read/forward lifecycle, not only a wrapper. Also price the serde name +boundary: `crates/http_headers/src/serde_impls.rs:482` uses +`String::deserialize(deserializer)?` even for well-known enum names; a +borrowed string visitor is an alternative to that temporary, not a request +for a new borrowed public deserialization API. Grammar checks and custom +budgets still apply. The shared route has broad reach, but actual caller +frequency and its fraction of request cost remain unknown. + +After the typed APIs settle, inspect source lookup, representation dispatch, +temporary result/list copies, singleton owned conversion, and budget-policy +boundaries separately from grammar parsing. Shared changes must be followed +by fresh inspection of every affected header, not just a canonical count. + +- `crates/http_headers/src/field.rs:403` — generic single-value borrowed forwarding +- `crates/http_headers/src/field.rs:420` — owned forwarding +- `crates/http_headers/src/source/field_lines.rs:465` — custom-source bounds +- `crates/http_headers/src/source/field_lines.rs:491` — source validation + +**Done when:** concrete source/result-copy candidates are measured across +all representations and affected headers; only per-case instruction wins +remain, with custom-name behavior, source/list limits, sharing, and defensive +checks preserved. + +**See also:** B2, B5 (deciding measurements). Preserve the independent +admission-policy oracles in `tests/source_limits.rs` and `tests/serde.rs`. + +**Decision:** B2's complete source/ownership/serde matrix and B5's ordinary +release layouts/copies decide acceptance or closure as not worthwhile under +B4; B6 bounds the potential request-level gain. SEC3 precedes forwarding +experiments, including registered methods and borrowed conditional tags. +Preserve direct/source and borrowed/owned admission parity in equivalent-work +comparisons. + +**Guideline connection:** [C-INTERMEDIATE](https://rust-lang.github.io/api-guidelines/flexibility.html#c-intermediate) +recommends reuse of useful intermediate results. Any private validation proof +must remain attached to the original immutable `FieldLines`; a later +`FieldSource::lines()` call is not the same proof-bearing input. Keep that +condition in correctness checks for any measured optimization. + +--- + + +### P24 — Reduce retained URI and authority metadata costs + +**Area:** Location, Host, and Access-Control-Allow-Origin decoding · **Priority:** Medium · **Effort:** Large + +**Evidence:** Structural · **Expected impact:** reduce retained result footprint or repeated boundary discovery per decode; layout, CPU magnitude and overall workload share unknown · **Risk:** larger fallback variants, normalization lifetimes, semantic/redaction drift; no new public facade API, unsafe only in the existing companion if required · **Scope:** 3 semantic header families, exhaustive; layout/codegen alternatives sampled + +Retained semantic metadata removes caller-side parsing but increases +decode-only work. Investigate compact component boundaries, result movement, +and classification costs without discarding parsed addresses, normalized +backing, exact port states, or original wire storage. Do not recover cheap +decoding by moving URI parsing, IDNA conversion, or allocation into getters. + +**Provenance constraint:** the comparison numerals in the next paragraph are +retained legacy backlog notes, not measurements established by this audit. +`crates/http_headers/docs/PERF.md:3` explicitly identifies its table as an +imported 0.1.0 snapshot. No current raw artifact/compiler/hardware provenance +for the newer numbers has been verified here; do not use them as a gate or +as evidence of the size of this candidate. + +Canonical Callgrind instruction counts against the existing `PERF.md` +baseline are Location borrowed 595 → 786 and owned 822 → 1029; Host borrowed +527 → 683 and owned 548 → 708; Allow-Origin borrowed 555 → 583 and owned +670 → 694. These are decode-only costs, not complete consumer workloads. +The semantic benchmarks must remain a separate acceptance axis; aggregate +improvements must not conceal a slower individual case. + +The structural trace is independent of those numbers: +`crates/http_headers/src/headers/location.rs:367` sets +`metadata: Metadata::from_simple(text)` after subset validation; +`crates/http_headers/src/headers/location/metadata.rs:51` begins further +`find_either` boundary searches. +That metadata stores `Range` and optional full-width boundaries +at lines 14–18. Compare emitting boundaries during the original safe scan +and compact checked offsets with a full-width fallback. Include Host/Origin +classification and retained normalization backing in the same layout study, +without replacing useful retained semantics with a lazy parser. Per-decode +work and result movement may shrink, but the common/general/relaxed mix and +the share of callers doing semantic reads are unknown. + +- `crates/http_headers/benches/http_headers_per_header.rs:355` — canonical Host decode workloads +- `crates/http_headers/benches/http_headers_per_header.rs:401` — canonical Location decode workloads +- `crates/http_headers/benches/http_headers_per_header.rs:272` — canonical Allow-Origin decode workloads +- `crates/http_headers/benches/http_headers_location_semantics.rs` — decode, repeated component reads, normalization, and forwarding +- `crates/http_headers/benches/http_headers_authority_semantics.rs` — retained authority semantics versus caller parsing + +**Done when:** fresh measurements with the baseline compiler/target flags +reduce retained-metadata overhead across both ownership modes without +regressing semantic workloads, grammar/error behavior, or allocations; +value layouts and generated-code effects explain the retained changes. +Keep `PERF.md` updates with the maintainer. + +**See also:** SEC3 (retained-marker prerequisite); P13, B4, B5. + +**Decision:** use B1's Location/authority decode, retained and 1/2/8-read +targets, B2 ownership/forwarding, and B5 layouts; B6 prices workload shares. +The Location caller baseline already retains its parsed URI, and both +authority comparison arms use the current decoder: neither is a historical +decoder baseline. Accept a measured nonregressing design or close as not +worthwhile. Preserve sensitive Host diagnostic redaction and SEC3 marker +propagation throughout. + +--- + + +### P25 — Compare streaming negotiation with bounded retained member metadata + +**Area:** Accept, Accept-Encoding, Accept-Language, Allow and Vary semantic lifecycles · **Priority:** Medium · **Effort:** Large + +**Evidence:** Structural · **Expected impact:** avoid repeated wire/member projection work proportional to reader count times wire length; no measured magnitude or known overall share · **Risk:** eager decode work, larger results, cache overflow complexity and semantic drift; no new public API or unsafe +**Scope:** 5 typed list families, exhaustive; proposed private representations are experimental + +Keep streaming as the control, but compare a bounded inline member index +collected during validation with today's wire-only retained lists and with +caller-retained yielded entries. Each new Accept iterator traverses raw +members and reconstructs media/quality/parameter boundaries. Encoding and +language entries similarly project token/range and quality; Allow/Vary +iterate validated tokens without allocating, but still rediscover members. +This is a representation/validation-to-reader boundary question, not P12's +smaller choice of delimiter-search primitive. + +- `crates/http_headers/src/headers/negotiation/accept.rs:111` — owned `.flat_map(...)` followed by `.map(AcceptEntry::from_validated)`; borrowed traversal starts at line 139 with `.validated_comma_items()` at line 141 +- `crates/http_headers/src/headers/negotiation/accept_entry.rs:59` — `bytes.iter().position(|byte| *byte == b';')` begins per-entry boundary reconstruction +- `crates/http_headers/src/headers/negotiation/quality.rs:120` — `from_validated` derives a compact or exact fractional view +- `crates/http_headers/src/headers/negotiation/allow.rs:78` — `self.items().map(MethodView::from_validated)` +- `crates/http_headers/src/headers/negotiation/vary.rs:129` — `self.items().map(VaryEntryView::from_validated)` +- `crates/http_headers/src/headers/negotiation/accept_scan.rs:367` — the acceptance-only DFA visits `bytes.chunks_exact(8)`; an unsupported shape is subsequently handled by the general parser + +For `r` fresh traversals of `n` wire bytes, the repeated projection component +is O(r × n). A bounded retained index might exchange this for one +validation/indexing pass and cheap member projections, but decode-only and +one-read consumers may lose. Read counts and input frequencies are unknown, +so no overall speedup or minimum useful cache size can be promised. Include +early-decline versus full acceptance-recognizer fallback as a smaller +pipeline alternative; it may add an unhelpful common-path branch. + +Do not build a policy engine, sort, deduplicate, round qualities, or use an +unbounded per-member allocation. Preserve exact relaxed fractions, duplicate +and wildcard behavior, quoted parameter/extension boundaries, physical lines, +error kinds/indices and budgets. An overflow path must retain today's +streaming semantics. Validated views must retain the original immutable +`FieldLines`, never reacquire a mutable source. + +**Done when:** B1/B3 extend `http_headers_negotiation_semantics` and +`http_headers_typed_tokens` with decode-only, 1/2/8 fresh traversals, +caller-retained entries, owned readers, first/last/missing matches, +quoted/relaxed/repeated and index-overflow controls. B5 records result layout +and downstream codegen; B6 establishes the read-count/workload crossover. +Accept a design only with B4's per-case instruction/allocation nonregression, +including decode-only, or close the experiment as not worthwhile with its +measurements. Preserve SEC3 behavior in inspect-and-forward comparisons. + +**See also:** SEC3 (retained-marker prerequisite); P12, P21, B1, B3, B5, B6. + +--- + + +### P26 — Evaluate inline general Content-Type metadata + +**Area:** nonliteral Content-Type decode, construction and clone · **Priority:** Medium · **Effort:** Medium + +**Evidence:** Structural · **Expected impact:** potentially remove one explicit metadata Box allocation per successful general parse/construction or deep metadata clone; time/footprint trade and overall nonliteral share unknown · **Risk:** larger common values/results and changed cache identity; no new API or unsafe +**Scope:** 1 metadata enum and 3 production Box-construction sites, exhaustive + +General Content-Type metadata is a fixed-size head with two inline parameter +ranges, yet it sits behind a box. Compare a compact inline general form with +the existing boxed form and literal variants. Do not assume saving the +allocation outweighs moving a larger owned/view result through every common +decode. This also differs from P10's recognizer and P11's lookup cursor. + +- `crates/http_headers/src/headers/content_type.rs:158` — `Parsed(Box)` +- `crates/http_headers/src/headers/content_type.rs:283` — construction uses `Box::new(ContentTypeHead { ... })` +- `crates/http_headers/src/headers/content_type.rs:677` — strict general parsing ends with `ContentTypeMetadata::Parsed(Box::new(head))` +- `crates/http_headers/src/headers/content_type.rs:692` — relaxed general fallback uses the same allocation + +The parser pays this once per successful nonliteral decode, including a +borrowed decode; literal hits avoid it. Derived metadata cloning also clones +the box. These are structural allocation sites, not current allocation +measurements or evidence that nonliteral inputs dominate deployed traffic. +That unknown fraction and subsequent read/clone frequency cap the overall +gain. Keep eager validation and constant-time scalar access; do not shift +parsing or allocation into getters to make decode look cheaper. + +**Done when:** use B1/B3's literal controls and nonliteral +0/1/2/3/16/49-parameter, quoted, relaxed and large-offset cases. B2 includes +direct construction and clone lifecycles; B5 compares layouts and emitted +copies under ordinary release and fat LTO. +Record allocations, allocated bytes, instructions and time per ownership +mode. Accept only B4-nonregressing results, with identical wire/semantic +identity and parameter behavior, or close as not worthwhile with the +comparison. Coordinate P11's resume cursor rather than creating competing +metadata layouts. + +**Guideline connection:** [M-AVOID-INDIRECTION](https://microsoft.github.io/rust-guidelines/guidelines/performance/#M-AVOID-INDIRECTION) +recommends avoiding needless indirection in hot types. Whether this box is +needless is unassessed until the common/general mix and result-layout +tradeoff are measured; the guideline does not mandate unconditional inlining +or justify a new cache identity. + +**See also:** P10, P11, B1, B2, B3, B5. + +--- + + +### P27 — Transfer completed Allow and Vary construction buffers + +**Area:** typed response-list construction into owned field storage · **Priority:** Low · **Effort:** Small + +**Evidence:** Structural · **Expected impact:** avoid the final O(n) wire copy and potentially its additional backing allocation above 64 bytes; frequency and overall share unknown · **Risk:** short-result layout regressions or added validation work; no new API or unsafe +**Scope:** 2 constructor bodies, exhaustive; Vary's name constructor delegates here + +`AllowOwned::from_methods` and `VaryOwned::from_entries` append into a String +and immediately construct final storage from its borrowed bytes. Long output +is copied into another backing allocation before that String is dropped. +Compare safe ownership transfer of the completed long buffer while retaining +the short inline result behavior; account for any new validation pass in the +comparison rather than presuming conversion is free. + +- `crates/http_headers/src/headers/negotiation/allow.rs:87` — `let mut wire = String::new()`; line 95 calls `FieldValue::from_validated_bytes(wire.as_bytes(), false)` +- `crates/http_headers/src/headers/negotiation/vary.rs:112` — the corresponding String builder; line 120 repeats the borrowed final conversion +- `crates/http_headers/src/field_value.rs:124` — the long `Repr::new` arm uses `Bytes::copy_from_slice(bytes)` +- `crates/http_headers/src/field_value.rs:636` — consuming String conversion already uses `Self::from_shared(Bytes::from(value))` + +This is once per typed list construction, not every request decode. The +constructed length grows with caller-supplied tokens; token lengths, list +sizes and construction frequency are unknown. Preserve empty lists, +spelling, order, duplicates and mixed wildcard/name semantics. Do not replace +the borrowed decoding API or introduce another grammar implementation. + +**Done when:** B2 measures construction inside the operation, including +0/1/many members, exact/unknown iterator hints, 63/64/65-byte results and +construct-to-sink lifecycles. Compare instructions, allocations and bytes +with B1 semantic/wire controls. Retain a B4-nonregressing transfer or close as +not worthwhile with the comparison, without changing validation or +sensitivity contracts. + +**See also:** B1, B2, B4. + +--- + +## Features + + +### F1 — Classify semantically sensitive typed headers with `data_privacy_core` + +**Area:** `http_headers` typed header values · **Priority:** Medium · **Effort:** Medium + +Integrate at the typed-header boundary, where the header name supplies stable +semantics, rather than on `FieldValue`. The raw value types can represent both +sensitive and nonsensitive data, while `Classified::data_class` must always +return a class. Handwritten implementations should depend on +`data_privacy_core`, not the full macro and redaction-engine crate. The +existing `is_sensitive` marker remains the transport and secret-safe `Debug` +signal used by `http::HeaderValue` interoperability. + +- `crates/http_headers/src/headers/authorization.rs` — Authorization descriptor and typed owned/view values +- `crates/http_headers/src/headers/set_cookie.rs` — Set-Cookie descriptor and typed owned/view values +- `crates/http_headers/src/headers/location.rs` — Location descriptor and typed owned/view values +- `crates/http_headers/src/field_value.rs` — representation-level Boolean sensitivity marker +- `crates/data_privacy_core/src/classified.rs` — `Classified` requires an unconditional `DataClass` + +**Done when:** an approved taxonomy defines classes for Authorization, +Set-Cookie, and Location; their borrowed and owned typed values implement +`Classified` through `data_privacy_core`; API tests cover every implementation; +and encoding still preserves the existing `is_sensitive` behavior. + +--- + + +### F2 — Add byte-oriented redaction before redacting typed headers + +**Area:** `data_privacy_core` redaction API and `http_headers` typed values · **Priority:** Medium · **Effort:** Large + +Do not expose policy-driven redaction for typed headers through the current +string-only API. HTTP field values can contain valid non-UTF-8 `obs-text`, and +Set-Cookie exposes its values as bytes. Converting such values lossily before +hashing or redaction would change their identity; rejecting them would leave a +valid sensitive-value path unsupported. + +- `crates/data_privacy_core/src/redactor.rs` — `Redactor::redact` accepts only `&str` and `fmt::Write` +- `crates/http_headers/src/headers/set_cookie.rs` — owned cookies iterate as byte-valued `FieldValue`s +- `crates/http_headers/src/headers/set_cookie.rs` — borrowed cookies iterate as byte-valued `FieldValueRef`s +- `crates/http_headers/src/headers/set_cookie.rs` — valid cookie field values are checked as bytes + +**Done when:** `data_privacy_core` provides an explicit byte-oriented +redaction contract, typed header redaction uses it without UTF-8 coercion, and +tests cover arbitrary valid field bytes including `obs-text` for pass-through, +replacement, and hashing policies. + +--- + + +### F3 — Convert `templated_uri::Uri` into `LocationOwned` + +**Area:** `templated_uri` integration with `http_headers::headers::LocationOwned` · **Priority:** Low · **Effort:** Small + +Provide an opt-in, one-way construction adapter for applications that build +classified redirect targets with `templated_uri`. Keep the integration in +`templated_uri` behind an optional dependency on `http_headers`; the +lower-level header crate should not acquire a dependency on the templating and +privacy stack. The adapter must materialize the URI for transmission and +produce a normal sensitive `LocationOwned`, without claiming to preserve the +template's per-component classifications after conversion. + +This conversion is deliberately not a replacement for `Location` validation. +`Location` accepts the complete RFC 3986 URI-reference grammar, including +relative, empty, and fragment-only references, while `templated_uri` models +HTTP request URIs and does not support fragments. + +- `crates/http_headers/src/headers/location.rs` — `Location` accepts RFC 3986 URI-references +- `crates/http_headers/src/headers/location.rs` — string conversion validates and marks the value sensitive +- `crates/templated_uri/src/lib.rs` — URI templates intentionally exclude fragments +- `crates/templated_uri/src/uri.rs` — `Uri` already materializes fallibly as `http::Uri` + +**Done when:** an optional `http_headers` integration feature in +`templated_uri` implements `TryFrom for LocationOwned` +with an error that preserves materialization and Location-validation failures; +converted values remain sensitive; and tests cover absolute and relative +path/query targets plus rejection or exclusion of unsupported fragment forms. + +--- + +## Documentation + + +### D1 — Document typed headers in the workspace HTTP APIs + +**Area:** `http_extensions`, `fetch`, and `rest_over_grpc` examples · **Priority:** Medium · **Effort:** Small + +Add concise, compiling examples showing how the workspace's HTTP-facing APIs +use `http_headers` without duplicating its parsing or encoding methods. +`http_extensions` and `fetch` should demonstrate typed fields on their built +`http::Request` and `http::Response` values and through builder +`headers_mut()` access. `rest_over_grpc` should use the explicit request and +response header-map accessors so the direction of each operation remains +clear. + +- `crates/http_extensions/src/http_request_builder.rs` — mutable request-builder headers +- `crates/http_extensions/src/http_response_builder.rs` — mutable response-builder headers +- `crates/fetch/src/lib.rs` — public `http` request, response, and header-map re-exports +- `crates/rest_over_grpc/src/context.rs` — explicit request-header access +- `crates/rest_over_grpc/src/context.rs` — explicit response-header access +- `crates/rest_over_grpc/src/http_response.rs` — mutable neutral-response headers + +**Done when:** checked examples decode Authorization or User-Agent from +requests and encode representative Location, ETag, and repeated Set-Cookie +response fields in each applicable crate, with feature requirements documented +and malformed input handled as `DecodeError`. + +--- + +## Benchmarks + + +### B1 — Cover semantic reads separately from decode-and-drop + +**Area:** per-header metabench coverage · **Priority:** High · **Effort:** Medium + +**Purpose:** maintain and improve. Several per-header rows only decode and +drop, so they cannot price repeated semantic reads. +Measure both `Field::view` and `Field::owned`, then separately consume semantic +outputs: weighted-list items, Content-Type parameters/lookups, Referrer +preference, CORS counts/wildcards, conditional tags, Range specs and WebSocket +extensions including parameters. Include cheap cached scalar/date getters as +controls. Use 1/2/8 reads and short versus multi-member inputs; do not infer +traffic weights from fixtures. + +Measure ETag construction inside the benchmark body; constructors used only +to prepare fixtures do not measure that operation. Cover short, long, obs-text, +and invalid opaque values to unblock P9 independently of decode benchmarks. + +Extend the typed-token, negotiation, URI, and authority semantic targets with +owned consumption and first/last/missing lookups rather than treating their +current fixtures as a complete matrix. Compare allocation-free streaming +negotiation members with a bounded inline metadata cache, including overflow +fallback, before choosing any additional decode-time indexing. + +Use P25's streaming, caller-retained-entry and bounded-index alternatives +with the same selection policy and exact output oracles. The existing +negotiation target already measures fixed three-offer selection, retained +entry scalar reads and select-and-forward; the Location target already +retains the caller's parsed URI. Extend their missing axes rather than +inventing a repeatedly reparsing comparison arm. Include decode-only, +owned/predecoded-owned reads and the complete decode/read/forward lifecycle +for Location, Host, Origin, Accept/Encoding/Language, Allow and Vary. + +Make the repeated-read multiplier observable. In the policy target, +`referrer_preferred_8`, `cache_max_age_8` and `auth_access_8` XOR the same pure +answer eight times and black-box only the final zero. This permits hoisting +or cancellation; it does not prove eight calls survive optimization. Guard +the per-read inputs/results and check B5's emitted operation boundary before +using those cases to price getter cost. + +- `crates/http_headers/benches/http_headers_operations.rs:44` — generic owned operation ends in `consume(header)` +- `crates/http_headers/benches/http_headers_operations.rs:52` — borrowed operation does the same +- `crates/http_headers/benches/http_headers_per_header.rs:221` — negotiation registrations use these operations +- `crates/http_headers/benches/http_headers_micro.rs:663` — existing Content-Type semantic cases are a starting point, not complete coverage +- `crates/http_headers/benches/http_headers_policy.rs:127` — `result ^= value.preferred().unwrap() as usize`; the cached-age and credential loops repeat this pattern at lines 143 and 164 +- `crates/http_headers/benches/http_headers_negotiation_semantics.rs:463` — `forward` and its caller-parsed comparison both start from current owned decoding at line 467 +- `crates/http_headers/benches/http_headers_location_semantics.rs:100` — caller-side URI retention in `caller_reads`, parsed once at line 102 + +**Done when:** cases report wall time, instructions and allocations separately +for decode-only and decode-plus-read, consume actual semantic values, and are +wired through B4 to fail on unexpected allocation increases or any instruction +increase against per-case baselines. + +This unblocks P9/P11/P12/P15/P24/P25/P26 independently of decode-only +microbenchmarks. Reader counts and allocations must be recorded, not inferred +from loop syntax or target names; reject a comparison whose two arms perform +different policy or ownership work. SEC3 precedes trusting sensitivity +preservation in forwarding controls. + +--- + + +### B2 — Measure source representations and ownership boundaries + +**Area:** source acquisition and raw conversion benchmarks · **Priority:** High · **Effort:** Medium + +**Purpose:** maintain and improve. Current per-header setup uses HTTP storage, +which is exempt from custom budgets and can share owned data. Add equivalent +`Single`, `Borrowed`, stored-`FieldValue` slice and HTTP cases. Cover 1/2/3/4/5/ +16/128 lines, 63/64/65-byte values, totals around 64 KiB and 1,024/1,025 items. +Measure canonical/noncanonical Accept-Ranges and repeated collectors explicitly. +Separately compare CORS `TryFrom`, `FromStr`, vector-taking conversion +and source `owned` without counting caller setup as parser allocation. + +The five shape targets already compare HTTP and borrowed raw sources; retain +those arms and fill the `Single`/stored-slice/size-limit intersections. +Compare default `FieldSink` and HTTP insertion, direct typed construction, +ownership transfer, cloning and destruction as separate operations and as a +complete lifecycle. Include Location/Host/Origin component construction, +Accept-family typed entries, and P27's Allow/Vary construction across the +inline boundary. For P26, distinguish the metadata box from wire storage. +For P21, add serde well-known/custom name deserialization and typed +deserialization with caller setup excluded. + +Keep setup and teardown explicit when measuring the shared storage operations +and bounded name corpora. Avoid substituting an artificial decode-and-leak +workload for a consumer that eventually drops its results. + +- `crates/http_headers/benches/http_headers_per_header.rs:32` — `fn map(...) -> HeaderMap` +- `crates/http_headers/tests/common/http_headers_storage_operations.rs:38` — shared storage operations return owned maps for post-measurement teardown +- `crates/http_headers/tests/common/http_headers_name_corpus.rs:76` — reusable name corpus setup; HTTP and crate-name corpora are at lines 86 and 97 +- `crates/http_headers/src/source/field_lines.rs:273` — representation-aware singleton acquisition +- `crates/http_headers/src/source/field_lines.rs:352` — repeated acquisition +- `crates/http_headers/src/headers/range/accept_ranges.rs:399` — ownership probes +- `crates/http_headers/benches/http_headers_shapes_common.rs:49` — fixture stores both HTTP and raw sources +- `crates/http_headers/src/serde_impls.rs:482` — temporary owned name before recognition; typed owned reconstruction is at line 594 +- `crates/http_headers/src/headers/negotiation/allow.rs:86` — typed constructor whose final buffer transfer needs a measured-body case; Vary's corresponding constructor is at `crates/http_headers/src/headers/negotiation/vary.rs:111` + +**Done when:** instruction, wall-time, allocation-count and allocated-byte +baselines cover validation and ownership transitions and preserve +sharing/singleton properties; expected limit errors remain distinct from +HTTP-exempt outcomes. Maintain cases fail through B4 on allocation increases +or any instruction increase. Include unknown size hints rather than testing +only exact slices. + +These cases unblock P21/P24/P26/P27 and protect SEC3 marker propagation, +wire order, sharing and fail-before-mutation behavior. Also gate unexpected +allocated-byte increases through B4; a lower allocation count alone can hide +a larger retained representation. Measure peak/live bytes with deliberate +setup/drop boundaries, without presenting that as an observed production +memory bottleneck. + +--- + + +### B3 — Cover nonliteral grammar and fallback distributions + +**Area:** parser success/fallback corpus · **Priority:** High · **Effort:** Large + +**Purpose:** maintain and improve. Add a stratified corpus rather than only +whole-line literals that coincide with fast paths. Measure wall time, +instructions and allocations, reporting each stratum without inventing +production frequencies. This supplies evidence for fallback paths and input +shapes that the common-case fixtures do not exercise. + +- `crates/http_headers/benches/http_headers_per_header.rs:331` — fixed Content-Length example `"348"` +- `crates/http_headers/benches/http_headers_per_header.rs:355` — Host is `example.com:8443` +- `crates/http_headers/benches/http_headers_micro.rs:69` — one three-parameter Content-Type fixture +- `crates/http_headers_simd/benches/http_headers_simd_no_std_dispatch.rs:51` — existing crossover-length scanner cases +- `crates/http_headers_simd/src/api.rs:16` — general UTF-8 validation has alignment-dependent work + +Include numeric 19/20-digit and leading-zero inputs, equal/conflicting +duplicates, long/obs-text tags, all date formats, range unit/casing/order forms, +and errors competing with malformed delimiters. Cover Content-Type common and +nonliteral forms, 0/1/2/3/16/49 parameters, first/last/missing/duplicate lookups +and compact-offset boundaries. Exercise late quoted Accept/WebSocket fallback, +weighted relaxed fractions, canonical/noncanonical HSTS, CORS wildcard/mixed-case +tokens and IPv6, Host percent/IPv6/IPvFuture/IDNA, and relaxed Location +backslashes. For Authorization, include scheme spacing, early/late decoded +colon, padding, fresh/warm extraction and success/failure reuse; retain all +zeroization in the measured lifecycle. Include large-to-small credentials and +lengths around the 64-KiB retention cap, recording retained capacity as well as +allocations. Scanner cases cover every tail length, invalid-byte position and +supported backend. Literal-recognition cases pair hits with equal-length misses +and casing variants. + +Add explicit slice-offset strata for ASCII/UTF-8 projection and vector +scanners, spanning at least one vector width. Record the input alignment and +compare matching offsets: byte-identical static fixtures can move when +unrelated code changes. Separate alignment effects from parser work through +raw-profile attribution rather than averaging away a slower case. + +Reuse the substantial existing `*_shapes` corpora: they already include +nonliteral, repeated, relaxed, absent and malformed HTTP/raw cases. The +outstanding work is the remaining boundary/reader/backend intersections +and trustworthy current baselines, not recreating those fixtures. +P25 needs strict and exact long-fraction qualities, duplicate policy, +quoted parameters/extensions and inline-index overflow; P26 needs literal +controls as well as every general metadata form. + +Separate cold credential extraction from a genuinely prewarmed measured +call. `crates/http_headers/benches/http_headers_policy.rs:184` calls +`value.extract(&mut credentials)` before another extraction at line 185 +inside the same measured function, so its current warm label includes both. +Keep a combined lifecycle case too, explicitly named, with zeroization. + +Record actual backend selection and compiler target flags. The companion's +`http_headers_simd_no_std_dispatch` name does not ensure `no_std`: the +configured all-features bench build and facade feature unification enable +`std`, while `.cargo/config.toml:3` enables x86-64-v3. Add isolated performance +strata for supported generic x86/std/no_std, available accelerated backends, +and AArch64, without executing unsupported instructions. This protects +threshold choices at `crates/http_headers_simd/src/dispatch.rs:17` and +line 30, including Authorization's separate local threshold. The isolated +no_std correctness checks do not replace these crossover measurements; +the no_std-only coverage/mutation exemptions remain unchanged. + +**Done when:** value and exact error-kind/index assertions cover these strata, +current baselines are recorded, and B4 gates selected maintain cases at any +unexpected allocation increase or any instruction increase. + +These axes unblock P1/P6/P9/P10/P11/P12/P13/P14/P18/P19/P24/P25/P26. +Also apply B4's allocated-byte gate. Retain observed/backend/configuration +provenance so compile-only success cannot be presented as measured fallback +performance or evidence of a vectorized ASCII helper. + +**Guideline connection:** [M-UNSAFE](https://microsoft.github.io/rust-guidelines/guidelines/correctness/#M-UNSAFE) +says performance-motivated unsafe should follow benchmarking. These +backend/fallback strata and B5's consumer profiles supply that evidence; +historical threshold comments alone do not. CON3 separately owns the +unassessed current Miri acceptance requirement. + +--- + + +### B4 — Gate parser instruction and allocation regressions + +**Area:** benchmark execution and regression enforcement · **Priority:** High · **Effort:** Medium + +**Purpose:** maintain. The generated benchmark check only compiles benchmarks; +it does not notice slower code. Wire a bounded, representative parser suite to +execute instruction/allocation comparisons against a current recorded baseline. +Preserve generated Anvil ownership: express supported configuration upstream +or add justified repository-specific performance automation, not a hand edit +to the generated recipe. + +- `justfiles/anvil/checks/bench.just:17` — `cargo ... bench ... --all-features --no-run` +- `justfiles/anvil/groups/scheduled-exhaustive.just:22` — scheduled tier invokes that compile-only recipe +- `crates/http_headers/docs/PERF.md:3` — timing table is explicitly an imported snapshot + +Protect short borrowed/owned singleton decode, retained shared HTTP storage, +all header-family semantic readers, custom-source validation and SIMD +crossovers. Proposed deterministic gates are any unexpected allocation-count +increase and any instruction increase per fixed case. Compare identical inputs, +compiler versions, target flags and operation boundaries; never offset a +regression with gains in other cases. Calibrate execution cost and wall-time +repeatability through B6 rather than claiming hosted timing is stable. + +Include allocated bytes in the zero-increase gate, not only allocation +count. Version the exact operation identity, consumed semantic work, +input/alignment, ownership and setup/drop boundaries, compiler, profile, +features, target flags and baseline artifacts together. Add B1's repeated-read +barriers before protecting a nominal eight-read workload. + +The report script is not a hidden regression gate: +`crates/http_headers/scripts/perf_report.rs:287` selects `--no-baseline`; +its `--check` mode compares rendered documentation with an existing artifact, +not old/new parser costs. The configured Anvil scheduled path remains +compile-only, and no successful runtime CI comparison has been observed +in this audit. Keep both required-check fan-in display names stable. + +**Done when:** an intentional regression fails the chosen automated entry point, +baseline/toolchain/hardware provenance and update rules are documented, and B1–B3 +maintain cases actually execute. A reporting-only job does not close this item. + +--- + + +### B5 — Measure downstream release code generation and layouts + +**Area:** consumer-profile instruction and layout evidence · **Priority:** Medium · **Effort:** Medium + +**Purpose:** improve, then maintain accepted wins. Compare a downstream caller +under ordinary release settings with the current fat-LTO benchmark profile. +Black-box input and semantic results. Inspect retained calls, indirect dispatch, +bounds/overflow checks, aggregate copies, type sizes and emitted error/panic +machinery; include retained Location/Host/origin metadata and Content-Type +layout controls. Account for monomorphized build +cost, code-size growth and hot-field/cache-line layout +instead of assuming all inlining is beneficial. +Include isolated before/after code-size attribution for bounded IPv6 origin +ASCII projection; whole-binary size differences containing other edits do not +establish that helper's footprint. + +Expand the retained representation controls to P25's bounded member index +and P26's inline/general metadata, with decode-only and semantic-reader +consumers so a larger result is not hidden by LTO. Inspect actual record +sizes/offsets, `Result` movement, surviving bounds checks, indirect calls, +panic paths and mono-item/code-size attribution. Do not infer any of these +from a Rust field order, missing inline annotation, or a comment claiming +vectorization. + +For the conditional recommendations in +[M-AVOID-INDIRECTION](https://microsoft.github.io/rust-guidelines/guidelines/performance/#M-AVOID-INDIRECTION), +[M-BOX-DST](https://microsoft.github.io/rust-guidelines/guidelines/performance/#M-BOX-DST) +and [M-SHRINK-TO-FIT](https://microsoft.github.io/rust-guidelines/guidelines/performance/#M-SHRINK-TO-FIT), +establish applicability before choosing a representation. Record whether +private sequences are immutable, frequently instantiated, large or +long-lived, and their retained capacity slack. Compare reduced indirection +or retained bytes against larger values, construction copies and allocation +costs. Reusable credential buffers and caller-visible mutable collections +are not blanket candidates for boxing or shrinking. + +Record clean/incremental consumer build time and code size for minimal, +selected-family, default, HTTP and serde features. The negotiation feature +selects IDNA compiled data at `crates/http_headers/Cargo.toml:73`; that is a +real configuration, not proof of unnecessary dependency cost. Compare a +generic supported CPU with repository x86-64-v3 and relevant AArch64 output. +Use isolated builds where feature unification would obscure the target +configuration; do not change the public defaults just to improve a benchmark. + +- `Cargo.toml:423` — release config specifies debug information, not fat LTO +- `Cargo.toml:428` — bench enables fat LTO and one codegen unit +- `crates/http_headers/src/headers/cors/shared.rs:911` — erased validator boundary +- `crates/http_headers/src/headers/content_type.rs:158` — boxed/general metadata tradeoff +- `crates/http_headers/src/headers/negotiation/host.rs:649` — retained host and port metadata +- `crates/http_headers/src/headers/cors/access_control_allow_origin.rs:942` — bounded IPv6 validation and ASCII projection + +**Done when:** each codegen hypothesis has actual layout/disassembly and +instruction/time/code-size evidence supporting acceptance or rejection. +Report-only experiments satisfy the improve portion; retained improvements +must gain B4 maintain cases with unchanged allocation expectations and the +zero-regression instruction gate. No unchecked access or removed invariant is an acceptable +substitute for compiler-visible proof. + +This is the explicit unblocker for currently unknown retained-call, +vectorization, bounds-check, panic-path, field-layout, monomorphization and +build-footprint questions. B3/B6 supply the relevant inputs and workload +shares; do not turn an emitted-code difference alone into an overall +performance claim. + +--- + + +### B6 — Establish workload shares and benchmark repeatability + +**Area:** request-shaped parsing and measurement calibration · **Priority:** Medium · **Effort:** Medium + +**Purpose:** improve. Existing per-header microbenchmarks cannot establish what +fraction of a real request each parser consumes, its input frequencies, or +tail-latency effects. Assemble a non-network request-shaped decode/read bundle +using explicitly sourced, sanitized workload distributions; until those are +available, report separate JSON-request, browser-negotiation and response-policy +strata without claiming a representative weighted average. Vary absent fields, +wire lengths, repeated lines, reader counts and borrowed/owned retention. + +- `crates/http_headers/benches/http_headers_micro.rs:87` — existing request fixtures can seed separate strata +- `crates/http_headers/benches/http_headers_per_header.rs:47` — 60-sample Criterion configuration does not establish current variance +- `crates/http_headers/docs/PERF.md:8` — imported per-header timing/instruction/allocation dimensions are not end-to-end shares + +Measure total bundle time, per-parser contribution, allocation traffic and +latency distribution; repeat across relevant architectures/profiles. Establish +run-to-run noise and suite duration before selecting timed gates or CI cadence. +This resolves currently unassessed workload shares, tail behavior, +instruction-cache tradeoffs, variance and execution-budget practicality. + +Keep attribution additive: measure total bundle cost and identified parser, +source, semantic-reader, construction and sink components against equivalent +work, while retaining the executable paths and raw profiles. Price the +reader-count crossover for P25 and common/nonliteral fractions for P26; +rank P21/P24/P27 only after their actual operation frequencies are known. +Current fixtures cannot supply those probabilities. Synthetic strata remain +report-only improve baselines, not a deployed latency SLO or weighted win. + +Make B1's semantic reads observable before collecting repeated runs. Record +setup, drop, cold/warm state and which allocation instrument is actually installed. +The companion's opt-in tracker at +`crates/http_headers_simd/src/tracking.rs:83` uses coherent process-wide +counters; its existence does not establish that every metabench target uses +it. Compare timing and allocation-counting modes separately rather than +attributing harness synchronization to production parsing. + +**Done when:** reporting baselines bound plausible end-to-end benefits for +parser optimizations, record distribution provenance and repeatability, and +determine a practical B4 suite/cadence. This is an improve-only measurement: a report is +sufficient here, but it does not replace B4's regression gate. + +**Guideline obligations:** [M-HOTPATH](https://microsoft.github.io/rust-guidelines/guidelines/performance/#M-HOTPATH) +recommends identifying, benchmarking and profiling performance-relevant +paths; [M-THROUGHPUT](https://microsoft.github.io/rust-guidelines/guidelines/performance/#M-THROUGHPUT) +recommends throughput and items-per-CPU-cycle evidence. Current compliance +with those measurement recommendations is unassessed, not established by +configured targets. Preserve revision, executable, compiler, target/ISA, +feature, workload and raw-profile provenance, including the retention +distribution needed for B5's conditional representation comparisons. + +--- + +## Testing + + +### T1 — Make counter-lock rendezvous tolerate spurious failure and abort safely + +**Area:** `http_headers_simd` allocation-counter concurrency tests · **Priority:** Medium · **Effort:** Small + +**Gap type:** Weak · **Would catch:** a broken retry loop that gives up after a spurious weak-CAS failure, without rejecting a correct retry · **Scope:** 1 rendezvous test and its attempt hook, exhaustive · **Blocks / blocked by:** none + +Make the test enforce the lock contract rather than guaranteed success of a +particular weak compare-exchange. The production loop correctly retries. +After releasing the guard, however, the controller asserts that the very +next attempt acquired the lock. A permitted spurious failure produces +`false`, so correct production code can fail this assertion. The worker +then waits for another resume message; unwinding through `thread::scope` +waits for that worker while the controller's channel endpoints remain alive. +This can also strand test cleanup and the shared test mutex. This is a +static permitted schedule, not an observed flake or production locking bug. + +- `crates/http_headers_simd/src/tracking.rs:43` — retry loop uses `compare_exchange_weak` at line 47; the hook only observes its result +- `crates/http_headers_simd/src/tracking.rs:417` — nearest and affected test creates three zero-capacity channels +- `crates/http_headers_simd/src/tracking.rs:427` — worker reports an attempt, then blocks awaiting resume +- `crates/http_headers_simd/src/tracking.rs:440` — controller requires immediate post-unlock success before sending the final resume +- `crates/http_headers_simd/src/tracking.rs:292` — the other concurrency test checks aggregate accounting, not this retry/abort schedule + +Keep this a deterministic unit test. Exercise an injected spurious failure +after unlock followed by success, and make controller failure disconnect or +cancel the worker before joining it. Preserve the assertion that acquisition +cannot complete while the original guard is held. The legal extra failure +is observably different from giving up or acquiring early; neither should be +hidden by weakening assertions or retrying the whole test. + +**Done when:** the controlled spurious-failure case completes with exactly one +successful acquisition after release, and fails if the retry is removed or +acquisition is allowed while the first guard is held. A deliberately aborted +controller also releases the worker and test lock with bounded, observable +completion instead of blocking scoped-thread cleanup. Correct production +weak-CAS behavior must not produce a test failure. diff --git a/crates/http_headers/examples/axum.rs b/crates/http_headers/examples/axum.rs new file mode 100644 index 000000000..3304227e8 --- /dev/null +++ b/crates/http_headers/examples/axum.rs @@ -0,0 +1,31 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Axum server using typed request and response headers. + +use std::error::Error; + +use tokio::net::TcpListener; + +#[path = "axum/app.rs"] +mod app; + +use app::router; + +#[tokio::main] +async fn main() -> Result<(), Box> { + let app = router(); + let testing = std::env::var_os("IS_TESTING").is_some(); + let address = if testing { "127.0.0.1:0" } else { "127.0.0.1:3000" }; + let listener = TcpListener::bind(address).await?; + println!("listening on http://{}", listener.local_addr()?); + + let server = axum::serve(listener, app); + if testing { + server.with_graceful_shutdown(async {}).await?; + } else { + server.await?; + } + + Ok(()) +} diff --git a/crates/http_headers/examples/axum/app.rs b/crates/http_headers/examples/axum/app.rs new file mode 100644 index 000000000..40a89957f --- /dev/null +++ b/crates/http_headers/examples/axum/app.rs @@ -0,0 +1,46 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::time::Duration; + +use axum::http::{HeaderMap, StatusCode}; +use axum::response::{IntoResponse, Response}; +use axum::routing::get; +use axum::{Json, Router}; +use http_headers::headers::{CacheControl, UserAgent}; +use http_headers::sink::FieldSinkExt; +use serde::Serialize; + +#[derive(Serialize)] +struct Greeting { + message: &'static str, + user_agent: String, +} + +type HandlerError = (StatusCode, String); + +async fn greeting(headers: HeaderMap) -> Result { + let user_agent = UserAgent::view(&headers) + .map_err(|error| (StatusCode::BAD_REQUEST, error.to_string()))? + .map(|value| value.as_str().map(str::to_owned)) + .transpose() + .map_err(|error| (StatusCode::BAD_REQUEST, error.to_string()))? + .unwrap_or_else(|| "unknown".to_owned()); + + let mut response = Json(Greeting { + message: "hello from http_headers", + user_agent, + }) + .into_response(); + + response + .headers_mut() + .set_cache_control(CacheControl::private().max_age(Duration::from_mins(1))) + .map_err(|error| (StatusCode::INTERNAL_SERVER_ERROR, error.to_string()))?; + + Ok(response) +} + +pub(crate) fn router() -> Router { + Router::new().route("/", get(greeting)) +} diff --git a/crates/http_headers/favicon.ico b/crates/http_headers/favicon.ico new file mode 100644 index 000000000..5d7bd15c9 --- /dev/null +++ b/crates/http_headers/favicon.ico @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3d2a6be9c244b42877fa86aaef97468175d8ecf64066e0bb6d48ef912e5261ba +size 432254 diff --git a/crates/http_headers/logo.png b/crates/http_headers/logo.png new file mode 100644 index 000000000..1523b6f74 --- /dev/null +++ b/crates/http_headers/logo.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b8fb949b8c6797dfa3489b4b3306977ed6d98f8c4da42a25171709c9c1285497 +size 114089 diff --git a/crates/http_headers/scripts/perf_report.rs b/crates/http_headers/scripts/perf_report.rs new file mode 100755 index 000000000..ca77f2f74 --- /dev/null +++ b/crates/http_headers/scripts/perf_report.rs @@ -0,0 +1,505 @@ +#!/usr/bin/env -S cargo +nightly -Zscript +--- +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +[package] +edition = "2024" + +[dependencies] +serde = { version = "1", features = ["derive"] } +serde_json = "1" +--- + +//! Render the per-header metabench results in `docs/PERF.md`. + +use std::collections::{BTreeMap, BTreeSet}; +use std::fmt::Write as _; +use std::path::{Path, PathBuf}; +use std::process::{Command, ExitCode}; +use std::{env, fs}; + +use serde::Deserialize; + +const SCHEMA_VERSION: u32 = 8; +const DEFAULT_ARTIFACT: &str = "target/metabench/http_headers_per_header/report.json"; +const DEFAULT_REPORT: &str = "crates/http_headers/docs/PERF.md"; + +const ROWS: &[(&str, &str, bool)] = &[ + ("accept", "Accept", false), + ("accept_encoding", "Accept-Encoding", false), + ("accept_language", "Accept-Language", false), + ("accept_ranges", "Accept-Ranges", true), + ( + "access_control_allow_credentials", + "Access-Control-Allow-Credentials", + true, + ), + ( + "access_control_allow_headers", + "Access-Control-Allow-Headers", + true, + ), + ( + "access_control_allow_methods", + "Access-Control-Allow-Methods", + true, + ), + ( + "access_control_allow_origin", + "Access-Control-Allow-Origin", + true, + ), + ( + "access_control_expose_headers", + "Access-Control-Expose-Headers", + true, + ), + ("access_control_max_age", "Access-Control-Max-Age", true), + ( + "access_control_request_headers", + "Access-Control-Request-Headers", + true, + ), + ( + "access_control_request_method", + "Access-Control-Request-Method", + true, + ), + ("allow", "Allow", true), + ("authorization_basic", "Authorization (Basic)", true), + ("authorization_bearer", "Authorization (Bearer)", true), + ("cache_control", "Cache-Control", true), + ("content_length", "Content-Length", true), + ("content_range", "Content-Range", true), + ("content_security_policy", "Content-Security-Policy", false), + ("content_type", "Content-Type", true), + ("etag", "ETag", true), + ("host", "Host", true), + ("if_match", "If-Match", true), + ("if_modified_since", "If-Modified-Since", true), + ("if_none_match", "If-None-Match", true), + ("if_range", "If-Range", true), + ("if_unmodified_since", "If-Unmodified-Since", true), + ("last_modified", "Last-Modified", true), + ("location", "Location", true), + ("range", "Range", true), + ("referrer_policy", "Referrer-Policy", true), + ("sec_websocket_accept", "Sec-WebSocket-Accept", true), + ( + "sec_websocket_extensions", + "Sec-WebSocket-Extensions", + false, + ), + ("sec_websocket_key", "Sec-WebSocket-Key", true), + ("sec_websocket_protocol", "Sec-WebSocket-Protocol", false), + ("sec_websocket_version", "Sec-WebSocket-Version", true), + ("server", "Server", true), + ("set_cookie", "Set-Cookie", true), + ( + "strict_transport_security", + "Strict-Transport-Security", + true, + ), + ("user_agent", "User-Agent", true), + ("vary", "Vary", true), + ("x_content_type_options", "X-Content-Type-Options", false), +]; + +/// Benchmarks the harness measures but the report omits, because they cover a +/// secondary shape of a header the report already lists. +const UNREPORTED: &[&str] = &["sec_websocket_version_advertisement"]; + +#[derive(Deserialize)] +struct Report { + schema_version: u32, + entries: Vec, +} + +#[derive(Deserialize)] +struct Entry { + identity: String, + results: BTreeMap, +} + +#[derive(Deserialize)] +struct EngineResult { + metrics: BTreeMap, +} + +#[derive(Deserialize)] +struct Metric { + value: MetricValue, + unit: Option, +} + +#[derive(Clone, Copy, Deserialize)] +#[serde(untagged)] +enum MetricValue { + Integer(u64), + Float(f64), +} + +#[derive(Clone, Copy)] +struct Measurements { + time_ns: f64, + instructions: u64, + allocations: u64, + allocated_bytes: u64, +} + +struct Arguments { + artifact: PathBuf, + report: PathBuf, + no_run: bool, + check: bool, +} + +fn main() -> ExitCode { + if env::args() + .skip(1) + .any(|argument| matches!(argument.as_str(), "--help" | "-h")) + { + println!("{}", usage()); + return ExitCode::SUCCESS; + } + match run() { + Ok(message) => { + println!("{message}"); + ExitCode::SUCCESS + } + Err(error) => { + eprintln!("perf_report: {error}"); + ExitCode::FAILURE + } + } +} + +fn run() -> Result { + let root = workspace_root()?; + let arguments = arguments(&root)?; + if !arguments.no_run && !arguments.check { + collect(&root, &arguments.artifact)?; + } + let report = render(&read_report(&arguments.artifact)?)?; + if arguments.check { + let existing = fs::read_to_string(&arguments.report) + .map_err(|error| format!("cannot read {}: {error}", arguments.report.display()))?; + if existing != report { + return Err(format!( + "{} differs from the metabench report", + arguments.report.display() + )); + } + return Ok(format!("{} is up to date", arguments.report.display())); + } + if let Some(parent) = arguments.report.parent() { + fs::create_dir_all(parent) + .map_err(|error| format!("cannot create {}: {error}", parent.display()))?; + } + fs::write(&arguments.report, report) + .map_err(|error| format!("cannot write {}: {error}", arguments.report.display()))?; + Ok(format!( + "wrote {} with {} header rows", + arguments.report.display(), + ROWS.len() + )) +} + +fn arguments(root: &Path) -> Result { + let mut artifact = root.join(DEFAULT_ARTIFACT); + let mut report = root.join(DEFAULT_REPORT); + let mut no_run = false; + let mut check = false; + let mut values = env::args().skip(1); + while let Some(value) = values.next() { + match value.as_str() { + "--artifact" => { + artifact = absolute(root, values.next().ok_or("missing artifact path")?) + } + "--output" => report = absolute(root, values.next().ok_or("missing output path")?), + "--no-run" => no_run = true, + "--check" => { + no_run = true; + check = true; + } + other => return Err(format!("unknown argument `{other}`\n{}", usage())), + } + } + Ok(Arguments { + artifact, + report, + no_run, + check, + }) +} + +fn usage() -> &'static str { + "usage: perf_report.rs [--no-run|--check] [--artifact PATH] [--output PATH]" +} + +fn absolute(root: &Path, value: String) -> PathBuf { + let path = PathBuf::from(value); + if path.is_absolute() { + path + } else { + root.join(path) + } +} + +fn workspace_root() -> Result { + let script = env::args() + .next() + .ok_or("missing script path")? + .parse::() + .map_err(|error| error.to_string())? + .canonicalize() + .map_err(|error| format!("cannot resolve script path: {error}"))?; + script + .parent() + .and_then(Path::parent) + .and_then(Path::parent) + .and_then(Path::parent) + .map(Path::to_path_buf) + .ok_or_else(|| "cannot locate workspace root".to_string()) +} + +fn collect(root: &Path, artifact: &Path) -> Result<(), String> { + if let Some(parent) = artifact.parent() { + fs::create_dir_all(parent) + .map_err(|error| format!("cannot create {}: {error}", parent.display()))?; + } + let status = Command::new("cargo") + .current_dir(root) + .args([ + "bench", + "-p", + "http_headers", + "--features", + "benchmarking,http", + "--bench", + "http_headers_per_header", + "--", + "--criterion", + "--gungraun", + "--allocations", + "--show-engine-output", + "--no-baseline", + "--export-json", + ]) + .arg(artifact) + .status() + .map_err(|error| format!("cannot run metabench: {error}"))?; + if status.success() { + Ok(()) + } else { + Err(format!("metabench failed with {status}")) + } +} + +fn read_report(path: &Path) -> Result { + let text = fs::read_to_string(path) + .map_err(|error| format!("cannot read {}: {error}", path.display()))?; + let report: Report = serde_json::from_str(&text) + .map_err(|error| format!("cannot parse {}: {error}", path.display()))?; + if report.schema_version != SCHEMA_VERSION { + return Err(format!( + "expected metabench schema {SCHEMA_VERSION}, got {}", + report.schema_version + )); + } + Ok(report) +} + +fn render(report: &Report) -> Result { + let entries = report + .entries + .iter() + .map(|entry| (entry.identity.as_str(), entry)) + .collect::>(); + if entries.len() != report.entries.len() { + return Err("metabench report contains duplicate identities".to_string()); + } + let expected = ROWS + .iter() + .flat_map(|(id, _, supported)| { + let mut identities = vec![ + format!("http_headers_per_header/per_header/http_headers_owned/{id}"), + format!("http_headers_per_header/per_header/http_headers_borrowed/{id}"), + ]; + if *supported { + identities.push(format!("http_headers_per_header/per_header/headers/{id}")); + } + identities + }) + .collect::>(); + let unreported = UNREPORTED + .iter() + .flat_map(|id| { + [ + format!("http_headers_per_header/per_header/http_headers_owned/{id}"), + format!("http_headers_per_header/per_header/http_headers_borrowed/{id}"), + format!("http_headers_per_header/per_header/headers/{id}"), + ] + }) + .collect::>(); + let actual = entries + .keys() + .map(|identity| (*identity).to_owned()) + .filter(|identity| !unreported.contains(identity)) + .collect(); + if actual != expected { + let missing = expected.difference(&actual).cloned().collect::>(); + let extra = actual.difference(&expected).cloned().collect::>(); + return Err(format!( + "metabench inventory mismatch; missing: {missing:?}; extra: {extra:?}" + )); + } + + let mut output = String::with_capacity(16 * 1024); + writeln!(output, "# Performance\n") + .map_err(|error| format!("cannot render report heading: {error}"))?; + writeln!( + output, + "This table is the comparative snapshot imported with `http_headers` 0.1.0. It\n\ + is not a promise of current timing on different hardware or dependency\n\ + versions; use the repository's metabench targets to measure the current\n\ + checkout.\n" + ) + .map_err(|error| format!("cannot render report disclaimer: {error}"))?; + writeln!( + output, + "Each cell reports metabench's median wall-clock time, Callgrind instruction\n\ + count, allocation count, and total allocated bytes for one typed decode and\n\ + read. Benchmarks consume the decoded value and force comparable semantic work\n\ + when `headers 0.4.1` defers parsing. Setup uses a prebuilt `HeaderMap` outside\n\ + the measured operation. Criterion uses a 1 s warm-up, 3 s measurement period,\n\ + and 60 samples per arm. Each cell is formatted as time, instructions, then\n\ + allocation count / allocated bytes. `n/a` means `headers 0.4.1` does not\n\ + provide that typed header.\n" + ) + .map_err(|error| format!("cannot render report description: {error}"))?; + writeln!( + output, + "| Header | `headers 0.4.1` | `http_headers` (owned) | `http_headers` (borrowed) |" + ) + .and_then(|()| writeln!(output, "|---|---:|---:|---:|")) + .map_err(|error| format!("cannot render table heading: {error}"))?; + + for (id, header, supported) in ROWS { + let theirs = if *supported { + format_measurements(measurements(required_entry(&entries, id, "headers")?)?) + } else { + "n/a".to_string() + }; + let owned = format_measurements(measurements(required_entry( + &entries, + id, + "http_headers_owned", + )?)?); + let borrowed = format_measurements(measurements(required_entry( + &entries, + id, + "http_headers_borrowed", + )?)?); + writeln!(output, "| {header} | {theirs} | {owned} | {borrowed} |") + .map_err(|error| format!("cannot render `{id}`: {error}"))?; + } + Ok(output) +} + +fn required_entry<'a>( + entries: &'a BTreeMap<&str, &Entry>, + group: &str, + benchmark: &str, +) -> Result<&'a Entry, String> { + let identity = format!("http_headers_per_header/per_header/{benchmark}/{group}"); + entries + .get(identity.as_str()) + .copied() + .ok_or_else(|| format!("missing benchmark `{identity}`")) +} + +fn measurements(entry: &Entry) -> Result { + let time = metric(entry, "criterion", "median")?; + if time.unit.as_deref() != Some("ns") { + return Err(format!( + "`{}` criterion median has unit {}, expected ns", + entry.identity, + time.unit.as_deref().unwrap_or("") + )); + } + Ok(Measurements { + time_ns: finite_nonnegative(entry, "criterion/median", float_value(time.value))?, + instructions: integer_metric(entry, "gungraun.callgrind", "Ir")?, + allocations: integer_metric(entry, "alloc_tracker", "Allocations")?, + allocated_bytes: integer_metric(entry, "alloc_tracker", "Allocated bytes")?, + }) +} + +fn metric<'a>(entry: &'a Entry, engine: &str, name: &str) -> Result<&'a Metric, String> { + entry + .results + .get(engine) + .and_then(|result| result.metrics.get(name)) + .ok_or_else(|| format!("`{}` lacks `{engine}/{name}`", entry.identity)) +} + +fn integer_metric(entry: &Entry, engine: &str, name: &str) -> Result { + match metric(entry, engine, name)?.value { + MetricValue::Integer(value) => Ok(value), + MetricValue::Float(_) => Err(format!( + "`{}` has non-integer `{engine}/{name}` value", + entry.identity + )), + } +} + +fn float_value(value: MetricValue) -> f64 { + match value { + MetricValue::Integer(value) => value as f64, + MetricValue::Float(value) => value, + } +} + +fn finite_nonnegative(entry: &Entry, name: &str, value: f64) -> Result { + if value.is_finite() && value >= 0.0 { + Ok(value) + } else { + Err(format!("`{}` has invalid `{name}` value", entry.identity)) + } +} + +fn format_measurements(value: Measurements) -> String { + format!( + "{}
{} instr
{} allocs / {}", + time(value.time_ns), + grouped(value.instructions), + grouped(value.allocations), + bytes(value.allocated_bytes) + ) +} + +fn time(nanoseconds: f64) -> String { + if nanoseconds >= 1_000.0 { + format!("{:.2} µs", nanoseconds / 1_000.0) + } else { + format!("{nanoseconds:.1} ns") + } +} + +fn grouped(value: u64) -> String { + let digits = value.to_string(); + let mut output = String::with_capacity(digits.len() + digits.len() / 3); + for (index, byte) in digits.bytes().enumerate() { + if index != 0 && (digits.len() - index).is_multiple_of(3) { + output.push(','); + } + output.push(char::from(byte)); + } + output +} + +fn bytes(value: u64) -> String { + format!("{} B", grouped(value)) +} diff --git a/crates/http_headers/src/decode_error.rs b/crates/http_headers/src/decode_error.rs new file mode 100644 index 000000000..327a9818f --- /dev/null +++ b/crates/http_headers/src/decode_error.rs @@ -0,0 +1,206 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Structured decode errors shared by every header parser. + +use std::error::Error; +use std::{fmt, mem}; + +use crate::FieldName; + +/// The reason a header could not be decoded. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::DecodeErrorKind; +/// +/// assert_eq!(DecodeErrorKind::MissingValue.to_string(), "missing value"); +/// ``` +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +#[non_exhaustive] +pub enum DecodeErrorKind { + /// No field value was present. + MissingValue, + + /// A singleton header contained more than one field value. + UnexpectedMultipleValues, + + /// The field value did not match the header grammar. + InvalidSyntax, + + /// A component that requires UTF-8 contained other bytes. + InvalidUtf8, + + /// A token contained a byte forbidden by the HTTP token grammar. + InvalidToken, + + /// A numeric component was invalid or out of range. + InvalidNumber, + + /// A quoted value was not terminated. + UnterminatedQuote, + + /// A source or typed-construction byte, field-line, or list-item budget was exceeded. + /// + /// This is an admission limit, not a statement that the field bytes or + /// grammar are invalid. + SourceLimitExceeded, +} + +/// An error produced while decoding a header. +/// +/// Raw field values are deliberately omitted because headers can contain +/// credentials and other secrets. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::{DecodeError, DecodeErrorKind, FieldName}; +/// +/// let error = DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax); +/// assert_eq!(error.header(), &FieldName::ContentType); +/// assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); +/// ``` +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct DecodeError { + header: &'static FieldName, + value_index: u32, + kind: DecodeErrorKind, +} + +#[cfg(target_pointer_width = "64")] +const _: [(); 16] = [(); mem::size_of::()]; + +impl DecodeError { + const NO_VALUE_INDEX: u32 = u32::MAX; + + /// Creates an error for `header`. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::{DecodeError, DecodeErrorKind, FieldName}; + /// + /// let error = DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax); + /// assert_eq!(error.value_index(), None); + /// ``` + #[must_use] + pub const fn new(header: &'static FieldName, kind: DecodeErrorKind) -> Self { + Self { + header, + value_index: Self::NO_VALUE_INDEX, + kind, + } + } + + /// Attaches the zero-based field-value index at which decoding failed. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::{DecodeError, DecodeErrorKind, FieldName}; + /// + /// let error = + /// DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax).at_value(2); + /// assert_eq!(error.value_index(), Some(2)); + /// ``` + #[must_use] + #[expect( + clippy::cast_possible_truncation, + reason = "the preceding bound check proves the index fits in the compact representation" + )] + pub const fn at_value(mut self, value_index: usize) -> Self { + self.value_index = if value_index >= Self::NO_VALUE_INDEX as usize { + Self::NO_VALUE_INDEX - 1 + } else { + value_index as u32 + }; + self + } + + /// Returns the affected header name. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::{DecodeError, DecodeErrorKind, FieldName}; + /// + /// let error = DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax); + /// assert_eq!(error.header(), &FieldName::ContentType); + /// ``` + #[must_use] + pub const fn header(&self) -> &'static FieldName { + self.header + } + + /// Returns the zero-based field-value index, when known. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::{DecodeError, DecodeErrorKind, FieldName}; + /// + /// let error = DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax); + /// assert_eq!(error.value_index(), None); + /// assert_eq!(error.at_value(1).value_index(), Some(1)); + /// ``` + #[must_use] + pub const fn value_index(&self) -> Option { + if self.value_index == Self::NO_VALUE_INDEX { + None + } else { + Some(self.value_index as usize) + } + } + + /// Returns the structured failure reason. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::{DecodeError, DecodeErrorKind, FieldName}; + /// + /// let error = DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax); + /// assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + /// ``` + #[must_use] + pub const fn kind(&self) -> DecodeErrorKind { + self.kind + } +} + +impl fmt::Display for DecodeError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "invalid {} header: {}", self.header.as_str(), KindDisplay(self.kind))?; + if let Some(index) = self.value_index() { + write!(f, " at value {index}")?; + } + Ok(()) + } +} + +impl Error for DecodeError {} + +impl fmt::Display for DecodeErrorKind { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(match self { + Self::MissingValue => "missing value", + Self::UnexpectedMultipleValues => "unexpected multiple values", + Self::InvalidSyntax => "invalid syntax", + Self::InvalidUtf8 => "invalid UTF-8", + Self::InvalidToken => "invalid token", + Self::InvalidNumber => "invalid number", + Self::UnterminatedQuote => "unterminated quoted string", + Self::SourceLimitExceeded => "source limit exceeded", + }) + } +} + +struct KindDisplay(DecodeErrorKind); + +impl fmt::Display for KindDisplay { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(f) + } +} diff --git a/crates/http_headers/src/field.rs b/crates/http_headers/src/field.rs new file mode 100644 index 000000000..208bac05a --- /dev/null +++ b/crates/http_headers/src/field.rs @@ -0,0 +1,581 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Core traits connecting fields, borrowed views, and wire encoding. + +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef}; + +/// Controls how field values are validated. +/// +/// Strict decoding follows the field grammar exactly. Relaxed decoding is +/// opt-in and permits only documented, field-specific interoperability +/// deviations; fields without such deviations behave exactly as in strict +/// mode. +/// +/// Relaxed decoding accepts these interoperability deviations: +/// +/// - flexible quality-value whitespace and precision for `Accept`, +/// `Accept-Encoding`, and `Accept-Language`; +/// - lowercase `w/` entity-tag prefixes; +/// - optional whitespace around `Content-Type`, `Range`, and `Content-Range` +/// delimiters; +/// - `UTC`, outer whitespace, and one-digit components in IMF-style HTTP +/// dates; +/// - UTF-8 internationalized `Host` names that pass IDNA conversion; and +/// - backslashes normalized to slashes while validating `Location`. +/// +/// Original field bytes are preserved. Relaxed mode still enforces numeric +/// bounds, token and list structure, entity-tag contents, range ordering, +/// valid IDNA and ports, and every framing, authorization, CORS, security, and +/// WebSocket invariant. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::DecodeMode; +/// +/// assert_eq!(DecodeMode::default(), DecodeMode::Strict); +/// ``` +#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)] +pub enum DecodeMode { + /// Requires the field's specified grammar. + #[default] + Strict, + /// Permits the explicitly documented interoperability relaxations. + Relaxed, +} + +/// Reads and writes one typed HTTP field. +/// +/// Prefer [`Field::view`] when the result can borrow from the source. Use +/// [`Field::owned`] when the result must outlive the source borrow. +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(all(feature = "http", feature = "headers-user-agent"))] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::headers::{UserAgent, UserAgentOwned}; +/// +/// let mut map = HeaderMap::new(); +/// UserAgent::insert(&mut map, UserAgentOwned::try_from_static("client/1")?)?; +/// assert!(UserAgent::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(all(feature = "http", feature = "headers-user-agent")))] +/// # fn main() {} +/// ``` +pub trait Field: Sized + 'static { + /// The borrowed value type returned by [`Field::view`]. + type View<'a> + where + Self: 'a; + + /// The owned value type returned by [`Field::owned`]. + type Owned: 'static; + + /// Returns the field name. + /// + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "headers-user-agent")] + /// # fn main() { + /// use http_headers::headers::UserAgent; + /// use http_headers::{Field, FieldName}; + /// + /// assert_eq!(::name(), &FieldName::UserAgent); + /// # } + /// # #[cfg(not(feature = "headers-user-agent"))] + /// # fn main() {} + /// ``` + fn name() -> &'static FieldName; + + /// Reads a borrowed typed view from `source`. + /// + /// This is the preferred read operation when the result does not need to + /// outlive `source`. Returns `Ok(None)` when the field is absent. + /// + /// # Errors + /// + /// Returns an error when a present field is malformed. + #[expect( + clippy::inline_always, + reason = "the strict-mode convenience must disappear from typed decode hot paths" + )] + #[inline(always)] + fn view(source: &S) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + Self::view_with(source, DecodeMode::Strict) + } + + /// Reads a borrowed typed view under an explicit validation policy. + /// + /// Returns `Ok(None)` when the field is absent. + /// + /// # Errors + /// + /// Returns an error when a present field violates the selected policy. + fn view_with(source: &S, mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized; + + /// Reads an independently owned field from `source`. + /// + /// Prefer [`Field::view`] unless the result must be retained after + /// `source` is released. Returns `Ok(None)` when the field is absent. + /// + /// # Errors + /// + /// Returns an error when a present field is malformed. + #[expect( + clippy::inline_always, + reason = "the strict-mode convenience must disappear from typed decode hot paths" + )] + #[inline(always)] + fn owned(source: &S) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + Self::owned_with(source, DecodeMode::Strict) + } + + /// Reads an independently owned field under an explicit validation policy. + /// + /// Returns `Ok(None)` when the field is absent. + /// + /// # Errors + /// + /// Returns an error when a present field violates the selected policy. + fn owned_with(source: &S, mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized; + + /// Inserts a field, replacing every existing value with that name. + /// + /// # Errors + /// + /// Returns an error without changing the sink when it cannot hold the + /// encoded values. + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized; + + /// Removes every field value stored for this field. + #[inline] + fn remove(sink: &mut S) + where + S: FieldSink + ?Sized, + { + sink.remove_values(Self::name()); + } +} + +/// Defines a custom field represented by exactly one validated field value. +/// +/// Implementing this trait also implements [`Field`], including source +/// lookup, rejection of multiple field lines, strict and relaxed reads, +/// insertion, and removal. Applications using built-in headers normally use +/// [`Field`] instead. +/// +/// The blanket [`Field`] implementation enforces custom-source byte and line +/// budgets. It cannot infer whether a downstream-defined grammar is +/// list-valued, so such implementations must enforce +/// [`MAX_CUSTOM_LIST_ITEMS`](crate::source::MAX_CUSTOM_LIST_ITEMS) in both +/// borrowed and owned decoding before accepting further parsed items. +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(all(feature = "http", feature = "headers-user-agent"))] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::SingleValueField; +/// use http_headers::headers::{UserAgent, UserAgentOwned}; +/// +/// let mut map = HeaderMap::new(); +/// let value = UserAgentOwned::try_from_static("client/1")?; +/// UserAgent::insert(&mut map, value)?; +/// assert_eq!( +/// ::name().as_str(), +/// "user-agent" +/// ); +/// assert!( +/// UserAgent::view(&map) +/// .expect("stored header is valid") +/// .is_some() +/// ); +/// # Ok(()) +/// # } +/// # #[cfg(not(all(feature = "http", feature = "headers-user-agent")))] +/// # fn main() {} +/// ``` +pub trait SingleValueField: Sized + 'static { + /// The borrowed view type. + type View<'a> + where + Self: 'a; + + /// The owned value type. + type Owned: 'static; + + /// Returns the field name. + /// + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "headers-user-agent")] + /// # fn main() { + /// use http_headers::SingleValueField; + /// use http_headers::headers::UserAgent; + /// + /// assert_eq!( + /// ::name().as_str(), + /// "user-agent" + /// ); + /// # } + /// # #[cfg(not(feature = "headers-user-agent"))] + /// # fn main() {} + /// ``` + fn name() -> &'static FieldName; + + /// Validates and constructs a borrowed view. + /// + /// # Errors + /// + /// Returns an error when `value` violates this field's grammar. + /// + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "headers-user-agent")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http_headers::headers::UserAgent; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = ::decode_view(FieldValueRef::new(b"client/1"))?; + /// assert_eq!(view.as_str()?, "client/1"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "headers-user-agent"))] + /// # fn main() {} + /// ``` + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError>; + + /// Validates and constructs a borrowed view under an explicit policy. + /// + /// Fields without documented interoperability relaxations apply strict + /// validation in either mode. + /// + /// # Errors + /// + /// Returns an error when `value` violates the selected policy. + fn decode_view_with(value: FieldValueRef<'_>, _mode: DecodeMode) -> Result, DecodeError> { + Self::decode_view(value) + } + + /// Validates and constructs the owned value from one field value. + /// + /// # Errors + /// + /// Returns an error when `value` violates this field's grammar. + /// + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "headers-user-agent")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http_headers::headers::UserAgent; + /// use http_headers::{FieldValue, SingleValueField}; + /// + /// let owned = ::decode_owned(FieldValue::from_static("client/1"))?; + /// assert_eq!( + /// ::as_field_value(&owned), + /// "client/1" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "headers-user-agent"))] + /// # fn main() {} + /// ``` + fn decode_owned(value: FieldValue) -> Result; + + /// Validates and constructs an owned value under an explicit policy. + /// + /// Fields without documented interoperability relaxations apply strict + /// validation in either mode. + /// + /// # Errors + /// + /// Returns an error when `value` violates the selected policy. + fn decode_owned_with(value: FieldValue, _mode: DecodeMode) -> Result { + Self::decode_owned(value) + } + + /// Returns the stored field value. + /// + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "headers-user-agent")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http_headers::SingleValueField; + /// use http_headers::headers::{UserAgent, UserAgentOwned}; + /// + /// let agent = UserAgentOwned::try_from_static("client/1")?; + /// assert_eq!( + /// ::as_field_value(&agent), + /// "client/1" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "headers-user-agent"))] + /// # fn main() {} + /// ``` + fn as_field_value(value: &Self::Owned) -> &FieldValue; + + /// Consumes the field and returns its stored field value. + /// + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "headers-user-agent")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http_headers::SingleValueField; + /// use http_headers::headers::{UserAgent, UserAgentOwned}; + /// + /// let agent = UserAgentOwned::try_from_static("client/1")?; + /// assert_eq!( + /// ::into_field_value(agent), + /// "client/1" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "headers-user-agent"))] + /// # fn main() {} + /// ``` + fn into_field_value(value: Self::Owned) -> FieldValue; +} + +fn validate_single_value_list_limit(lines: &FieldLines<'_>) -> Result<(), DecodeError> { + if T::name() == &FieldName::Range { + lines.validate_list_item_limit(b',', true) + } else if T::name() == &FieldName::StrictTransportSecurity { + lines.validate_list_item_limit(b';', false) + } else { + Ok(()) + } +} + +fn validate_single_value_view(lines: &FieldLines<'_>) -> Result<(), DecodeError> { + if T::name() == &FieldName::Range || T::name() == &FieldName::StrictTransportSecurity { + validate_single_value_list_limit::(lines) + } else { + lines.validate_custom_source() + } +} + +impl Field for T +where + T: SingleValueField, +{ + type View<'a> + = T::View<'a> + where + T: 'a; + type Owned = T::Owned; + + #[inline] + fn name() -> &'static FieldName { + T::name() + } + + #[expect( + clippy::inline_always, + reason = "the single-value adapter must disappear from borrowed decode hot paths" + )] + #[inline(always)] + fn view_with(source: &S, mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(T::name()) else { + return Ok(None); + }; + validate_single_value_view::(&lines)?; + let value = lines.exactly_one()?; + T::decode_view_with(value, mode).map(Some) + } + + #[expect( + clippy::inline_always, + reason = "the single-value adapter must disappear from owned decode hot paths" + )] + #[inline(always)] + fn owned_with(source: &S, mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(T::name()) else { + return Ok(None); + }; + validate_single_value_list_limit::(&lines)?; + let owned = lines.exactly_one_owned()?; + T::decode_owned_with(owned, mode).map(Some) + } + + #[inline] + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(T::name(), EncodedValues::single(T::into_field_value(value))) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{DecodeMode, Field, SingleValueField}; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, TestSink}; + + struct FixtureHeader; + + impl SingleValueField for FixtureHeader { + type View<'a> = FieldValueRef<'a>; + type Owned = FieldValue; + + fn name() -> &'static FieldName { + &FieldName::UserAgent + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + if value.as_bytes() == b"strict" { + Ok(value) + } else { + Err(DecodeError::new(::name(), DecodeErrorKind::InvalidSyntax)) + } + } + + fn decode_view_with(value: FieldValueRef<'_>, mode: DecodeMode) -> Result, DecodeError> { + if mode == DecodeMode::Relaxed { + Ok(value) + } else { + Self::decode_view(value) + } + } + + fn decode_owned(value: FieldValue) -> Result { + if value.as_bytes() == b"strict" { + Ok(value) + } else { + Err(DecodeError::new(::name(), DecodeErrorKind::InvalidSyntax)) + } + } + + fn decode_owned_with(value: FieldValue, mode: DecodeMode) -> Result { + if mode == DecodeMode::Relaxed { + Ok(value) + } else { + Self::decode_owned(value) + } + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value + } + } + + struct Source { + values: Vec, + } + + impl FieldSource for Source { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_slice(name, &self.values) + } + } + + #[test] + fn single_value_adapter_handles_absence_modes_cardinality_insert_and_remove() { + let absent = Source { values: vec![] }; + assert_eq!(FixtureHeader::view(&absent).expect("absence is valid"), None); + assert_eq!(FixtureHeader::owned(&absent).expect("absence is valid"), None); + + let strict = Source { + values: vec![FieldValue::from_static("strict")], + }; + assert_eq!( + FixtureHeader::view(&strict) + .expect("strict value decodes") + .expect("strict value is present"), + "strict" + ); + assert_eq!( + FixtureHeader::owned(&strict) + .expect("strict value decodes") + .expect("strict value is present"), + "strict" + ); + + let relaxed = Source { + values: vec![FieldValue::from_static("relaxed")], + }; + assert_eq!( + FixtureHeader::view(&relaxed) + .expect_err("strict view rejects relaxed fixture") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + FixtureHeader::owned(&relaxed) + .expect_err("strict owned decode rejects relaxed fixture") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + FixtureHeader::view_with(&relaxed, DecodeMode::Relaxed) + .expect("relaxed view decodes") + .expect("relaxed value is present"), + "relaxed" + ); + assert_eq!( + FixtureHeader::owned_with(&relaxed, DecodeMode::Relaxed) + .expect("relaxed owned decode succeeds") + .expect("relaxed value is present"), + "relaxed" + ); + + let multiple = Source { + values: vec![FieldValue::from_static("strict"), FieldValue::from_static("strict")], + }; + assert_eq!( + FixtureHeader::view(&multiple).expect_err("multiple values are rejected").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!( + FixtureHeader::owned(&multiple) + .expect_err("multiple owned values are rejected") + .kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + + let mut sink = TestSink::new(); + let owned = FieldValue::from_static("strict"); + assert_eq!(::as_field_value(&owned), "strict"); + FixtureHeader::insert(&mut sink, owned).expect("insertion succeeds"); + assert!(sink.contains(&FieldName::UserAgent)); + FixtureHeader::remove(&mut sink); + assert!(!sink.contains(&FieldName::UserAgent)); + } +} diff --git a/crates/http_headers/src/field_name.rs b/crates/http_headers/src/field_name.rs new file mode 100644 index 000000000..9d16b9821 --- /dev/null +++ b/crates/http_headers/src/field_name.rs @@ -0,0 +1,830 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! The name a field is stored and looked up under. + +use std::error::Error; +use std::hash::{Hash, Hasher}; +use std::sync::Arc; +use std::{fmt, mem, str}; + +const STACK_NORMALIZATION_CAPACITY: usize = 64; + +/// An error produced when bytes cannot form an HTTP field name. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::{FieldName, InvalidFieldName}; +/// +/// let error: InvalidFieldName = FieldName::try_from_bytes(b"bad header").unwrap_err(); +/// assert_eq!(error.to_string(), "invalid HTTP field name"); +/// ``` +#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)] +pub struct InvalidFieldName; + +impl fmt::Display for InvalidFieldName { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("invalid HTTP field name") + } +} + +impl Error for InvalidFieldName {} + +macro_rules! known_headers { + ($(($variant:ident, $konst:ident, $text:literal),)+) => { + /// An owned HTTP field name. + /// + /// Every name in [`FieldName::ALL_KNOWN`] has its own variant. Other + /// valid names use [`FieldName::Custom`], allowing this type to + /// represent extension and application-defined headers as well. + /// + /// Names compare and hash by their lowercase bytes. Recognition is + /// ASCII case-insensitive, so a parsed `Accept` and the constant + /// [`FieldName::Accept`] are the same value. + /// + /// Runtime names support validation, comparison, and conversion to + /// container-specific names. The [`crate::source::FieldSource`] and + /// [`crate::sink::FieldSink`] traits instead require static descriptors; + /// a locally constructed name cannot be used at that boundary. Use + /// the container's native API for dynamic name operations. + /// + /// # Examples + /// + /// ```rust + /// use std::sync::LazyLock; + /// + /// use http_headers::FieldName; + /// + /// static CUSTOM: LazyLock = + /// LazyLock::new(|| FieldName::from_static("x-trace-id")); + /// + /// assert_eq!(FieldName::UserAgent.as_str(), "user-agent"); + /// assert_eq!(FieldName::try_from_bytes(b"User-Agent")?, FieldName::UserAgent); + /// assert!(FieldName::UserAgent.index().is_some_and(|index| index < FieldName::COUNT)); + /// assert_eq!(CUSTOM.index(), None); + /// # Ok::<(), http_headers::InvalidFieldName>(()) + /// ``` + #[derive(Clone)] + #[non_exhaustive] + pub enum FieldName { + $( + #[doc = concat!("The `", $text, "` header.")] + $variant, + )+ + /// A field name this crate does not recognize. + /// + /// Use [`FieldName::from_static`] for a static lowercase name or + /// [`FieldName::try_from_bytes`] for runtime input. + /// + /// The shared representation makes cloning a custom name cheap. + #[non_exhaustive] + Custom(Arc), + } + + impl FieldName { + /// Every well-known name, ordered by index. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// + /// assert_eq!(FieldName::ALL_KNOWN.len(), FieldName::COUNT); + /// ``` + pub const ALL_KNOWN: &'static [Self] = &[$(Self::$variant,)+]; + + /// The number of well-known names. + /// + /// # Examples + /// + /// ```rust + /// assert!(http_headers::FieldName::COUNT > 0); + /// ``` + pub const COUNT: usize = Self::ALL_KNOWN.len(); + + /// Returns the name's position in [`FieldName::ALL_KNOWN`]. + /// + /// A custom name has no index. The returned value is always + /// smaller than [`FieldName::COUNT`] and is suitable for indexing + /// a table with one entry per well-known name. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// + /// assert!(FieldName::Accept.index().is_some_and(|index| index < FieldName::COUNT)); + /// assert_eq!(FieldName::from_static("x-trace-id").index(), None); + /// ``` + #[must_use] + #[inline] + pub fn index(&self) -> Option { + #[repr(usize)] + enum Index { + $($variant,)+ + } + + match self { + $(Self::$variant => Some(Index::$variant as usize),)+ + Self::Custom(_) => None, + } + } + + /// Returns the lowercase field name. + /// + /// # Examples + /// + /// ```rust + /// assert_eq!(http_headers::FieldName::Accept.as_str(), "accept"); + /// ``` + #[must_use] + #[inline] + pub fn as_str(&self) -> &str { + match self { + $(Self::$variant => $text,)+ + Self::Custom(name) => name, + } + } + + /// Creates a name from a static lowercase field name. + /// + /// A recognized name becomes its corresponding well-known variant. + /// + /// # Panics + /// + /// Panics when `name` is not a lowercase HTTP token or is longer + /// than 65,535 bytes. + /// + /// # Examples + /// + /// ```rust + /// use std::sync::LazyLock; + /// + /// use http_headers::FieldName; + /// + /// static TRACE_ID: LazyLock = + /// LazyLock::new(|| FieldName::from_static("x-trace-id")); + /// + /// assert_eq!(TRACE_ID.as_str(), "x-trace-id"); + /// assert_eq!(FieldName::from_static("accept"), FieldName::Accept); + /// ``` + #[must_use] + pub fn from_static(name: &'static str) -> Self { + assert!( + is_lowercase_name(name.as_bytes()) && name.len() <= MAX_NAME_LEN, + "invalid static HTTP field name: {name:?}" + ); + Self::known_from_bytes(name.as_bytes()) + .unwrap_or_else(|| Self::Custom(Arc::from(name))) + } + + /// Returns the static `http` name of a well-known header. + /// + /// A custom name has no `http` constant. Converting the + /// [`FieldName`] itself works for both well-known and custom + /// names. + /// + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() { + /// use http_headers::FieldName; + /// + /// assert_eq!(FieldName::Accept.http_name(), Some(&http::header::ACCEPT)); + /// assert_eq!(FieldName::from_static("x-trace-id").http_name(), None); + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + #[cfg(feature = "http")] + #[must_use] + #[inline] + pub fn http_name(&self) -> Option<&'static http::HeaderName> { + match self { + $(Self::$variant => Some(&http::header::$konst),)+ + Self::Custom(_) => None, + } + } + } + + /// The lowercase text of every well-known name, ordered by index. + /// + /// This is the key table the length-bucketed recognition dispatch + /// below is derived from at compile time; it is generated from the + /// same list as the variants, so the two can never drift. + const KNOWN_TEXTS: &[&str] = &[$($text,)+]; + }; +} + +known_headers! { + (Accept, ACCEPT, "accept"), + (AcceptCharset, ACCEPT_CHARSET, "accept-charset"), + (AcceptEncoding, ACCEPT_ENCODING, "accept-encoding"), + (AcceptLanguage, ACCEPT_LANGUAGE, "accept-language"), + (AcceptRanges, ACCEPT_RANGES, "accept-ranges"), + (AccessControlAllowCredentials, ACCESS_CONTROL_ALLOW_CREDENTIALS, "access-control-allow-credentials"), + (AccessControlAllowHeaders, ACCESS_CONTROL_ALLOW_HEADERS, "access-control-allow-headers"), + (AccessControlAllowMethods, ACCESS_CONTROL_ALLOW_METHODS, "access-control-allow-methods"), + (AccessControlAllowOrigin, ACCESS_CONTROL_ALLOW_ORIGIN, "access-control-allow-origin"), + (AccessControlExposeHeaders, ACCESS_CONTROL_EXPOSE_HEADERS, "access-control-expose-headers"), + (AccessControlMaxAge, ACCESS_CONTROL_MAX_AGE, "access-control-max-age"), + (AccessControlRequestHeaders, ACCESS_CONTROL_REQUEST_HEADERS, "access-control-request-headers"), + (AccessControlRequestMethod, ACCESS_CONTROL_REQUEST_METHOD, "access-control-request-method"), + (Age, AGE, "age"), + (Allow, ALLOW, "allow"), + (AltSvc, ALT_SVC, "alt-svc"), + (Authorization, AUTHORIZATION, "authorization"), + (CacheControl, CACHE_CONTROL, "cache-control"), + (CacheStatus, CACHE_STATUS, "cache-status"), + (CdnCacheControl, CDN_CACHE_CONTROL, "cdn-cache-control"), + (Connection, CONNECTION, "connection"), + (ContentDisposition, CONTENT_DISPOSITION, "content-disposition"), + (ContentEncoding, CONTENT_ENCODING, "content-encoding"), + (ContentLanguage, CONTENT_LANGUAGE, "content-language"), + (ContentLength, CONTENT_LENGTH, "content-length"), + (ContentLocation, CONTENT_LOCATION, "content-location"), + (ContentRange, CONTENT_RANGE, "content-range"), + (ContentSecurityPolicy, CONTENT_SECURITY_POLICY, "content-security-policy"), + (ContentSecurityPolicyReportOnly, CONTENT_SECURITY_POLICY_REPORT_ONLY, "content-security-policy-report-only"), + (ContentType, CONTENT_TYPE, "content-type"), + (Cookie, COOKIE, "cookie"), + (Dnt, DNT, "dnt"), + (Date, DATE, "date"), + (Etag, ETAG, "etag"), + (Expect, EXPECT, "expect"), + (Expires, EXPIRES, "expires"), + (Forwarded, FORWARDED, "forwarded"), + (From, FROM, "from"), + (Host, HOST, "host"), + (IfMatch, IF_MATCH, "if-match"), + (IfModifiedSince, IF_MODIFIED_SINCE, "if-modified-since"), + (IfNoneMatch, IF_NONE_MATCH, "if-none-match"), + (IfRange, IF_RANGE, "if-range"), + (IfUnmodifiedSince, IF_UNMODIFIED_SINCE, "if-unmodified-since"), + (LastModified, LAST_MODIFIED, "last-modified"), + (Link, LINK, "link"), + (Location, LOCATION, "location"), + (MaxForwards, MAX_FORWARDS, "max-forwards"), + (Origin, ORIGIN, "origin"), + (Pragma, PRAGMA, "pragma"), + (ProxyAuthenticate, PROXY_AUTHENTICATE, "proxy-authenticate"), + (ProxyAuthorization, PROXY_AUTHORIZATION, "proxy-authorization"), + (PublicKeyPins, PUBLIC_KEY_PINS, "public-key-pins"), + (PublicKeyPinsReportOnly, PUBLIC_KEY_PINS_REPORT_ONLY, "public-key-pins-report-only"), + (Range, RANGE, "range"), + (Referer, REFERER, "referer"), + (ReferrerPolicy, REFERRER_POLICY, "referrer-policy"), + (Refresh, REFRESH, "refresh"), + (RetryAfter, RETRY_AFTER, "retry-after"), + (SecWebSocketAccept, SEC_WEBSOCKET_ACCEPT, "sec-websocket-accept"), + (SecWebSocketExtensions, SEC_WEBSOCKET_EXTENSIONS, "sec-websocket-extensions"), + (SecWebSocketKey, SEC_WEBSOCKET_KEY, "sec-websocket-key"), + (SecWebSocketProtocol, SEC_WEBSOCKET_PROTOCOL, "sec-websocket-protocol"), + (SecWebSocketVersion, SEC_WEBSOCKET_VERSION, "sec-websocket-version"), + (Server, SERVER, "server"), + (SetCookie, SET_COOKIE, "set-cookie"), + (StrictTransportSecurity, STRICT_TRANSPORT_SECURITY, "strict-transport-security"), + (Te, TE, "te"), + (Trailer, TRAILER, "trailer"), + (TransferEncoding, TRANSFER_ENCODING, "transfer-encoding"), + (UserAgent, USER_AGENT, "user-agent"), + (Upgrade, UPGRADE, "upgrade"), + (UpgradeInsecureRequests, UPGRADE_INSECURE_REQUESTS, "upgrade-insecure-requests"), + (Vary, VARY, "vary"), + (Via, VIA, "via"), + (Warning, WARNING, "warning"), + (WwwAuthenticate, WWW_AUTHENTICATE, "www-authenticate"), + (XContentTypeOptions, X_CONTENT_TYPE_OPTIONS, "x-content-type-options"), + (XDnsPrefetchControl, X_DNS_PREFETCH_CONTROL, "x-dns-prefetch-control"), + (XFrameOptions, X_FRAME_OPTIONS, "x-frame-options"), + (XXssProtection, X_XSS_PROTECTION, "x-xss-protection"), +} + +/// The longest field name a [`FieldName`] can hold. +/// +/// This implementation limit bounds storage and adapter conversion costs. +const MAX_NAME_LEN: usize = (1 << 16) - 1; + +#[cfg_attr(coverage_nightly, coverage(off))] +#[cfg_attr(test, mutants::skip)] // `>` to `>=` is equivalent when assigning the same length. +const fn max_known_len() -> usize { + let mut longest = 0; + let mut index = 0; + while index < KNOWN_TEXTS.len() { + let length = KNOWN_TEXTS[index].len(); + if length > longest { + longest = length; + } + index += 1; + } + longest +} + +/// The length of the longest well-known name. +const MAX_KNOWN_LEN: usize = max_known_len(); + +/// Where the well-known names of each length start in [`KNOWN_BY_LENGTH`]. +/// +/// `LENGTH_BUCKETS[n]..LENGTH_BUCKETS[n + 1]` is the range of slots holding +/// the names of length `n`, so the table carries one extra terminating entry. +const LENGTH_BUCKETS: [usize; MAX_KNOWN_LEN + 2] = { + let mut starts = [0; MAX_KNOWN_LEN + 2]; + let mut index = 0; + while index < KNOWN_TEXTS.len() { + starts[KNOWN_TEXTS[index].len() + 1] += 1; + index += 1; + } + let mut length = 1; + while length < starts.len() { + starts[length] += starts[length - 1]; + length += 1; + } + starts +}; + +/// Well-known name indexes, grouped by name length. +const KNOWN_BY_LENGTH: [usize; KNOWN_TEXTS.len()] = { + let mut grouped = [0; KNOWN_TEXTS.len()]; + let mut cursors = LENGTH_BUCKETS; + let mut index = 0; + while index < KNOWN_TEXTS.len() { + let length = KNOWN_TEXTS[index].len(); + grouped[cursors[length]] = index; + cursors[length] += 1; + index += 1; + } + grouped +}; + +/// Returns the well-known names that are `length` bytes long. +/// +/// The result is empty for a length no well-known name has, and `None` for a +/// length longer than every well-known name. +#[inline] +fn known_candidates(length: usize) -> Option<&'static [usize]> { + let start = *LENGTH_BUCKETS.get(length)?; + let end = *LENGTH_BUCKETS.get(length + 1)?; + KNOWN_BY_LENGTH.get(start..end) +} + +/// Proves the recognition tables address the same names the enum does. +/// +/// Every slot of a bucket holds a name of that bucket's length, and the +/// buckets cover every well-known name exactly once, so recognition can +/// compare only the names of the input's length and still see all of them. +const _: () = { + assert!( + KNOWN_TEXTS.len() == FieldName::COUNT, + "KNOWN_TEXTS length must equal FieldName::COUNT" + ); + assert!( + LENGTH_BUCKETS[MAX_KNOWN_LEN + 1] == KNOWN_TEXTS.len(), + "the final length bucket must end at KNOWN_TEXTS.len()" + ); + + let mut length = 0; + while length <= MAX_KNOWN_LEN { + let mut slot = LENGTH_BUCKETS[length]; + assert!( + slot <= LENGTH_BUCKETS[length + 1], + "field-name length bucket offsets must be nondecreasing" + ); + while slot < LENGTH_BUCKETS[length + 1] { + assert!( + KNOWN_TEXTS[KNOWN_BY_LENGTH[slot]].len() == length, + "each known field name must occupy the bucket matching its length" + ); + slot += 1; + } + length += 1; + } +}; + +impl FieldName { + /// Creates a name from arbitrary bytes, lowercasing them. + /// + /// A name is at most 65,535 bytes long. + /// + /// The returned name is owned, not a static descriptor. It can be compared + /// and converted for a container's native API, but a local value cannot be + /// passed to [`crate::source::FieldSource`] or [`crate::sink::FieldSink`]. + /// + /// # Errors + /// + /// Returns an error when `bytes` is not an HTTP token, or is longer than + /// 65,535 bytes. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// + /// assert_eq!(FieldName::try_from_bytes(b"Accept")?, FieldName::Accept); + /// assert!(FieldName::try_from_bytes(b"bad name").is_err()); + /// # Ok::<(), http_headers::InvalidFieldName>(()) + /// ``` + pub fn try_from_bytes(bytes: impl AsRef<[u8]>) -> Result { + let bytes = bytes.as_ref(); + if bytes.is_empty() || bytes.len() > MAX_NAME_LEN || !bytes.iter().copied().all(crate::validate::token_byte) { + return Err(InvalidFieldName); + } + if let Some(known) = Self::known_from_bytes(bytes) { + return Ok(known); + } + if bytes.iter().all(|byte| !byte.is_ascii_uppercase()) { + let lowercase = str::from_utf8(bytes).map_err(|_invalid| InvalidFieldName)?; + return Ok(Self::Custom(Arc::from(lowercase))); + } + Ok(Self::Custom(lowercase_custom_name(bytes))) + } + + /// Returns the well-known name these bytes denote, if any. + /// + /// The comparison is ASCII case-insensitive. Only the well-known names + /// whose length matches `bytes` are compared: [`LENGTH_BUCKETS`] maps a + /// length to a range of [`KNOWN_BY_LENGTH`], so recognition never walks + /// the whole table. + fn known_from_bytes(bytes: &[u8]) -> Option { + for &index in known_candidates(bytes.len())? { + let text = KNOWN_TEXTS.get(index)?; + if eq_ignore_ascii_case(text, bytes) { + return Self::ALL_KNOWN.get(index).cloned(); + } + } + + None + } + + /// Returns the lowercase field-name bytes. + /// + /// # Examples + /// + /// ```rust + /// assert_eq!(http_headers::FieldName::Accept.as_bytes(), b"accept"); + /// ``` + #[must_use] + #[inline] + pub fn as_bytes(&self) -> &[u8] { + self.as_str().as_bytes() + } + + /// Creates a name from bytes another HTTP implementation already + /// validated as a lowercase field name. + /// + /// The bytes are recognized but never revalidated. The `http` crate + /// accepts one byte this crate's own parser does not — `"`, which its + /// HTTP/2 name table admits — so revalidating a name that crate produced + /// would reject a name it considers valid. + #[cfg(feature = "http")] + fn from_validated_lowercase(name: &str) -> Self { + Self::known_from_bytes(name.as_bytes()).unwrap_or_else(|| Self::Custom(Arc::from(name))) + } + + /// Converts this name into an `http::HeaderName` without panicking. + /// + /// # Errors + /// + /// Returns an error when the name is not one the `http` crate accepts. + /// + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::InvalidFieldName> { + /// use http_headers::FieldName; + /// + /// assert_eq!( + /// FieldName::Accept.try_to_http_header_name()?, + /// http::header::ACCEPT + /// ); + /// # Ok::<(), http_headers::InvalidFieldName>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + #[cfg(feature = "http")] + pub fn try_to_http_header_name(&self) -> Result { + if let Some(known) = self.http_name() { + return Ok(known.clone()); + } + let bytes = self.as_bytes(); + http::HeaderName::from_bytes(bytes) + // A name that came from the `http` crate can hold `"`, which only + // that crate's HTTP/2 table accepts. + .or_else(|_invalid| http::HeaderName::from_lowercase(bytes)) + .map_err(|_invalid| InvalidFieldName) + } +} + +/// Keeps ordinary names stack-backed while bounding the constructor's stack frame. +fn lowercase_custom_name(bytes: &[u8]) -> Arc { + if bytes.len() <= STACK_NORMALIZATION_CAPACITY { + let mut lowercase = [0; STACK_NORMALIZATION_CAPACITY]; + for (destination, source) in lowercase.iter_mut().zip(bytes) { + *destination = source.to_ascii_lowercase(); + } + let lowercase = str::from_utf8(&lowercase[..bytes.len()]).expect("validated HTTP field-name bytes are always ASCII"); + return Arc::from(lowercase); + } + + let mut lowercase = bytes.to_vec(); + lowercase.make_ascii_lowercase(); + Arc::from(String::from_utf8(lowercase).expect("validated HTTP field-name bytes are always ASCII")) +} + +impl fmt::Debug for FieldName { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + fmt::Debug::fmt(self.as_str(), f) + } +} + +impl fmt::Display for FieldName { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } +} + +impl PartialEq for FieldName { + #[inline] + fn eq(&self, other: &Self) -> bool { + // Distinct well-known variants always carry distinct names, so comparing + // them needs no text. A custom name may still spell a well-known one, + // which only the text can settle. + match (self, other) { + (Self::Custom(_), _) | (_, Self::Custom(_)) => self.as_str().eq_ignore_ascii_case(other.as_str()), + _ => mem::discriminant(self) == mem::discriminant(other), + } + } +} + +impl Eq for FieldName {} + +impl Hash for FieldName { + fn hash(&self, state: &mut H) { + self.as_str().len().hash(state); + for byte in self.as_bytes() { + byte.to_ascii_lowercase().hash(state); + } + } +} + +impl AsRef for FieldName { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl AsRef<[u8]> for FieldName { + fn as_ref(&self) -> &[u8] { + self.as_bytes() + } +} + +impl PartialEq for FieldName { + fn eq(&self, other: &str) -> bool { + self.as_str().eq_ignore_ascii_case(other) + } +} + +impl PartialEq<&str> for FieldName { + fn eq(&self, other: &&str) -> bool { + self.as_str().eq_ignore_ascii_case(other) + } +} + +impl PartialEq for str { + fn eq(&self, other: &FieldName) -> bool { + self.eq_ignore_ascii_case(other.as_str()) + } +} + +impl TryFrom<&[u8]> for FieldName { + type Error = InvalidFieldName; + + fn try_from(bytes: &[u8]) -> Result { + Self::try_from_bytes(bytes) + } +} + +impl TryFrom<&str> for FieldName { + type Error = InvalidFieldName; + + fn try_from(name: &str) -> Result { + Self::try_from_bytes(name.as_bytes()) + } +} + +/// Compares an already lowercase name with arbitrary bytes, ignoring case. +const fn eq_ignore_ascii_case(lowercase: &str, bytes: &[u8]) -> bool { + let lowercase = lowercase.as_bytes(); + if lowercase.len() != bytes.len() { + return false; + } + let mut index = 0; + while index < lowercase.len() { + if lowercase[index] != bytes[index].to_ascii_lowercase() { + return false; + } + index += 1; + } + true +} + +/// Returns whether every byte is permitted in a lowercase field name. +const fn is_lowercase_name(bytes: &[u8]) -> bool { + if bytes.is_empty() { + return false; + } + let mut index = 0; + while index < bytes.len() { + let byte = bytes[index]; + if byte.is_ascii_uppercase() || !crate::validate::token_byte(byte) { + return false; + } + index += 1; + } + true +} + +#[cfg(feature = "http")] +mod http_conversions { + use super::FieldName; + + impl From<&FieldName> for http::HeaderName { + /// Converts a name into the `http` crate's own name type. + fn from(name: &FieldName) -> Self { + name.try_to_http_header_name().unwrap_or_else(|_invalid| invalid_custom_name()) + } + } + + impl From for http::HeaderName { + /// Converts a name into the `http` crate's own name type. + fn from(name: FieldName) -> Self { + Self::from(&name) + } + } + + impl From<&http::HeaderName> for FieldName { + /// Converts an `http` name, preserving unrecognized names as + /// [`FieldName::Custom`]. + fn from(name: &http::HeaderName) -> Self { + Self::from_validated_lowercase(name.as_str()) + } + } + + impl From for FieldName { + fn from(name: http::HeaderName) -> Self { + Self::from(&name) + } + } + + /// Reports a `FieldName::Custom` that violates its structural invariant. + /// + /// The `Custom` variant is `#[non_exhaustive]`, so it is built only by + /// this crate's own constructors, each of which enforces a lowercase + /// field-name token of at most 65,535 bytes — exactly what the `http` + /// name type accepts. No name a caller can construct reaches this, and + /// in-crate code that built one would have broken the type's invariant. + #[expect( + clippy::panic, + reason = "an infallible conversion cannot report a broken in-crate invariant any other way" + )] + fn invalid_custom_name() -> ! { + panic!( + "`FieldName::Custom` holds bytes that are not a valid HTTP field name; use `FieldName::try_to_http_header_name` to convert without panicking" + ); + } + + #[cfg(test)] + #[cfg_attr(coverage_nightly, coverage(off))] + mod tests { + #[test] + fn invariant_failure_has_an_explicit_diagnostic() { + let panic = std::panic::catch_unwind(super::invalid_custom_name); + assert!(panic.is_err()); + } + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::collections::hash_map::DefaultHasher; + use std::error::Error; + use std::hash::{Hash, Hasher}; + use std::sync::Arc; + + use super::{FieldName, InvalidFieldName, STACK_NORMALIZATION_CAPACITY}; + + #[test] + fn known_names_are_dense_case_insensitive_and_round_trip() { + assert_eq!(FieldName::ALL_KNOWN.len(), FieldName::COUNT); + for (index, name) in FieldName::ALL_KNOWN.iter().enumerate() { + assert_eq!(name.index(), Some(index)); + assert_eq!( + FieldName::try_from_bytes(name.as_str().to_ascii_uppercase().as_bytes()).expect("known name parses"), + *name + ); + assert_eq!(name.as_bytes(), name.as_str().as_bytes()); + #[cfg(feature = "http")] + assert_eq!(name.http_name().expect("known http constant").as_str(), name.as_str()); + } + + let disguised_known = FieldName::Custom(Arc::from("accept")); + assert_eq!(disguised_known, FieldName::Accept); + assert_eq!(FieldName::Accept, disguised_known); + assert_eq!( + FieldName::Custom(Arc::from("ACCEPT")), + FieldName::Accept, + "a custom name spelling a well-known one stays equal regardless of case" + ); + assert_ne!(disguised_known, FieldName::AcceptEncoding); + assert_ne!(FieldName::Accept, FieldName::AcceptEncoding); + assert_ne!( + FieldName::Custom(Arc::from("x-trace-id")), + FieldName::Custom(Arc::from("x-request-id")) + ); + assert_eq!(disguised_known.index(), None); + #[cfg(feature = "http")] + assert_eq!(disguised_known.http_name(), None); + } + + #[test] + fn custom_names_validation_formatting_hashing_and_conversions_are_consistent() { + let error = InvalidFieldName; + assert_eq!(error.to_string(), "invalid HTTP field name"); + let error: &dyn Error = &error; + assert!(error.source().is_none()); + + let custom = FieldName::try_from_bytes(b"X-Trace-ID").expect("valid custom name"); + assert_eq!(custom.as_str(), "x-trace-id"); + assert_eq!(custom.index(), None); + assert_eq!(FieldName::from_static("x-trace-id"), custom); + assert_eq!(format!("{custom:?}"), "\"x-trace-id\""); + let lowercase = FieldName::try_from_bytes(b"x-trace-id").expect("valid lowercase custom"); + assert_eq!(lowercase, custom); + assert_eq!(lowercase.as_str(), "x-trace-id"); + assert_eq!(custom.to_string(), "x-trace-id"); + assert_eq!(AsRef::::as_ref(&custom), "x-trace-id"); + assert_eq!(AsRef::<[u8]>::as_ref(&custom), b"x-trace-id"); + assert!(PartialEq::::eq(&custom, "X-TRACE-ID")); + assert!(PartialEq::<&str>::eq(&custom, &"X-TRACE-ID")); + assert!(>::eq("X-TRACE-ID", &custom)); + assert_eq!(FieldName::try_from(b"x-trace-id".as_slice()).expect("valid"), custom); + assert_eq!(FieldName::try_from("x-trace-id").expect("valid"), custom); + + let mixed = FieldName::Custom(Arc::from("X-Trace-ID")); + let mut lower_hash = DefaultHasher::new(); + custom.hash(&mut lower_hash); + let mut mixed_hash = DefaultHasher::new(); + mixed.hash(&mut mixed_hash); + assert_eq!(lower_hash.finish(), mixed_hash.finish()); + + FieldName::try_from_bytes(b"").expect_err("empty name is invalid"); + FieldName::try_from_bytes(b"bad name").expect_err("space is invalid"); + std::panic::catch_unwind(|| FieldName::from_static("")).expect_err("empty static name panics"); + std::panic::catch_unwind(|| FieldName::from_static("Upper")).expect_err("uppercase static name panics"); + std::panic::catch_unwind(|| FieldName::from_static("bad name")).expect_err("invalid static token panics"); + + for length in [STACK_NORMALIZATION_CAPACITY, STACK_NORMALIZATION_CAPACITY + 1] { + let mut name = vec![b'a'; length]; + name[0] = b'X'; + assert_eq!( + FieldName::try_from_bytes(&name).expect("valid mixed-case custom name").as_bytes(), + name.iter().map(u8::to_ascii_lowercase).collect::>() + ); + } + } + + #[cfg(feature = "http")] + #[test] + fn http_name_conversions_cover_known_and_custom_owned_and_borrowed_paths() { + let known_ref = http::HeaderName::from(&FieldName::Accept); + assert_eq!(known_ref, http::header::ACCEPT); + let known_owned = http::HeaderName::from(FieldName::Accept); + assert_eq!(known_owned, http::header::ACCEPT); + + let custom = FieldName::from_static("x-trace-id"); + let custom_http = http::HeaderName::from(&custom); + assert_eq!(custom_http.as_str(), "x-trace-id"); + assert_eq!(FieldName::from(&custom_http), custom); + assert_eq!(FieldName::from(custom_http), custom); + + let invalid_custom = FieldName::Custom(Arc::from("bad name")); + std::panic::catch_unwind(|| http::HeaderName::from(&invalid_custom)) + .expect_err("only a broken in-crate invariant reaches the panic"); + invalid_custom + .try_to_http_header_name() + .expect_err("the fallible conversion reports it instead of panicking"); + } + + #[test] + fn case_insensitive_comparison_rejects_differing_lengths() { + assert!(std::hint::black_box(super::eq_ignore_ascii_case)("accept", b"Accept")); + assert!(!std::hint::black_box(super::eq_ignore_ascii_case)("accept", b"Accept-")); + assert!(!std::hint::black_box(super::eq_ignore_ascii_case)("accept", b"Accep")); + assert!(!std::hint::black_box(super::eq_ignore_ascii_case)("accept", b"Reject")); + } +} diff --git a/crates/http_headers/src/field_value.rs b/crates/http_headers/src/field_value.rs new file mode 100644 index 000000000..15f7f3a6e --- /dev/null +++ b/crates/http_headers/src/field_value.rs @@ -0,0 +1,1399 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Owned and borrowed HTTP field values. + +use std::cmp::Ordering; +use std::error::Error; +use std::fmt; +use std::hash::{Hash, Hasher}; +use std::str::{self, FromStr, Utf8Error}; + +use bytes::Bytes; + +use crate::sink::FieldSensitivity; +use crate::validate; + +/// An error produced when bytes cannot form an HTTP field value. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::{FieldValue, InvalidFieldValue}; +/// +/// let error: InvalidFieldValue = FieldValue::from_bytes(b"line\r\nbreak").unwrap_err(); +/// assert_eq!(error.to_string(), "invalid HTTP field value"); +/// ``` +#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)] +pub struct InvalidFieldValue; + +impl fmt::Display for InvalidFieldValue { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("invalid HTTP field value") + } +} + +impl Error for InvalidFieldValue {} + +/// An owned, validated HTTP field value. +/// +/// Construction validates the [RFC 9110 field-value grammar]. +/// +/// This type is the `field-value` production itself: the bytes after the colon +/// on a single field line. RFC 9110 additionally uses "field value" for the +/// comma-joined combination of every line sharing a name. [`comma_items`] +/// walks complete items across those lines without materializing them. +/// Combining lines yields bytes that are themselves a valid `field-value`, so +/// one type serves both senses. +/// +/// Mark sensitive values to keep their bytes out of [`Debug`] output. The +/// sensitivity marker does not affect equality, ordering, or hashing. +/// +/// [RFC 9110 field-value grammar]: https://www.rfc-editor.org/rfc/rfc9110#section-5.5 +/// [`comma_items`]: crate::source::FieldLines::comma_items +/// +/// # Examples +/// +/// ```rust +/// use http_headers::FieldValue; +/// +/// let value = FieldValue::from_static("gzip"); +/// assert_eq!(value.as_bytes(), b"gzip"); +/// assert!(FieldValue::from_bytes(b"line\r\nbreak").is_err()); +/// ``` +#[derive(Clone)] +pub struct FieldValue { + repr: Repr, +} + +impl Default for FieldValue { + fn default() -> Self { + Self::from_static("") + } +} + +/// The longest value stored without allocating. +/// +/// Sized from measurement rather than from the shared variant's footprint. +/// Inlining beats sharing for values this short — a copy of a few dozen bytes +/// costs less than the atomic clone and promotion an owner reference needs — +/// and at this length the wider type is still free: no header measured slower +/// against a 38-byte buffer. Past roughly 112 bytes every header starts paying +/// for the larger moves, which is where sharing takes over instead. +const INLINE_CAPACITY: usize = 64; + +#[derive(Clone)] +enum Repr { + Inline { + len: u8, + sensitive: bool, + buf: [u8; INLINE_CAPACITY], + }, + Shared { + bytes: Bytes, + sensitive: bool, + }, + /// A value retained from an `http::HeaderMap` without copying its bytes. + /// + /// `http` keeps the `Bytes` behind a `HeaderValue` private, so sharing its + /// buffer means holding the `HeaderValue` itself. It is 40 bytes, well + /// inside what the inline variant already needs, so the extra + /// representation costs nothing in size. Sensitivity is stored separately + /// to keep the accessors `const`. + #[cfg(feature = "http")] + Http { + value: http::HeaderValue, + sensitive: bool, + }, +} + +impl Repr { + /// Stores `bytes` inline when short enough, otherwise in shared storage. + fn new(bytes: &[u8], sensitive: bool) -> Self { + if bytes.len() <= INLINE_CAPACITY { + let mut buf = [0; INLINE_CAPACITY]; + buf[..bytes.len()].copy_from_slice(bytes); + #[expect(clippy::cast_possible_truncation, reason = "guarded by the INLINE_CAPACITY check above")] + Self::Inline { + len: bytes.len() as u8, + sensitive, + buf, + } + } else { + Self::Shared { + bytes: Bytes::copy_from_slice(bytes), + sensitive, + } + } + } + + /// Stores an `http` value inline when short, and otherwise retains it. + #[cfg(feature = "http")] + #[inline] + fn from_http(value: &http::HeaderValue) -> Self { + let bytes = value.as_bytes(); + if bytes.len() <= INLINE_CAPACITY { + Self::new(bytes, value.is_sensitive()) + } else { + Self::retain_http(value.clone(), value.is_sensitive()) + } + } + + /// Retains an `http` value, refcounting its buffer rather than copying it. + #[cfg(feature = "http")] + #[cold] + #[inline(never)] + fn retain_http(value: http::HeaderValue, sensitive: bool) -> Self { + Self::Http { value, sensitive } + } + + /// Stores owned bytes inline when short enough, otherwise shares their allocation. + fn from_owner_bytes(bytes: impl AsRef<[u8]> + Into, sensitive: bool) -> Self { + let slice = bytes.as_ref(); + if slice.len() <= INLINE_CAPACITY { + Self::new(slice, sensitive) + } else { + Self::Shared { + bytes: bytes.into(), + sensitive, + } + } + } + + /// Shares `owner`'s stable projection when the value is too long to inline. + #[cfg(feature = "http")] + fn from_owner(owner: T, sensitive: bool) -> Self + where + T: AsRef<[u8]> + Send + 'static, + { + Self::from_owner_bytes(Bytes::from_owner(owner), sensitive) + } + + #[inline] + fn as_bytes(&self) -> &[u8] { + match self { + Self::Inline { len, buf, .. } => &buf[..*len as usize], + Self::Shared { bytes, .. } => bytes, + #[cfg(feature = "http")] + Self::Http { value, .. } => value.as_bytes(), + } + } + + #[inline] + const fn sensitive(&self) -> bool { + match self { + Self::Inline { sensitive, .. } | Self::Shared { sensitive, .. } => *sensitive, + #[cfg(feature = "http")] + Self::Http { sensitive, .. } => *sensitive, + } + } + + #[inline] + const fn set_sensitive(&mut self, value: bool) { + match self { + Self::Inline { sensitive, .. } | Self::Shared { sensitive, .. } => *sensitive = value, + #[cfg(feature = "http")] + Self::Http { sensitive, .. } => *sensitive = value, + } + } + + /// Consumes the value and returns shared storage, allocating when the + /// bytes were held inline. + fn into_shared(self) -> Bytes { + match self { + Self::Inline { len, buf, .. } => Bytes::copy_from_slice(&buf[..len as usize]), + Self::Shared { bytes, .. } => bytes, + #[cfg(feature = "http")] + Self::Http { value, .. } => Bytes::from_owner(value), + } + } +} + +impl FieldValue { + /// Creates a field value from bytes that already satisfy the field-value + /// grammar. + /// + /// Callers must validate `bytes` with [`crate::validate::field_value`] + /// before calling this helper. + #[inline] + pub(crate) fn from_validated_bytes(bytes: &[u8], sensitive: bool) -> Self { + debug_assert!(crate::validate::field_value(bytes)); + Self { + repr: Repr::new(bytes, sensitive), + } + } + + /// Adopts a validated buffer, copying inline only when short enough. + /// + /// Callers must establish the field-value grammar before handing over the + /// buffer. Longer values retain its allocation without copying the bytes. + #[cfg(any(test, feature = "headers-negotiation"))] + pub(crate) fn from_validated_owned_bytes(bytes: Vec, sensitive: bool) -> Self { + debug_assert!(validate::field_value(&bytes)); + Self { + repr: Repr::from_owner_bytes(bytes, sensitive), + } + } + + /// Creates a value from a static string. + /// + /// # Panics + /// + /// Panics when `value` is not a valid field value. Use + /// [`FieldValue::try_from_static`] to detect that without panicking. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// + /// assert_eq!(FieldValue::from_static("gzip").as_bytes(), b"gzip"); + /// ``` + #[must_use] + pub const fn from_static(value: &'static str) -> Self { + assert!(is_field_value(value.as_bytes()), "invalid static HTTP field value"); + Self { + repr: Repr::Shared { + bytes: Bytes::from_static(value.as_bytes()), + sensitive: false, + }, + } + } + + /// Creates a value from a static string without panicking. + /// + /// # Errors + /// + /// Returns an error when `value` is not a valid field value. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// + /// assert!(FieldValue::try_from_static("gzip").is_ok()); + /// assert!(FieldValue::try_from_static("\r\n").is_err()); + /// ``` + pub const fn try_from_static(value: &'static str) -> Result { + if is_field_value(value.as_bytes()) { + Ok(Self::from_static(value)) + } else { + Err(InvalidFieldValue) + } + } + + /// Creates a value from `bytes`. + /// + /// # Errors + /// + /// Returns an error when `bytes` is not a valid field value. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// + /// assert_eq!(FieldValue::from_bytes(b"gzip")?.as_bytes(), b"gzip"); + /// # Ok::<(), http_headers::InvalidFieldValue>(()) + /// ``` + pub fn from_bytes(bytes: impl AsRef<[u8]>) -> Result { + let bytes = bytes.as_ref(); + if validate::field_value(bytes) { + Ok(Self::from_validated_bytes(bytes, false)) + } else { + Err(InvalidFieldValue) + } + } + + /// Creates a value from `value`. + /// + /// # Errors + /// + /// Returns an error when `value` is not a valid field value. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// + /// assert_eq!(FieldValue::from_str("gzip")?.as_bytes(), b"gzip"); + /// # Ok::<(), http_headers::InvalidFieldValue>(()) + /// ``` + #[expect( + clippy::should_implement_trait, + reason = "the inherent constructor mirrors the conventional field-value API and `FromStr` is implemented as well" + )] + pub fn from_str(value: impl AsRef) -> Result { + Self::from_bytes(value.as_ref().as_bytes()) + } + + /// Creates a value from [`Bytes`]. + /// + /// # Errors + /// + /// Returns an error when `bytes` is not a valid field value. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// + /// let value = FieldValue::from_shared(bytes::Bytes::from_static(b"gzip"))?; + /// assert_eq!(value.as_bytes(), b"gzip"); + /// # Ok::<(), http_headers::InvalidFieldValue>(()) + /// ``` + pub fn from_shared(bytes: Bytes) -> Result { + if validate::field_value(&bytes) { + Ok(Self { + repr: Repr::Shared { bytes, sensitive: false }, + }) + } else { + Err(InvalidFieldValue) + } + } + + /// Creates a value that shares `owner`'s buffer instead of copying it. + /// + /// This is the hook for a [`crate::source::FieldSource`] whose storage is already + /// refcounted. Short values are copied inline, which is cheaper than + /// sharing; longer ones retain `owner`, so cloning the resulting value + /// costs a refcount bump rather than a copy of its bytes, whatever the + /// owning type is. + /// + /// # Errors + /// + /// Returns an error when `owner` does not hold a valid field value. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// + /// let value = FieldValue::from_owner(vec![b'g', b'z', b'i', b'p'])?; + /// assert_eq!(value.as_bytes(), b"gzip"); + /// # Ok::<(), http_headers::InvalidFieldValue>(()) + /// ``` + pub fn from_owner(owner: T) -> Result + where + T: AsRef<[u8]> + Send + 'static, + { + let bytes = Bytes::from_owner(owner); + if validate::field_value(&bytes) { + Ok(Self { + repr: Repr::from_owner_bytes(bytes, false), + }) + } else { + Err(InvalidFieldValue) + } + } + + /// Returns the wire bytes. + /// + /// # Examples + /// + /// ```rust + /// assert_eq!( + /// http_headers::FieldValue::from_static("gzip").as_bytes(), + /// b"gzip" + /// ); + /// ``` + #[must_use] + #[inline] + pub fn as_bytes(&self) -> &[u8] { + self.repr.as_bytes() + } + + /// Returns a borrowed view of this value. + /// + /// # Examples + /// + /// ```rust + /// let value = http_headers::FieldValue::from_static("gzip"); + /// assert_eq!(value.as_field_value_ref().as_bytes(), b"gzip"); + /// ``` + #[must_use] + #[inline] + pub fn as_field_value_ref(&self) -> FieldValueRef<'_> { + FieldValueRef::new(self.as_bytes()).with_sensitive(self.is_sensitive()) + } + + /// Returns the value as UTF-8. + /// + /// # Errors + /// + /// Returns an error when the value is not UTF-8. + /// + /// # Examples + /// + /// ```rust + /// assert_eq!( + /// http_headers::FieldValue::from_static("gzip").try_as_str()?, + /// "gzip" + /// ); + /// # Ok::<(), std::str::Utf8Error>(()) + /// ``` + #[inline] + pub fn try_as_str(&self) -> Result<&str, Utf8Error> { + str::from_utf8(self.as_bytes()) + } + + /// Returns the number of wire bytes. + /// + /// # Examples + /// + /// ```rust + /// assert_eq!(http_headers::FieldValue::from_static("gzip").len(), 4); + /// ``` + #[must_use] + #[inline] + pub fn len(&self) -> usize { + self.as_bytes().len() + } + + /// Returns whether the value is empty. + /// + /// # Examples + /// + /// ```rust + /// assert!(http_headers::FieldValue::from_static("").is_empty()); + /// ``` + #[must_use] + #[inline] + pub fn is_empty(&self) -> bool { + self.as_bytes().is_empty() + } + + /// Returns whether the value was marked as carrying sensitive data. + /// + /// # Examples + /// + /// ```rust + /// assert!(!http_headers::FieldValue::from_static("gzip").is_sensitive()); + /// ``` + #[must_use] + #[inline] + pub const fn is_sensitive(&self) -> bool { + self.repr.sensitive() + } + + /// Sets the value's sensitivity classification. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldSensitivity; + /// + /// let mut value = http_headers::FieldValue::from_static("secret"); + /// value.set_sensitivity(FieldSensitivity::Sensitive); + /// assert!(value.is_sensitive()); + /// ``` + #[inline] + pub const fn set_sensitivity(&mut self, sensitivity: FieldSensitivity) { + self.repr.set_sensitive(sensitivity.is_sensitive()); + } + + pub(crate) const fn set_sensitive(&mut self, sensitive: bool) { + self.repr.set_sensitive(sensitive); + } + + /// Returns the value with its sensitivity classification applied. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldSensitivity; + /// + /// let value = http_headers::FieldValue::from_static("secret") + /// .with_sensitivity(FieldSensitivity::Sensitive); + /// assert!(value.is_sensitive()); + /// ``` + #[must_use] + #[inline] + pub fn with_sensitivity(mut self, sensitivity: FieldSensitivity) -> Self { + self.set_sensitivity(sensitivity); + self + } + + pub(crate) fn with_sensitive(mut self, sensitive: bool) -> Self { + self.set_sensitive(sensitive); + self + } + + /// Consumes the value and returns its bytes. + /// + /// The sensitivity marker is not represented in the returned [`Bytes`]. + /// + /// # Examples + /// + /// ```rust + /// let bytes = http_headers::FieldValue::from_static("gzip").into_shared(); + /// assert_eq!(bytes.as_ref(), b"gzip"); + /// ``` + #[must_use] + #[inline] + pub fn into_shared(self) -> Bytes { + self.repr.into_shared() + } +} + +/// Returns whether every byte is permitted in an HTTP field value. +/// +/// This is the `const` counterpart of [`crate::validate::field_value`], which +/// the accelerated runtime path uses; both accept exactly the same bytes. +const fn is_field_value(bytes: &[u8]) -> bool { + let mut index = 0; + while index < bytes.len() { + let byte = bytes[index]; + if byte != b'\t' && (byte < b' ' || byte == 0x7f) { + return false; + } + index += 1; + } + true +} + +macro_rules! impl_from_integer { + ($($integer:ty),+ $(,)?) => { + $( + impl From<$integer> for FieldValue { + /// Formats the number, which is always a valid field value. + fn from(value: $integer) -> Self { + let mut buffer = itoa::Buffer::new(); + let formatted = buffer.format(value); + Self { + repr: Repr::new(formatted.as_bytes(), false), + } + } + } + )+ + }; +} + +impl_from_integer!(i16, i32, i64, isize, u16, u32, u64, usize); + +impl fmt::Debug for FieldValue { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + if self.is_sensitive() { + f.write_str("FieldValue(Sensitive)") + } else { + write!(f, "FieldValue({:?})", ByteStr(self.as_bytes())) + } + } +} + +impl PartialEq for FieldValue { + fn eq(&self, other: &Self) -> bool { + self.as_bytes() == other.as_bytes() + } +} + +impl Eq for FieldValue {} + +impl Ord for FieldValue { + fn cmp(&self, other: &Self) -> Ordering { + self.as_bytes().cmp(other.as_bytes()) + } +} + +impl PartialOrd for FieldValue { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl Hash for FieldValue { + fn hash(&self, state: &mut H) { + self.as_bytes().hash(state); + } +} + +impl AsRef<[u8]> for FieldValue { + fn as_ref(&self) -> &[u8] { + self.as_bytes() + } +} + +impl FromStr for FieldValue { + type Err = InvalidFieldValue; + + fn from_str(value: &str) -> Result { + Self::from_bytes(value.as_bytes()) + } +} + +impl TryFrom<&[u8]> for FieldValue { + type Error = InvalidFieldValue; + + fn try_from(value: &[u8]) -> Result { + Self::from_bytes(value) + } +} + +impl TryFrom<&str> for FieldValue { + type Error = InvalidFieldValue; + + fn try_from(value: &str) -> Result { + Self::from_bytes(value.as_bytes()) + } +} + +impl TryFrom> for FieldValue { + type Error = InvalidFieldValue; + + fn try_from(value: Vec) -> Result { + Self::from_shared(Bytes::from(value)) + } +} + +impl TryFrom for FieldValue { + type Error = InvalidFieldValue; + + fn try_from(value: String) -> Result { + Self::from_shared(Bytes::from(value)) + } +} + +impl TryFrom for FieldValue { + type Error = InvalidFieldValue; + + fn try_from(value: Bytes) -> Result { + Self::from_shared(value) + } +} + +impl From for Bytes { + fn from(value: FieldValue) -> Self { + value.into_shared() + } +} + +/// A borrowed HTTP field value. +/// +/// This is the borrowed counterpart of [`FieldValue`]. Field `*View` types +/// use it to expose field bytes for as long as the source remains borrowed. +/// +/// Like [`FieldValue`], it can be marked sensitive. The marker is preserved +/// when converting to an owned value and does not affect equality, ordering, +/// or hashing. [`Debug`] never includes the borrowed bytes, regardless of the +/// marker. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::FieldValueRef; +/// +/// let value = FieldValueRef::new(b"gzip"); +/// assert_eq!(value.to_str()?, "gzip"); +/// # Ok::<(), std::str::Utf8Error>(()) +/// ``` +#[derive(Clone, Copy, Default)] +pub struct FieldValueRef<'a> { + bytes: &'a [u8], + sensitive: bool, +} + +impl<'a> FieldValueRef<'a> { + /// Creates a borrowed field value. + /// + /// This constructor accepts arbitrary bytes. Use + /// [`FieldValueRef::try_to_field_value`] when converting untrusted bytes + /// into a validated [`FieldValue`]. + /// + /// # Examples + /// + /// ```rust + /// assert_eq!( + /// http_headers::FieldValueRef::new(b"gzip").as_bytes(), + /// b"gzip" + /// ); + /// ``` + #[must_use] + #[inline] + pub const fn new(bytes: &'a [u8]) -> Self { + Self { bytes, sensitive: false } + } + + /// Returns the value with its sensitivity classification applied. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldSensitivity; + /// + /// let value = + /// http_headers::FieldValueRef::new(b"secret").with_sensitivity(FieldSensitivity::Sensitive); + /// assert!(value.is_sensitive()); + /// ``` + #[must_use] + #[inline] + pub const fn with_sensitivity(mut self, sensitivity: FieldSensitivity) -> Self { + self.sensitive = sensitivity.is_sensitive(); + self + } + + pub(crate) const fn with_sensitive(mut self, sensitive: bool) -> Self { + self.sensitive = sensitive; + self + } + + /// Returns whether the value was marked as carrying sensitive data. + /// + /// # Examples + /// + /// ```rust + /// assert!(!http_headers::FieldValueRef::new(b"gzip").is_sensitive()); + /// ``` + #[must_use] + #[inline] + pub const fn is_sensitive(self) -> bool { + self.sensitive + } + + /// Returns the wire bytes. + /// + /// # Examples + /// + /// ```rust + /// assert_eq!( + /// http_headers::FieldValueRef::new(b"gzip").as_bytes(), + /// b"gzip" + /// ); + /// ``` + #[must_use] + #[inline] + pub const fn as_bytes(self) -> &'a [u8] { + self.bytes + } + + /// Returns the value as UTF-8. + /// + /// Checks the bytes for valid UTF-8 without allocating. This does not + /// validate the HTTP field-value grammar. + /// + /// # Errors + /// + /// Returns an error when the value is not UTF-8. + /// + /// # Examples + /// + /// ```rust + /// assert_eq!(http_headers::FieldValueRef::new(b"gzip").to_str()?, "gzip"); + /// # Ok::<(), std::str::Utf8Error>(()) + /// ``` + #[inline] + pub const fn to_str(self) -> Result<&'a str, Utf8Error> { + str::from_utf8(self.bytes) + } + + /// Returns the number of wire bytes. + /// + /// # Examples + /// + /// ```rust + /// assert_eq!(http_headers::FieldValueRef::new(b"gzip").len(), 4); + /// ``` + #[must_use] + #[inline] + pub const fn len(self) -> usize { + self.bytes.len() + } + + /// Returns whether the value is empty. + /// + /// # Examples + /// + /// ```rust + /// assert!(http_headers::FieldValueRef::new(b"").is_empty()); + /// ``` + #[must_use] + #[inline] + pub const fn is_empty(self) -> bool { + self.bytes.is_empty() + } + + /// Creates a validated owned value from the borrowed bytes. + /// + /// The sensitivity flag travels with the bytes. + /// + /// # Errors + /// + /// Returns an error when the borrowed bytes are not a valid HTTP field + /// value. [`FieldValueRef::new`] accepts arbitrary bytes, so a reference + /// a caller built by hand can hold bytes this validated owned type + /// rejects. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValueRef; + /// + /// let owned = FieldValueRef::new(b"gzip").try_to_field_value()?; + /// assert_eq!(owned.as_bytes(), b"gzip"); + /// assert!( + /// FieldValueRef::new(b"bad\nvalue") + /// .try_to_field_value() + /// .is_err() + /// ); + /// # Ok::<(), http_headers::InvalidFieldValue>(()) + /// ``` + #[inline] + pub fn try_to_field_value(self) -> Result { + FieldValue::from_bytes(self.bytes).map(|owned| owned.with_sensitive(self.sensitive)) + } + + #[inline] + pub(crate) fn to_validated_field_value(self) -> FieldValue { + FieldValue::from_validated_bytes(self.bytes, self.sensitive) + } +} + +impl fmt::Debug for FieldValueRef<'_> { + /// Formats the value without exposing its bytes. + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + if self.sensitive { + f.write_str("FieldValueRef(Sensitive)") + } else { + write!(f, "FieldValueRef({} bytes)", self.bytes.len()) + } + } +} + +impl PartialEq for FieldValueRef<'_> { + #[inline] + fn eq(&self, other: &Self) -> bool { + self.bytes == other.bytes + } +} + +impl Eq for FieldValueRef<'_> {} + +impl Ord for FieldValueRef<'_> { + #[inline] + fn cmp(&self, other: &Self) -> Ordering { + self.bytes.cmp(other.bytes) + } +} + +impl PartialOrd for FieldValueRef<'_> { + #[inline] + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl Hash for FieldValueRef<'_> { + fn hash(&self, state: &mut H) { + self.bytes.hash(state); + } +} + +impl AsRef<[u8]> for FieldValueRef<'_> { + fn as_ref(&self) -> &[u8] { + self.bytes + } +} + +impl<'a> From<&'a FieldValue> for FieldValueRef<'a> { + fn from(value: &'a FieldValue) -> Self { + value.as_field_value_ref() + } +} + +impl<'a> From<&'a [u8]> for FieldValueRef<'a> { + fn from(bytes: &'a [u8]) -> Self { + Self::new(bytes) + } +} + +impl TryFrom> for FieldValue { + type Error = InvalidFieldValue; + + /// Creates a validated owned value from the borrowed bytes. + /// + /// # Errors + /// + /// Returns an error when the borrowed bytes are not a valid HTTP field + /// value. + fn try_from(value: FieldValueRef<'_>) -> Result { + value.try_to_field_value() + } +} + +/// Formats bytes as a string when possible and as an escaped list otherwise. +struct ByteStr<'a>(&'a [u8]); + +impl fmt::Debug for ByteStr<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + match str::from_utf8(self.0) { + Ok(text) => fmt::Debug::fmt(text, f), + Err(_invalid) => fmt::Debug::fmt(self.0, f), + } + } +} + +macro_rules! impl_eq_bytes { + ($owner:ty, $other:ty) => { + impl PartialEq<$other> for $owner { + #[inline] + fn eq(&self, other: &$other) -> bool { + AsRef::<[u8]>::as_ref(self) == AsRef::<[u8]>::as_ref(other) + } + } + }; +} + +impl_eq_bytes!(FieldValue, str); +impl_eq_bytes!(FieldValue, &str); +impl_eq_bytes!(FieldValue, String); +impl_eq_bytes!(FieldValue, [u8]); +impl_eq_bytes!(FieldValue, &[u8]); +impl_eq_bytes!(FieldValue, Vec); +impl_eq_bytes!(FieldValue, FieldValueRef<'_>); +impl_eq_bytes!(FieldValueRef<'_>, str); +impl_eq_bytes!(FieldValueRef<'_>, &str); +impl_eq_bytes!(FieldValueRef<'_>, String); +impl_eq_bytes!(FieldValueRef<'_>, [u8]); +impl_eq_bytes!(FieldValueRef<'_>, &[u8]); +impl_eq_bytes!(FieldValueRef<'_>, Vec); +impl_eq_bytes!(FieldValueRef<'_>, FieldValue); + +impl PartialEq for str { + #[inline] + fn eq(&self, other: &FieldValue) -> bool { + self.as_bytes() == other.as_bytes() + } +} + +impl PartialEq> for str { + #[inline] + fn eq(&self, other: &FieldValueRef<'_>) -> bool { + self.as_bytes() == other.as_bytes() + } +} + +impl PartialEq for &str { + #[inline] + fn eq(&self, other: &FieldValue) -> bool { + self.as_bytes() == other.as_bytes() + } +} + +impl PartialEq> for &str { + #[inline] + fn eq(&self, other: &FieldValueRef<'_>) -> bool { + self.as_bytes() == other.as_bytes() + } +} + +#[cfg(feature = "http")] +mod http_conversions { + use http::HeaderValue; + + use super::{FieldValue, Repr}; + + impl From<&HeaderValue> for FieldValue { + /// Converts an `http` field value while preserving its sensitivity marker. + /// + /// A value too long to inline shares the `HeaderValue`'s buffer rather + /// than copying it. `http` keeps its `Bytes` private, so the retained + /// owner is a clone of the `HeaderValue` itself — a refcount bump. + fn from(value: &HeaderValue) -> Self { + Self { + repr: Repr::from_http(value), + } + } + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::collections::hash_map::DefaultHasher; + use std::error::Error; + use std::hash::{Hash, Hasher}; + use std::ptr; + use std::sync::Arc; + use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering}; + + use bytes::Bytes; + + use super::{FieldValue, FieldValueRef, INLINE_CAPACITY, InvalidFieldValue, Repr}; + + #[test] + fn short_values_are_stored_without_allocating_and_do_not_grow_the_type() { + assert_eq!(size_of::(), 72); + + let inline = FieldValue::from_bytes([b'a'; INLINE_CAPACITY]).expect("valid"); + assert!(matches!(inline.repr, Repr::Inline { .. })); + assert_eq!(inline.as_bytes(), &[b'a'; INLINE_CAPACITY]); + + let shared = FieldValue::from_bytes([b'a'; INLINE_CAPACITY + 1]).expect("valid"); + assert!(matches!(shared.repr, Repr::Shared { .. })); + assert_eq!(shared.as_bytes(), &[b'a'; INLINE_CAPACITY + 1]); + + for mut value in [inline, shared] { + let bytes = value.as_bytes().to_vec(); + value.set_sensitive(true); + assert!(value.is_sensitive()); + assert_eq!(value.clone().into_shared(), bytes); + } + } + + #[test] + fn validated_owned_buffers_stay_inline_through_the_capacity_boundary() { + for length in [0, 1, INLINE_CAPACITY - 1, INLINE_CAPACITY] { + for sensitive in [false, true] { + let mut bytes = Vec::with_capacity(INLINE_CAPACITY * 2); + bytes.extend((0..length).map(|index| [b'\t', b' ', b'~', 0x80, 0xff][index % 5])); + let expected = bytes.clone(); + + let value = FieldValue::from_validated_owned_bytes(bytes, sensitive); + + assert!(matches!(value.repr, Repr::Inline { .. })); + assert_eq!(value.as_bytes(), expected); + assert_eq!(value.is_sensitive(), sensitive); + } + } + } + + #[test] + fn validated_owned_buffers_retain_payload_allocations_with_or_without_spare_capacity() { + for length in [INLINE_CAPACITY + 1, INLINE_CAPACITY * 2, INLINE_CAPACITY * 2 + 1] { + for spare in [0, 37] { + for sensitive in [false, true] { + let mut bytes = Vec::with_capacity(length + spare); + bytes.extend((0..length).map(|index| [b'\t', b' ', b'~', 0x80, 0xff][index % 5])); + let expected = bytes.clone(); + let address = bytes.as_ptr(); + + let value = FieldValue::from_validated_owned_bytes(bytes, sensitive); + + assert!(matches!(value.repr, Repr::Shared { .. })); + assert!(ptr::eq(value.as_bytes().as_ptr(), address)); + assert_eq!(value.as_bytes(), expected); + assert_eq!(value.is_sensitive(), sensitive); + + let retained = value.clone(); + drop(value); + assert!(ptr::eq(retained.as_bytes().as_ptr(), address)); + assert_eq!(retained.as_bytes(), expected); + assert_eq!(retained.is_sensitive(), sensitive); + + let shared = retained.into_shared(); + assert!(ptr::eq(shared.as_ptr(), address)); + assert_eq!(shared.as_ref(), expected); + } + } + } + } + + #[cfg(debug_assertions)] + #[test] + fn validated_owned_buffers_debug_check_the_field_value_invariant() { + for invalid in [b"bad\nvalue".to_vec(), vec![0x7f; INLINE_CAPACITY + 1]] { + std::panic::catch_unwind(|| FieldValue::from_validated_owned_bytes(invalid, false)).unwrap_err(); + } + } + + #[test] + fn constructors_accessors_sensitivity_and_debug_validate_real_bytes() { + let error = InvalidFieldValue; + assert_eq!(error.to_string(), "invalid HTTP field value"); + let error: &dyn Error = &error; + assert!(error.source().is_none()); + + let default = FieldValue::default(); + assert!(default.is_empty()); + assert_eq!(default.len(), 0); + assert_eq!(FieldValue::try_from_static("gzip").expect("valid"), "gzip"); + FieldValue::try_from_static("\r\n").expect_err("CRLF is invalid"); + std::panic::catch_unwind(|| FieldValue::from_static("\n")).expect_err("invalid static value panics"); + + let copied = FieldValue::from_bytes(b"\t visible \xff").expect("valid field bytes"); + assert_eq!(copied.as_bytes(), b"\t visible \xff"); + copied.try_as_str().expect_err("obs-text is not UTF-8"); + assert_eq!( + format!("{copied:?}"), + "FieldValue([9, 32, 118, 105, 115, 105, 98, 108, 101, 32, 255])" + ); + FieldValue::from_bytes(b"\0").expect_err("NUL is invalid"); + FieldValue::from_bytes(b"\x7f").expect_err("DEL is invalid"); + FieldValue::from_str("line\nbreak").expect_err("newline is invalid"); + + let shared = FieldValue::from_shared(Bytes::from_static(b"shared")).expect("valid"); + assert_eq!(shared.try_as_str().expect("UTF-8"), "shared"); + FieldValue::from_shared(Bytes::from_static(b"\r")).expect_err("carriage return is invalid"); + assert_eq!(shared.into_shared(), Bytes::from_static(b"shared")); + + let mut sensitive = FieldValue::from_static("secret"); + sensitive.set_sensitive(true); + assert!(sensitive.is_sensitive()); + assert_eq!(format!("{sensitive:?}"), "FieldValue(Sensitive)"); + sensitive.set_sensitive(false); + assert!(!sensitive.is_sensitive()); + assert!(sensitive.with_sensitive(true).is_sensitive()); + } + + #[test] + fn numeric_conversion_comparison_hashing_and_owned_conversions_are_consistent() { + assert_eq!(FieldValue::from(i16::MIN), i16::MIN.to_string()); + assert_eq!(FieldValue::from(i32::MIN), i32::MIN.to_string()); + assert_eq!(FieldValue::from(i64::MIN), i64::MIN.to_string()); + assert_eq!(FieldValue::from(isize::MIN), isize::MIN.to_string()); + assert_eq!(FieldValue::from(u16::MAX), u16::MAX.to_string()); + assert_eq!(FieldValue::from(u32::MAX), u32::MAX.to_string()); + assert_eq!(FieldValue::from(u64::MAX), u64::MAX.to_string()); + assert_eq!(FieldValue::from(usize::MAX), usize::MAX.to_string()); + + let value = FieldValue::from_static("abc"); + let same = "abc".parse::().expect("valid"); + let later = FieldValue::from_static("abd"); + assert_eq!(value, same); + assert!(value < later); + assert_eq!(value.partial_cmp(&later), Some(std::cmp::Ordering::Less)); + assert_eq!(value.as_ref(), b"abc"); + + let mut first_hash = DefaultHasher::new(); + value.hash(&mut first_hash); + let mut second_hash = DefaultHasher::new(); + same.with_sensitive(true).hash(&mut second_hash); + assert_eq!(first_hash.finish(), second_hash.finish()); + + assert_eq!(FieldValue::try_from(b"abc".as_slice()).expect("valid"), "abc"); + assert_eq!(FieldValue::try_from("abc").expect("valid"), "abc"); + assert_eq!(FieldValue::try_from(b"abc".to_vec()).expect("valid"), "abc"); + assert_eq!(FieldValue::try_from(String::from("abc")).expect("valid"), "abc"); + assert_eq!(FieldValue::try_from(Bytes::from_static(b"abc")).expect("valid"), "abc"); + let bytes: Bytes = FieldValue::from_static("abc").into(); + assert_eq!(bytes, Bytes::from_static(b"abc")); + } + + #[test] + fn borrowed_values_and_symmetric_comparisons_match_wire_bytes() { + let owned = FieldValue::from_static("abc"); + let borrowed = FieldValueRef::from(&owned); + assert_eq!(borrowed.as_bytes(), b"abc"); + assert_eq!(borrowed.to_str().expect("UTF-8"), "abc"); + assert_eq!(borrowed.len(), 3); + assert!(!borrowed.is_empty()); + assert_eq!(borrowed.as_ref(), b"abc"); + assert_eq!(format!("{borrowed:?}"), "FieldValueRef(3 bytes)"); + + let empty = FieldValueRef::default(); + assert!(empty.is_empty()); + let from_bytes = FieldValueRef::from(b"abc".as_slice()); + assert_eq!(from_bytes.try_to_field_value().expect("valid bytes"), owned); + assert_eq!(FieldValue::try_from(from_bytes).expect("valid bytes"), owned); + FieldValueRef::new(b"\xff").to_str().expect_err("obs-text is not UTF-8"); + assert_eq!(format!("{:?}", FieldValueRef::new(b"\xff")), "FieldValueRef(1 bytes)"); + + let string = String::from("abc"); + let vector = b"abc".to_vec(); + let slice = b"abc".as_slice(); + assert!(PartialEq::::eq(&owned, "abc")); + assert!(PartialEq::<&str>::eq(&owned, &"abc")); + assert!(PartialEq::::eq(&owned, &string)); + assert!(PartialEq::<[u8]>::eq(&owned, b"abc")); + assert!(PartialEq::<&[u8]>::eq(&owned, &slice)); + assert!(PartialEq::>::eq(&owned, &vector)); + assert!(PartialEq::>::eq(&owned, &borrowed)); + assert!(PartialEq::::eq(&borrowed, "abc")); + assert!(PartialEq::<&str>::eq(&borrowed, &"abc")); + assert!(PartialEq::::eq(&borrowed, &string)); + assert!(PartialEq::<[u8]>::eq(&borrowed, b"abc")); + assert!(PartialEq::<&[u8]>::eq(&borrowed, &slice)); + assert!(PartialEq::>::eq(&borrowed, &vector)); + assert!(PartialEq::::eq(&borrowed, &owned)); + assert!("abc".eq(&owned)); + assert!("abc".eq(&borrowed)); + assert!((&"abc").eq(&owned)); + assert!((&"abc").eq(&borrowed)); + } + + #[test] + fn borrowed_values_redact_debug_and_carry_sensitivity_into_owned_values() { + let secret = FieldValue::from_static("Bearer credential").with_sensitive(true); + let borrowed = secret.as_field_value_ref(); + assert!(borrowed.is_sensitive()); + assert_eq!(format!("{borrowed:?}"), "FieldValueRef(Sensitive)"); + + let owned = borrowed.try_to_field_value().expect("valid bytes"); + assert!(owned.is_sensitive()); + assert_eq!(format!("{owned:?}"), "FieldValue(Sensitive)"); + assert_eq!(owned.as_bytes(), b"Bearer credential"); + + let plain = FieldValueRef::new(b"Bearer credential"); + assert!(!plain.is_sensitive()); + let rendered = format!("{plain:?}"); + assert!( + !rendered.contains("credential"), + "a borrowed value must never render its bytes: {rendered}" + ); + assert!(plain.with_sensitive(true).is_sensitive()); + + assert_eq!(plain, borrowed); + assert!(plain <= borrowed); + let mut plain_hash = DefaultHasher::new(); + plain.hash(&mut plain_hash); + let mut sensitive_hash = DefaultHasher::new(); + borrowed.hash(&mut sensitive_hash); + assert_eq!(plain_hash.finish(), sensitive_hash.finish()); + } + + #[cfg(feature = "http")] + #[test] + fn http_conversions_preserve_bytes_and_sensitivity() { + let mut http = http::HeaderValue::from_static("secret"); + http.set_sensitive(true); + let borrowed_owned = FieldValue::from(&http); + assert_eq!(borrowed_owned, "secret"); + assert!(borrowed_owned.is_sensitive()); + + let moved_owned = FieldValue::from(http); + let round_trip = http::HeaderValue::try_from(moved_owned).expect("valid"); + assert_eq!(round_trip, "secret"); + assert!(round_trip.is_sensitive()); + + let borrowed = FieldValueRef::from(&round_trip); + assert_eq!(borrowed, "secret"); + assert!(borrowed.is_sensitive()); + assert_eq!(format!("{borrowed:?}"), "FieldValueRef(Sensitive)"); + let converted = http::HeaderValue::try_from(borrowed).expect("valid"); + assert_eq!(converted, "secret"); + assert!(converted.is_sensitive()); + assert!(borrowed.try_to_field_value().expect("valid").is_sensitive()); + http::HeaderValue::try_from(FieldValueRef::new(b"\n")).expect_err("newline is invalid"); + } + + #[cfg(feature = "http")] + #[test] + fn a_long_http_value_is_shared_rather_than_copied() { + let long = "a".repeat(INLINE_CAPACITY + 1); + let http = http::HeaderValue::from_str(&long).expect("valid"); + let mut shared = FieldValue::from(&http); + assert_eq!(shared.as_bytes(), long.as_bytes()); + assert!( + std::ptr::eq(shared.as_bytes().as_ptr(), http.as_bytes().as_ptr()), + "a value too long to inline must retain the source buffer, not copy it" + ); + shared.set_sensitive(true); + let round_trip = http::HeaderValue::try_from(shared.clone()).expect("retained HTTP value is valid"); + assert!(round_trip.is_sensitive()); + assert!( + std::ptr::eq(round_trip.as_bytes().as_ptr(), http.as_bytes().as_ptr()), + "round-tripping a retained HTTP value must preserve its shared buffer" + ); + assert_eq!(shared.into_shared().as_ref(), long.as_bytes()); + + let short = http::HeaderValue::from_static("client/1"); + let inlined = FieldValue::from(&short); + assert_eq!(inlined.as_bytes(), b"client/1"); + assert!( + !std::ptr::eq(inlined.as_bytes().as_ptr(), short.as_bytes().as_ptr()), + "a short value is copied inline, which costs less than sharing" + ); + } + + #[test] + fn an_owner_is_retained_only_when_the_value_is_too_long_to_inline() { + let long = vec![b'a'; INLINE_CAPACITY + 1]; + let address = long.as_ptr(); + let shared = FieldValue::from_owner(long).expect("valid"); + assert!(std::ptr::eq(shared.as_bytes().as_ptr(), address)); + assert!(matches!(shared.repr, Repr::Shared { .. })); + + let short = vec![b'a'; INLINE_CAPACITY]; + let inlined = FieldValue::from_owner(short).expect("valid"); + assert!(matches!(inlined.repr, Repr::Inline { .. })); + assert_eq!(inlined.as_bytes(), &[b'a'; INLINE_CAPACITY]); + + FieldValue::from_owner(vec![b'\n']).expect_err("a newline is not a field value"); + } + + #[test] + fn owner_validation_and_storage_use_one_stable_projection() { + struct StatefulOwner { + valid: Vec, + invalid: Vec, + calls: Arc, + } + + impl AsRef<[u8]> for StatefulOwner { + fn as_ref(&self) -> &[u8] { + if self.calls.fetch_add(1, AtomicOrdering::SeqCst) == 0 { + &self.valid + } else { + &self.invalid + } + } + } + + let calls = Arc::new(AtomicUsize::new(0)); + let owner = StatefulOwner { + valid: vec![b'a'; INLINE_CAPACITY + 1], + invalid: { + let mut bytes = vec![b'a'; INLINE_CAPACITY + 1]; + bytes[INLINE_CAPACITY] = b'\n'; + bytes + }, + calls: Arc::clone(&calls), + }; + let value = FieldValue::from_owner(owner).expect("the captured projection is valid"); + assert_eq!(value.as_bytes(), &[b'a'; INLINE_CAPACITY + 1]); + assert_eq!(calls.load(AtomicOrdering::SeqCst), 1); + } +} + +#[cfg(feature = "http")] +mod http_conversions_owned { + use http::HeaderValue; + + use super::{FieldValue, FieldValueRef, InvalidFieldValue, Repr}; + + impl From for FieldValue { + /// Converts an `http` field value while preserving its sensitivity marker. + /// + /// A value too long to inline is retained rather than copied. + fn from(value: HeaderValue) -> Self { + let sensitive = value.is_sensitive(); + Self { + repr: Repr::from_owner(value, sensitive), + } + } + } + + impl TryFrom for HeaderValue { + type Error = InvalidFieldValue; + + /// Converts the value while preserving its sensitivity marker. + /// + /// A value retained from an `http::HeaderMap` is handed back as it + /// was, without rebuilding or revalidating it. + fn try_from(value: FieldValue) -> Result { + let sensitive = value.is_sensitive(); + let mut converted = match value.repr { + Repr::Http { value, .. } => value, + repr => Self::from_maybe_shared(repr.into_shared()).map_err(|_invalid| InvalidFieldValue)?, + }; + converted.set_sensitive(sensitive); + Ok(converted) + } + } + + impl TryFrom> for HeaderValue { + type Error = InvalidFieldValue; + + fn try_from(value: FieldValueRef<'_>) -> Result { + let mut converted = Self::from_bytes(value.as_bytes()).map_err(|_invalid| InvalidFieldValue)?; + converted.set_sensitive(value.is_sensitive()); + Ok(converted) + } + } + + impl<'a> From<&'a HeaderValue> for FieldValueRef<'a> { + /// Borrows the value's bytes, carrying its sensitivity flag along. + fn from(value: &'a HeaderValue) -> Self { + Self::new(value.as_bytes()).with_sensitive(value.is_sensitive()) + } + } +} diff --git a/crates/http_headers/src/headers/authorization.rs b/crates/http_headers/src/headers/authorization.rs new file mode 100644 index 000000000..81940146f --- /dev/null +++ b/crates/http_headers/src/headers/authorization.rs @@ -0,0 +1,1362 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Basic and bearer `Authorization` header types. + +use std::fmt; +use std::io::Write as _; +use std::marker::PhantomData; +#[cfg(test)] +use std::sync::{Arc, Mutex}; + +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD; +use base64::write::EncoderWriter; +use zeroize::Zeroize; + +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, SingleValueField, validate}; + +/// Caps reusable decoded credential storage at 64 KiB to retain common credentials without +/// indefinitely holding attacker-sized allocations; increasing it trades memory for reuse. +const DEFAULT_CREDENTIAL_RETAIN_LIMIT: usize = 64 * 1024; + +#[cfg(test)] +type ZeroizationLog = Arc)>>>; + +/// The Bearer authorization scheme. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::AuthorizationOwned::::bearer( +/// "abc.def", +/// )?; +/// assert_eq!(value.token()?, b"abc.def"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct Bearer; + +/// The Basic authorization scheme. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{AuthorizationOwned, Basic}; +/// +/// let value = AuthorizationOwned::::basic(b"Aladdin", b"open sesame")?; +/// assert_eq!( +/// value.encoded_credentials()?, +/// b"QWxhZGRpbjpvcGVuIHNlc2FtZQ==" +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct Basic; + +/// Defines the `Authorization` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 11.6.2](https://www.rfc-editor.org/rfc/rfc9110#section-11.6.2). +/// The Basic and Bearer schemes are defined by +/// [RFC 7617 section 2](https://www.rfc-editor.org/rfc/rfc7617#section-2) and +/// [RFC 6750 section 2.1](https://www.rfc-editor.org/rfc/rfc6750#section-2.1). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{Authorization, AuthorizationOwned, Bearer}; +/// +/// let mut map = HeaderMap::new(); +/// Authorization::::insert(&mut map, AuthorizationOwned::::bearer("abc.def")?)?; +/// assert!(Authorization::::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct Authorization { + _private: PhantomData S>, +} + +/// Owned value for the `Authorization` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 11.6.2], with the Basic scheme specified by +/// [RFC 7617 section 2] and the Bearer scheme specified by [RFC 6750 section 2.1]. +/// +/// # Examples +/// +/// ```rust +/// let authorization = +/// http_headers::headers::AuthorizationOwned::::bearer( +/// "abc.def", +/// )?; +/// assert_eq!(authorization.token()?, b"abc.def"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Authorization: Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ==` carries Basic +/// credentials, while the Bearer scheme carries a token68 credential. +/// +/// [RFC 9110 section 11.6.2]: https://www.rfc-editor.org/rfc/rfc9110#section-11.6.2 +/// [RFC 7617 section 2]: https://www.rfc-editor.org/rfc/rfc7617#section-2 +/// [RFC 6750 section 2.1]: https://www.rfc-editor.org/rfc/rfc6750#section-2.1 +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct AuthorizationOwned { + value: FieldValue, + scheme: PhantomData, +} + +/// Borrowed value for the `Authorization` header. +#[derive(Clone, Copy, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{Authorization, AuthorizationView, Basic}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let view: AuthorizationView<'_, Basic> = +/// as SingleValueField>::decode_view(FieldValueRef::new( +/// b"Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ==", +/// ))?; +/// assert_eq!(view.credentials(), b"QWxhZGRpbjpvcGVuIHNlc2FtZQ=="); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct AuthorizationView<'a, S> { + value: FieldValueRef<'a>, + credentials: &'a [u8], + scheme: PhantomData, +} + +/// Reusable storage for decoded Basic credentials. +/// +/// Existing credentials are zeroized before every extraction and when this +/// value is dropped. Only initialized credential storage is wiped; retained +/// spare capacity contains no credentials that were not wiped first. Capacity +/// is reused up to a configurable retention limit. +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{Authorization, AuthorizationOwned, Basic, BasicCredentials}; +/// +/// let mut map = HeaderMap::new(); +/// Authorization::::insert( +/// &mut map, +/// AuthorizationOwned::::basic(b"user", b"password")?, +/// )?; +/// let authorization = Authorization::::view(&map)?.expect("authorization present"); +/// let mut credentials = BasicCredentials::new(); +/// let decoded = authorization.extract(&mut credentials)?; +/// assert_eq!(decoded.username(), b"user"); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +pub struct BasicCredentials { + bytes: Vec, + username_end: usize, + password_start: usize, + retain_limit: usize, + #[cfg(test)] + zeroization_observer: Option, +} + +impl fmt::Debug for AuthorizationOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("AuthorizationOwned") + .field("sensitive", &true) + .finish_non_exhaustive() + } +} + +impl fmt::Debug for AuthorizationView<'_, S> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("AuthorizationView") + .field("sensitive", &true) + .finish_non_exhaustive() + } +} + +impl fmt::Debug for BasicCredentials { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("BasicCredentials").field("sensitive", &true).finish_non_exhaustive() + } +} + +impl Default for BasicCredentials { + fn default() -> Self { + Self::new() + } +} + +impl std::str::FromStr for AuthorizationOwned +where + Authorization: SingleValueField, +{ + type Err = DecodeError; + + fn from_str(value: &str) -> Result { + let value = FieldValue::from_str(value).map_err(|_invalid| super::invalid_syntax(&FieldName::Authorization))?; + as SingleValueField>::decode_owned(value) + } +} + +impl Drop for BasicCredentials { + fn drop(&mut self) { + self.zeroize_initialized(); + } +} + +impl AuthorizationOwned { + /// Constructs a Bearer authorization value. + /// + /// # Errors + /// + /// Returns an error when `token` is not valid Bearer token68 syntax. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::AuthorizationOwned::::bearer( + /// "abc.def", + /// )?; + /// assert_eq!(value.token()?, b"abc.def"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn bearer(token: impl AsRef) -> Result { + let token = token.as_ref(); + if !token68(token.as_bytes()) { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + build_authorization("Bearer", token.as_bytes()) + } + + /// Returns the Bearer token. + /// # Errors + /// + /// Returns an error if the stored range and wire value disagree. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::AuthorizationOwned::::bearer( + /// "abc.def", + /// )?; + /// assert_eq!(value.token()?, b"abc.def"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn token(&self) -> Result<&[u8], DecodeError> { + parse_scheme(self.value.as_field_value_ref(), &BEARER).map(|(_start, credentials)| credentials) + } +} + +impl AuthorizationOwned { + /// Constructs Basic credentials from arbitrary username and password bytes. + /// + /// # Errors + /// + /// Returns an error if the username contains `:` or encoded sizing + /// overflows. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{AuthorizationOwned, Basic}; + /// + /// let value = AuthorizationOwned::::basic(b"Aladdin", b"open sesame")?; + /// assert_eq!( + /// value.encoded_credentials()?, + /// b"QWxhZGRpbjpvcGVuIHNlc2FtZQ==" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn basic(username: impl AsRef<[u8]>, password: impl AsRef<[u8]>) -> Result { + let username = username.as_ref(); + let password = password.as_ref(); + if username.contains(&b':') { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + let total_len = basic_wire_len(username.len(), password.len())?; + let prefix = b"Basic "; + let mut wire = Vec::with_capacity(total_len); + wire.extend_from_slice(prefix); + encode_basic_credentials(&mut wire, username, password); + debug_assert_eq!( + wire.len(), + total_len, + "Basic authorization wire length must match the precomputed capacity" + ); + let mut value = super::value_from_bytes(&FieldName::Authorization, wire)?; + value.set_sensitive(true); + Ok(Self { + value, + scheme: PhantomData, + }) + } + + /// Returns the base64-encoded credential bytes. + /// # Errors + /// + /// Returns an error if the stored range and wire value disagree. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{AuthorizationOwned, Basic}; + /// + /// let value = AuthorizationOwned::::basic(b"Aladdin", b"open sesame")?; + /// assert_eq!( + /// value.encoded_credentials()?, + /// b"QWxhZGRpbjpvcGVuIHNlc2FtZQ==" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn encoded_credentials(&self) -> Result<&[u8], DecodeError> { + parse_scheme(self.value.as_field_value_ref(), &BASIC).map(|(_start, credentials)| credentials) + } + + /// Decodes the username and password into reusable credential storage. + /// + /// # Errors + /// + /// Returns an error if the stored authorization value is invalid. + pub fn extract<'a>(&self, output: &'a mut BasicCredentials) -> Result<&'a BasicCredentials, DecodeError> { + output.fill(self.encoded_credentials()?) + } +} + +impl AuthorizationOwned { + /// Returns the complete sensitive field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{AuthorizationOwned, Bearer}; + /// + /// let value = AuthorizationOwned::::bearer("abc.def")?; + /// assert_eq!(value.as_field_value().as_bytes(), b"Bearer abc.def"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Returns reusable sensitive wire storage. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{AuthorizationOwned, Basic}; + /// + /// let value = AuthorizationOwned::::basic(b"Aladdin", b"open sesame")?; + /// let field_value = value.into_field_value(); + /// assert_eq!( + /// field_value.as_bytes(), + /// b"Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ==" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +impl From> for FieldValue { + #[inline] + fn from(value: AuthorizationOwned) -> Self { + value.value + } +} + +impl<'a, S> AuthorizationView<'a, S> { + /// Returns the encoded credential bytes after the authorization scheme. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{Authorization, Basic}; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = as SingleValueField>::decode_view(FieldValueRef::new( + /// b"Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ==", + /// ))?; + /// assert_eq!(view.credentials(), b"QWxhZGRpbjpvcGVuIHNlc2FtZQ=="); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn credentials(self) -> &'a [u8] { + self.credentials + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{Authorization, Bearer}; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let value = FieldValueRef::new(b"Bearer abc.def"); + /// let view = as SingleValueField>::decode_view(value)?; + /// assert_eq!(view.as_field_value().as_bytes(), b"Bearer abc.def"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::wrong_self_convention, + reason = "borrowed views are Copy and expose value-style accessors consistently" + )] + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value.with_sensitive(true) + } +} + +impl AuthorizationView<'_, Basic> { + /// Decodes the username and password into reusable credential storage. + /// + /// # Errors + /// + /// Returns an error if the stored authorization value is invalid. + pub fn extract<'a>(&self, output: &'a mut BasicCredentials) -> Result<&'a BasicCredentials, DecodeError> { + output.fill(self.credentials) + } +} + +impl<'a> AuthorizationView<'a, Bearer> { + /// Returns the borrowed Bearer token. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::AuthorizationOwned::::bearer( + /// "abc.def", + /// )?; + /// assert_eq!(value.token()?, b"abc.def"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn token(self) -> &'a [u8] { + self.credentials + } +} + +impl BasicCredentials { + /// Creates empty reusable credential storage. + #[must_use] + pub const fn new() -> Self { + Self::with_retain_limit(DEFAULT_CREDENTIAL_RETAIN_LIMIT) + } + + /// Creates credential storage with a custom retained-capacity limit. + #[must_use] + pub const fn with_retain_limit(retain_limit: usize) -> Self { + Self { + bytes: Vec::new(), + username_end: 0, + password_start: 0, + retain_limit, + #[cfg(test)] + zeroization_observer: None, + } + } + + fn fill(&mut self, encoded: &[u8]) -> Result<&Self, DecodeError> { + self.clear(); + let result = (|| { + STANDARD + .decode_vec(encoded, &mut self.bytes) + .map_err(|_invalid| super::invalid_syntax(&FieldName::Authorization))?; + self.bytes + .iter() + .position(|byte| *byte == b':') + .ok_or_else(|| super::invalid_syntax(&FieldName::Authorization)) + })(); + let colon = match result { + Ok(colon) => colon, + Err(error) => { + self.clear(); + return Err(error); + } + }; + self.username_end = colon; + self.password_start = colon + 1; + Ok(self) + } + + /// Returns the decoded username bytes. + #[must_use] + pub fn username(&self) -> &[u8] { + &self.bytes[..self.username_end] + } + + /// Returns the decoded password bytes. + #[must_use] + pub fn password(&self) -> &[u8] { + &self.bytes[self.password_start..] + } + + /// Zeroizes the decoded credentials and applies the retention limit. + pub fn clear(&mut self) { + self.zeroize_initialized(); + self.bytes.clear(); + self.username_end = 0; + self.password_start = 0; + if self.bytes.capacity() > self.retain_limit { + self.bytes.shrink_to(self.retain_limit); + } + } + + /// Returns the currently allocated credential capacity. + #[must_use] + pub fn capacity(&self) -> usize { + self.bytes.capacity() + } + + // `decode_vec` extends `bytes` before writing, including on errors, so the + // Vec length covers every initialized credential byte. Reuse always wipes + // that range before clearing or reallocating, leaving spare capacity free + // of historical credentials. + fn zeroize_initialized(&mut self) { + self.bytes.as_mut_slice().zeroize(); + #[cfg(test)] + if let Some(observer) = &self.zeroization_observer { + observer + .lock() + .unwrap() + .push((self.bytes.len(), self.bytes.capacity(), self.bytes.clone())); + } + } +} + +impl SingleValueField for Authorization { + type View<'a> = AuthorizationView<'a, Bearer>; + type Owned = AuthorizationOwned; + + fn name() -> &'static FieldName { + &FieldName::Authorization + } + + #[inline] + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + let (_credential_start, credentials) = parse_scheme(value, &BEARER)?; + if !token68(credentials) { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + Ok(AuthorizationView { + value, + credentials, + scheme: PhantomData, + }) + } + + #[expect( + clippy::inline_always, + reason = "measured: Criterion otherwise outlines this conversion while Callgrind inlines it" + )] + #[inline(always)] + fn decode_owned(value: FieldValue) -> Result { + let (_credential_start, credentials) = parse_scheme(value.as_field_value_ref(), &BEARER)?; + if !token68(credentials) { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + Ok(authorization_owned_from_value(value)) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +impl SingleValueField for Authorization { + type View<'a> = AuthorizationView<'a, Basic>; + type Owned = AuthorizationOwned; + + fn name() -> &'static FieldName { + &FieldName::Authorization + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + let (_credential_start, credentials) = parse_scheme(value, &BASIC)?; + validate_basic(credentials)?; + Ok(AuthorizationView { + value, + credentials, + scheme: PhantomData, + }) + } + + fn decode_owned(value: FieldValue) -> Result { + let (_credential_start, credentials) = parse_scheme(value.as_field_value_ref(), &BASIC)?; + validate_basic(credentials)?; + Ok(authorization_owned_from_value(value)) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +fn encode_basic_credentials(wire: &mut Vec, username: &[u8], password: &[u8]) { + let mut encoder = EncoderWriter::new(wire, &STANDARD); + encoder + .write_all(username) + .and_then(|()| encoder.write_all(b":")) + .and_then(|()| encoder.write_all(password)) + .and_then(|()| encoder.finish().map(|_wire| ())) + .expect("base64 encoding into a Vec cannot fail"); +} + +#[inline] +fn basic_wire_len(username_len: usize, password_len: usize) -> Result { + let input_len = username_len + .checked_add(1) + .and_then(|length| length.checked_add(password_len)) + .ok_or_else(authorization_size_error)?; + let encoded_len = base64::encoded_len(input_len, true).ok_or_else(authorization_size_error)?; + let total_len = b"Basic ".len().checked_add(encoded_len).ok_or_else(authorization_size_error)?; + Ok(total_len) +} + +#[inline] +fn authorization_wire_len(scheme_len: usize, credentials_len: usize) -> Result { + scheme_len + .checked_add(1) + .and_then(|length| length.checked_add(credentials_len)) + .ok_or_else(authorization_size_error) +} + +#[cold] +fn authorization_size_error() -> DecodeError { + DecodeError::new(&FieldName::Authorization, DecodeErrorKind::InvalidNumber) +} + +fn authorization_owned_from_value(value: FieldValue) -> AuthorizationOwned { + let mut value = value; + value.set_sensitive(true); + AuthorizationOwned { + value, + scheme: PhantomData, + } +} + +fn build_authorization(scheme: &str, credentials: &[u8]) -> Result, DecodeError> { + let total_len = authorization_wire_len(scheme.len(), credentials.len())?; + let mut wire = Vec::with_capacity(total_len); + wire.extend_from_slice(scheme.as_bytes()); + wire.push(b' '); + wire.extend_from_slice(credentials); + let mut value = super::value_from_bytes(&FieldName::Authorization, wire)?; + value.set_sensitive(true); + Ok(AuthorizationOwned { + value, + scheme: PhantomData, + }) +} + +/// A case-insensitive scheme prefix matched eight bytes at a time. +struct Scheme { + name: &'static [u8], + /// Lowercases the scheme bytes of a little-endian eight-byte load. + lower: u64, + /// Retains the scheme bytes plus the delimiting space. + keep: u64, + /// The lowercased scheme followed by the delimiting space. + expected: u64, +} + +impl Scheme { + const fn new(name: &'static [u8]) -> Self { + assert!(name.len() < 8, "scheme plus its space must fit in a word"); + let mut lower = [0_u8; 8]; + let mut keep = [0_u8; 8]; + let mut expected = [0_u8; 8]; + let mut index = 0; + while index < name.len() { + lower[index] = 0x20; + keep[index] = 0xff; + expected[index] = name[index]; + index += 1; + } + keep[index] = 0xff; + expected[index] = b' '; + Self { + name, + lower: u64::from_le_bytes(lower), + keep: u64::from_le_bytes(keep), + expected: u64::from_le_bytes(expected), + } + } +} + +const BASIC: Scheme = Scheme::new(b"basic"); +const BEARER: Scheme = Scheme::new(b"bearer"); + +/// Length at which the accelerated `token68` validator beats the local scalar +/// loop, matching the dispatch threshold of the acceleration crate. +const TOKEN68_SIMD_THRESHOLD: usize = 32; + +/// Maps every byte to zero when it is `token68` data and `0xff` otherwise. +const TOKEN68_INVALID: [u8; 256] = { + let mut table = [0xff_u8; 256]; + let mut index = 0; + while index < 256 { + #[expect(clippy::cast_possible_truncation, reason = "index is below 256")] + let byte = index as u8; + if matches!( + byte, + b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'.' | b'_' | b'~' | b'+' | b'/' + ) { + table[index] = 0; + } + index += 1; + } + table +}; + +/// Returns whether `bytes` is an RFC 9110 `token68` value. +/// +/// Equivalent to [`validate::token68`], but short credentials skip the +/// dispatch preamble and validate through a branchless table scan. +fn token68(bytes: &[u8]) -> bool { + if bytes.len() >= TOKEN68_SIMD_THRESHOLD { + return token68_accelerated(bytes); + } + if all_token68_data(bytes) { + return !bytes.is_empty(); + } + token68_with_padding(bytes) +} + +/// Keeps the vector implementation out of line so that short credentials +/// validate without its register pressure and stack frame. +#[inline(never)] +fn token68_accelerated(bytes: &[u8]) -> bool { + validate::token68(bytes) +} + +/// Validates values that carry trailing `=` padding. +#[inline(never)] +fn token68_with_padding(bytes: &[u8]) -> bool { + let Some(last_data) = bytes.iter().rposition(|byte| *byte != b'=') else { + return false; + }; + all_token68_data(&bytes[..=last_data]) +} + +/// Returns whether every byte of `bytes` is `token68` data. +/// +/// The scan walks eight-byte windows and finishes with a window aligned to the +/// end of the slice. That final window overlaps the previous one, which is +/// harmless because the scan only accumulates rejection bits. +#[expect( + clippy::inline_always, + reason = "measured: folding the short token68 scan into Bearer decoding saves 17 Ir" +)] +#[inline(always)] +fn all_token68_data(bytes: &[u8]) -> bool { + let Some(head) = bytes.first_chunk::<8>() else { + let mut invalid = 0_u8; + for byte in bytes { + invalid |= TOKEN68_INVALID[usize::from(*byte)]; + } + return invalid == 0; + }; + let mut invalid = token68_invalid_bits(head); + let mut rest = &bytes[8..]; + while let Some((chunk, tail)) = rest.split_first_chunk::<8>() { + invalid |= token68_invalid_bits(chunk); + rest = tail; + } + if !rest.is_empty() { + let tail = bytes.last_chunk::<8>().unwrap_or(head); + invalid |= token68_invalid_bits(tail); + } + invalid == 0 +} + +#[inline] +#[expect( + clippy::trivially_copy_pass_by_ref, + reason = "copying the window costs more instructions than indexing it in place" +)] +fn token68_invalid_bits(chunk: &[u8; 8]) -> u8 { + let mut invalid = 0_u8; + for byte in chunk { + invalid |= TOKEN68_INVALID[usize::from(*byte)]; + } + invalid +} + +fn parse_scheme<'a>(value: FieldValueRef<'a>, scheme: &Scheme) -> Result<(usize, &'a [u8]), DecodeError> { + let bytes = value.as_bytes(); + let Some(head) = bytes.first_chunk::<8>() else { + return parse_short_scheme(bytes, scheme); + }; + if (u64::from_le_bytes(*head) | scheme.lower) & scheme.keep != scheme.expected { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + credentials_at(bytes, scheme.name.len() + 1) +} + +/// Matches schemes in values too short for the eight-byte load. +#[cold] +fn parse_short_scheme<'a>(bytes: &'a [u8], scheme: &Scheme) -> Result<(usize, &'a [u8]), DecodeError> { + let scheme_length = scheme.name.len(); + if bytes.len() <= scheme_length || bytes[scheme_length] != b' ' || !bytes[..scheme_length].eq_ignore_ascii_case(scheme.name) { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + credentials_at(bytes, scheme_length + 1) +} + +/// Skips optional whitespace at `start` and returns the credential bytes. +#[inline] +fn credentials_at(bytes: &[u8], start: usize) -> Result<(usize, &[u8]), DecodeError> { + let rest = &bytes[start..]; + let offset = rest + .iter() + .position(|byte| *byte != b' ') + .ok_or_else(|| super::invalid_syntax(&FieldName::Authorization))?; + Ok((start + offset, &rest[offset..])) +} + +/// Sentinel bit marking a byte that is not part of the base64 alphabet. +const NOT_SEXTET: u8 = 0x80; + +/// Maps every byte to its base64 sextet, or to [`NOT_SEXTET`]. +const BASE64_SEXTET: [u8; 256] = { + let mut table = [NOT_SEXTET; 256]; + let mut index = 0; + while index < 256 { + #[expect(clippy::cast_possible_truncation, reason = "index is below 256")] + let byte = index as u8; + table[index] = match byte { + b'A'..=b'Z' => byte - b'A', + b'a'..=b'z' => byte - b'a' + 26, + b'0'..=b'9' => byte - b'0' + 52, + b'+' => 62, + b'/' => 63, + _ => NOT_SEXTET, + }; + index += 1; + } + table +}; + +fn validate_basic(encoded: &[u8]) -> Result<(), DecodeError> { + if encoded.is_empty() || !encoded.len().is_multiple_of(4) { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + let (body, last) = encoded.split_at(encoded.len() - 4); + let mut colons = 0_u32; + for chunk in body.chunks_exact(4) { + let &[first, second, third, fourth] = <&[u8; 4]>::try_from(chunk).expect("chunks_exact(4) always yields four bytes"); + let first = BASE64_SEXTET[usize::from(first)]; + let second = BASE64_SEXTET[usize::from(second)]; + let third = BASE64_SEXTET[usize::from(third)]; + let fourth = BASE64_SEXTET[usize::from(fourth)]; + if (first | second | third | fourth) & NOT_SEXTET != 0 { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + let group = (u32::from(first) << 18) | (u32::from(second) << 12) | (u32::from(third) << 6) | u32::from(fourth); + colons |= colon_marks(group); + } + validate_basic_last(last, colons != 0) +} + +/// Marks which of the three bytes packed into `group` decode to `:`. +#[inline] +const fn colon_marks(group: u32) -> u32 { + let difference = group ^ 0x003a_3a3a; + difference.wrapping_sub(0x0001_0101) & !difference & 0x0080_8080 +} + +/// Validates the final base64 quantum, which alone may carry `=` padding. +fn validate_basic_last(chunk: &[u8], mut has_colon: bool) -> Result<(), DecodeError> { + let &[first, second, third, fourth] = <&[u8; 4]>::try_from(chunk).expect("the final quantum always holds four bytes"); + let first_sextet = BASE64_SEXTET[usize::from(first)]; + let second_sextet = BASE64_SEXTET[usize::from(second)]; + if (first_sextet | second_sextet) & NOT_SEXTET != 0 { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + has_colon |= (first_sextet << 2) | (second_sextet >> 4) == b':'; + + let third_sextet = BASE64_SEXTET[usize::from(third)]; + if third_sextet & NOT_SEXTET != 0 { + if third != b'=' || fourth != b'=' || second_sextet & 0x0f != 0 { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + } else { + has_colon |= (second_sextet << 4) | (third_sextet >> 2) == b':'; + + let fourth_sextet = BASE64_SEXTET[usize::from(fourth)]; + if fourth_sextet & NOT_SEXTET != 0 { + if fourth != b'=' || third_sextet & 0x03 != 0 { + return Err(super::invalid_syntax(&FieldName::Authorization)); + } + } else { + has_colon |= (third_sextet << 6) | fourth_sextet == b':'; + } + } + + if has_colon { + Ok(()) + } else { + Err(super::invalid_syntax(&FieldName::Authorization)) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #![expect( + clippy::assertions_on_result_states, + reason = "tests classify parser outcomes without needing successful values" + )] + + use std::sync::{Arc, Mutex}; + + use base64::Engine as _; + use base64::engine::general_purpose::STANDARD; + + use super::{ + Authorization, AuthorizationOwned, AuthorizationView, BASIC, BEARER, Basic, BasicCredentials, Bearer, Scheme, + TOKEN68_SIMD_THRESHOLD, all_token68_data, authorization_owned_from_value, authorization_wire_len, basic_wire_len, + build_authorization, credentials_at, parse_scheme, parse_short_scheme, token68, token68_invalid_bits, token68_with_padding, + validate_basic, validate_basic_last, + }; + use crate::sink::FieldSink; + use crate::source::FieldSource; + use crate::{DecodeErrorKind, FieldName, FieldValue, SingleValueField, TestSink}; + + type ZeroizationObserver = Arc)>>>; + + // Both sizes preserve base64 quantum alignment and heap-backed header values. + const CREDENTIAL_BLOCK_LEN: usize = if cfg!(miri) { 64 } else { 1_024 }; + + fn observed_credentials() -> (BasicCredentials, ZeroizationObserver) { + let observer = Arc::new(Mutex::new(Vec::new())); + let mut credentials = BasicCredentials::new(); + credentials.zeroization_observer = Some(Arc::clone(&observer)); + (credentials, observer) + } + + fn assert_zeroized(snapshot: &(usize, usize, Vec), expected_len: usize, lifecycle: &str) { + assert_eq!(snapshot.0, expected_len, "{lifecycle} observed an unexpected initialized length"); + assert!(snapshot.0 <= snapshot.1, "{lifecycle} initialized length exceeds capacity"); + let snapshot = &snapshot.2; + assert_eq!( + snapshot.len(), + expected_len, + "{lifecycle} must cover all initialized credential bytes" + ); + assert!(snapshot.iter().all(|byte| *byte == 0), "{lifecycle} left credential bytes"); + } + + #[test] + fn basic_validator_matches_the_decoder_and_colon_oracle() { + let valid = b"dXNlcjpwYXNzd29yZA=="; + for index in 0..valid.len() - 4 { + for byte in crate::test_support::substitution_bytes(valid[index], index, valid.len() - 4) { + let mut candidate = valid.to_vec(); + candidate[index] = byte; + let oracle = STANDARD.decode(&candidate).is_ok_and(|decoded| decoded.contains(&b':')); + assert_eq!(validate_basic(&candidate).is_ok(), oracle, "index {index}, byte {byte:#04x}"); + } + } + } + + #[test] + fn bearer_construction_parsing_and_views_cover_token68_paths() { + for token in ["a", "abc.def", "YWJj=", "abcdefghijklmnopqrstuvwxyz0123456789"] { + let owned = AuthorizationOwned::::bearer(token).expect("valid bearer token"); + assert_eq!(owned.token(), Ok(token.as_bytes())); + assert!(owned.as_field_value().is_sensitive()); + assert!(format!("{owned:?}").contains("sensitive")); + + let parsed = owned + .as_field_value() + .try_as_str() + .expect("ASCII authorization") + .parse::>() + .expect("FromStr accepts constructed value"); + assert_eq!(parsed.token(), Ok(token.as_bytes())); + + let mut table = TestSink::new(); + Authorization::::insert(&mut table, owned).expect("table accepts bearer"); + let view = Authorization::::view(&table).expect("valid bearer").expect("present"); + assert_eq!(view.token(), token.as_bytes()); + assert_eq!(view.credentials(), token.as_bytes()); + assert_eq!( + view.as_field_value().as_bytes(), + table + .lines(&FieldName::Authorization) + .expect("authorization stored") + .repeated() + .next() + .expect("one value") + .as_bytes() + ); + assert!(format!("{view:?}").contains("sensitive")); + assert_eq!( + Authorization::::owned(&table) + .expect("valid owned bearer") + .expect("present") + .token(), + Ok(token.as_bytes()) + ); + table.remove_values(&FieldName::Authorization); + assert!(Authorization::::view(&table).expect("absence is valid").is_none()); + } + + for token in ["", "=", "==", "ab=c", "ab c"] { + assert!(AuthorizationOwned::::bearer(token).is_err(), "{token:?}"); + } + assert!(token68(b"short")); + assert!(!token68(b"")); + assert!(token68(b"data==")); + assert!(!token68(b"====")); + assert!(!token68(b"data=x")); + assert!(token68(&[b'a'; TOKEN68_SIMD_THRESHOLD])); + assert!(!token68(&[b' '; TOKEN68_SIMD_THRESHOLD])); + assert!(all_token68_data(b"1234567")); + assert!(all_token68_data(b"12345678")); + assert!(all_token68_data(b"123456789")); + assert!(all_token68_data(b"1234567890123456")); + assert!(!all_token68_data(b"1234567 ")); + assert_eq!(token68_invalid_bits(b"12345678"), 0); + assert_ne!(token68_invalid_bits(b"1234567 "), 0); + } + + #[test] + fn basic_credentials_cover_construction_extraction_errors_and_retention() { + assert!(AuthorizationOwned::::basic(b"user:name", b"secret").is_err()); + let owned = AuthorizationOwned::::basic(b"user", b"password").expect("valid credentials"); + assert_eq!(owned.encoded_credentials(), Ok(b"dXNlcjpwYXNzd29yZA==".as_slice())); + assert!(owned.as_field_value().is_sensitive()); + assert!(format!("{owned:?}").contains("sensitive")); + let parsed = owned + .as_field_value() + .try_as_str() + .expect("ASCII authorization") + .parse::>() + .expect("basic FromStr"); + assert_eq!(parsed.encoded_credentials(), owned.encoded_credentials()); + assert!(parsed.into_field_value().is_sensitive()); + + let mut credentials = BasicCredentials::default(); + let extracted = owned.extract(&mut credentials).expect("valid base64"); + assert_eq!(extracted.username(), b"user"); + assert_eq!(extracted.password(), b"password"); + assert!(format!("{extracted:?}").contains("sensitive")); + assert!(credentials.capacity() >= b"user:password".len()); + credentials.clear(); + assert_eq!(credentials.username(), b""); + assert_eq!(credentials.password(), b""); + + let mut table = TestSink::new(); + Authorization::::insert(&mut table, owned).expect("table accepts basic"); + let view = Authorization::::view(&table).expect("valid basic").expect("present"); + assert_eq!(view.credentials(), b"dXNlcjpwYXNzd29yZA=="); + let extracted = view.extract(&mut credentials).expect("valid view credentials"); + assert_eq!(extracted.username(), b"user"); + assert_eq!(extracted.password(), b"password"); + assert!(Authorization::::owned(&table).expect("valid owned basic").is_some()); + + let mut bounded = BasicCredentials::with_retain_limit(0); + AuthorizationOwned::::basic([b'a'; 256], b"") + .expect("valid large credentials") + .extract(&mut bounded) + .expect("valid extraction"); + assert!(bounded.capacity() >= 257); + bounded.clear(); + assert_eq!(bounded.capacity(), 0); + + for encoded in [b"!!!!".as_slice(), b"bm9jb2xvbg=="] { + assert!(bounded.fill(encoded).is_err()); + assert_eq!(bounded.username(), b""); + assert_eq!(bounded.password(), b""); + } + } + + #[test] + fn basic_credential_capacity_tracks_decoded_output_instead_of_encoded_input() { + let username = [b'a'; CREDENTIAL_BLOCK_LEN * 4]; + let authorization = AuthorizationOwned::::basic(username, b"").unwrap(); + let encoded_len = authorization.encoded_credentials().unwrap().len(); + let mut credentials = BasicCredentials::new(); + authorization.extract(&mut credentials).unwrap(); + + assert!(credentials.capacity() > username.len()); + assert!(credentials.capacity() < encoded_len); + } + + #[test] + fn basic_credential_clear_zeroizes_the_initialized_range() { + let (mut credentials, observer) = observed_credentials(); + AuthorizationOwned::::basic([b'a'; CREDENTIAL_BLOCK_LEN], [b'b'; CREDENTIAL_BLOCK_LEN]) + .unwrap() + .extract(&mut credentials) + .unwrap(); + let credential_len = credentials.username().len() + 1 + credentials.password().len(); + observer.lock().unwrap().clear(); + + credentials.clear(); + + let snapshots = observer.lock().unwrap(); + assert_eq!(snapshots.len(), 1); + assert_zeroized(&snapshots[0], credential_len, "clear"); + } + + #[test] + fn failed_basic_credential_decode_zeroizes_old_and_partial_output() { + let (mut credentials, observer) = observed_credentials(); + AuthorizationOwned::::basic([b'a'; CREDENTIAL_BLOCK_LEN], b"secret") + .unwrap() + .extract(&mut credentials) + .unwrap(); + let old_len = credentials.username().len() + 1 + credentials.password().len(); + observer.lock().unwrap().clear(); + + let malformed = [STANDARD.encode([b'x'; CREDENTIAL_BLOCK_LEN]), "!".into()].concat(); + assert!(credentials.fill(malformed.as_bytes()).is_err()); + + let snapshots = observer.lock().unwrap(); + assert_eq!(snapshots.len(), 2, "failed fill must erase old storage and decoder output"); + assert_zeroized(&snapshots[0], old_len, "failed decode old credentials"); + assert!(!snapshots[1].2.is_empty()); + assert_zeroized(&snapshots[1], snapshots[1].0, "failed decode partial output"); + } + + #[test] + fn failed_basic_credential_parse_zeroizes_decoded_output() { + let (mut credentials, observer) = observed_credentials(); + + assert!(credentials.fill(b"bm9jb2xvbg==").is_err()); + + let snapshots = observer.lock().unwrap(); + assert_eq!(snapshots.len(), 2, "failed parse must erase decoded output"); + assert_zeroized(&snapshots[0], 0, "failed parse empty input"); + assert_zeroized(&snapshots[1], b"nocolon".len(), "failed parse decoded output"); + } + + #[test] + fn basic_credential_reuse_bounds_later_wipes_to_the_short_output() { + let (mut credentials, observer) = observed_credentials(); + AuthorizationOwned::::basic([b'a'; CREDENTIAL_BLOCK_LEN * 2], [b'b'; CREDENTIAL_BLOCK_LEN * 2]) + .unwrap() + .extract(&mut credentials) + .unwrap(); + let old_len = credentials.username().len() + 1 + credentials.password().len(); + let old_capacity = credentials.capacity(); + observer.lock().unwrap().clear(); + + credentials.fill(b"dTpw").unwrap(); + + { + let snapshots = observer.lock().unwrap(); + assert_eq!(snapshots.len(), 1); + assert_zeroized(&snapshots[0], old_len, "reuse"); + assert!(snapshots[0].0 < old_capacity, "reuse must not wipe spare capacity"); + } + assert_eq!(credentials.username(), b"u"); + assert_eq!(credentials.password(), b"p"); + + observer.lock().unwrap().clear(); + credentials.clear(); + let snapshots = observer.lock().unwrap(); + assert_eq!(snapshots.len(), 1); + assert_zeroized(&snapshots[0], 3, "short reuse"); + assert_eq!(snapshots[0].1, old_capacity, "reuse must retain the high-water allocation"); + } + + #[test] + fn basic_credential_growth_wipes_before_reallocation_and_on_drop() { + let observer = { + let (mut credentials, observer) = observed_credentials(); + credentials.fill(b"dTpw").unwrap(); + let old_capacity = credentials.capacity(); + observer.lock().unwrap().clear(); + + let decoded = [b"expanded:".as_slice(), &[b'x'; CREDENTIAL_BLOCK_LEN * 8]].concat(); + let encoded = STANDARD.encode(&decoded); + credentials.fill(encoded.as_bytes()).unwrap(); + assert!(credentials.capacity() > old_capacity, "larger credentials must reallocate"); + + { + let snapshots = observer.lock().unwrap(); + assert_eq!(snapshots.len(), 1); + assert_zeroized(&snapshots[0], 3, "pre-reallocation"); + assert_eq!(snapshots[0].1, old_capacity); + } + observer.lock().unwrap().clear(); + observer + }; + + let snapshots = observer.lock().unwrap(); + assert_eq!(snapshots.len(), 1); + assert_zeroized( + &snapshots[0], + b"expanded:".len() + CREDENTIAL_BLOCK_LEN * 8, + "post-reallocation drop", + ); + } + + #[test] + fn dropping_basic_credentials_zeroizes_the_initialized_range() { + let observer = { + let (mut credentials, observer) = observed_credentials(); + AuthorizationOwned::::basic([b'a'; CREDENTIAL_BLOCK_LEN], b"secret") + .unwrap() + .extract(&mut credentials) + .unwrap(); + observer.lock().unwrap().clear(); + observer + }; + + let snapshots = observer.lock().unwrap(); + assert_eq!(snapshots.len(), 1); + assert!(!snapshots[0].2.is_empty()); + assert_zeroized(&snapshots[0], snapshots[0].0, "drop"); + } + + #[test] + fn scheme_and_basic_quantum_validation_cover_boundary_cases() { + let short = FieldValue::from_static("Basic X"); + assert_eq!( + parse_scheme(short.as_field_value_ref(), &BASIC).expect("short scheme matches").1, + b"X" + ); + for (wire, scheme) in [("Basic", &BASIC), ("Basic ", &BASIC), ("Bearer", &BEARER), ("Digest x", &BASIC)] { + assert!( + parse_scheme(FieldValue::from_str(wire).expect("legal field value").as_field_value_ref(), scheme).is_err(), + "{wire}" + ); + } + + for valid in [b"Og==".as_slice(), b"YTo=", b"YWI6", b"YWJjOg=="] { + assert!(validate_basic(valid).is_ok(), "{valid:?}"); + } + for invalid in [b"".as_slice(), b"abc", b"Oh==", b"YTp=", b"YWJj", b"YWJj====", b"!!!!"] { + assert!(validate_basic(invalid).is_err(), "{invalid:?}"); + } + + let malformed = FieldValue::from_static("Basic bm9jb2xvbg=="); + assert_eq!( + as SingleValueField>::decode_owned(malformed) + .expect_err("decoded credentials require a colon") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + let malformed = FieldValue::from_static("Bearer ab=c"); + assert!( as SingleValueField>::decode_view(malformed.as_field_value_ref()).is_err()); + assert!( as SingleValueField>::decode_owned(malformed).is_err()); + assert!("\n".parse::>().is_err()); + assert!("\n".parse::>().is_err()); + + let bearer = FieldValue::from_static("Bearer abc"); + let bearer_owned = as SingleValueField>::decode_owned(bearer.clone()).expect("valid owned bearer"); + assert_eq!( as SingleValueField>::as_field_value(&bearer_owned), &bearer); + let basic = FieldValue::from_static("Basic dTpw"); + let basic_view = as SingleValueField>::decode_view(basic.as_field_value_ref()).expect("valid borrowed basic"); + assert_eq!(basic_view.credentials(), b"dTpw"); + let basic_owned = as SingleValueField>::decode_owned(basic.clone()).expect("valid owned basic"); + assert_eq!( as SingleValueField>::as_field_value(&basic_owned), &basic); + } + + #[test] + fn private_scheme_and_base64_helpers_run_directly() { + let runtime_name = std::hint::black_box(b"basic".as_slice()); + let runtime_scheme = Scheme::new(runtime_name); + assert_eq!(runtime_scheme.name, b"basic"); + assert_eq!( + parse_short_scheme(b"BaSiC credential", &runtime_scheme).expect("runtime scheme").1, + b"credential" + ); + assert_eq!( + credentials_at(b"Basic credential", 6) + .expect("credentials after optional spaces") + .1, + b"credential" + ); + assert!(credentials_at(b"Basic ", 6).is_err()); + + assert!(token68_with_padding(b"abc==")); + assert!(!token68_with_padding(b"===")); + assert!(!token68_with_padding(b"ab=!=")); + + assert!(validate_basic_last(b"Og==", false).is_ok()); + assert!(validate_basic_last(b"YTo=", false).is_ok()); + assert!(validate_basic_last(b"YWI6", false).is_ok()); + assert!(validate_basic_last(b"YQ==", true).is_ok()); + assert!(validate_basic_last(b"YQ=A", false).is_err()); + assert!(validate_basic_last(b"YWF=", false).is_err()); + } + + #[test] + fn private_owned_paths_and_credential_reuse_clear_failures() { + let custom = build_authorization("Custom", b"credential").expect("valid wire value"); + assert_eq!(custom.as_field_value().as_bytes(), b"Custom credential"); + assert!(custom.as_field_value().is_sensitive()); + assert!(custom.token().is_err()); + + let wrapped = authorization_owned_from_value::(FieldValue::from_static("Bearer x")); + assert_eq!(wrapped.token(), Ok(b"x".as_slice())); + let raw = as SingleValueField>::into_field_value(wrapped); + assert!(raw.is_sensitive()); + + let empty = AuthorizationOwned::::basic(b"", b"").expect("empty username and password"); + let mut credentials = BasicCredentials::new(); + let decoded = empty.extract(&mut credentials).expect("colon-only credentials"); + assert_eq!(decoded.username(), b""); + assert_eq!(decoded.password(), b""); + + let binary = AuthorizationOwned::::basic(b"user", [0, 0xff]).expect("basic credentials accept arbitrary password bytes"); + let decoded = binary.extract(&mut credentials).expect("binary password decodes"); + assert_eq!(decoded.username(), b"user"); + assert_eq!(decoded.password(), [0, 0xff]); + let retained_capacity = credentials.capacity(); + assert!(credentials.fill(b"!!!!").is_err()); + assert_eq!(credentials.username(), b""); + assert_eq!(credentials.password(), b""); + assert_eq!(credentials.capacity(), retained_capacity); + + let malformed = AuthorizationOwned:: { + value: FieldValue::from_static("Bearer x"), + scheme: std::marker::PhantomData, + }; + assert!(malformed.encoded_credentials().is_err()); + assert!(malformed.extract(&mut credentials).is_err()); + + let direct = FieldValue::from_static("Bearer direct"); + let view = AuthorizationView:: { + value: direct.as_field_value_ref(), + credentials: b"direct", + scheme: std::marker::PhantomData, + }; + assert_eq!(view.token(), b"direct"); + } + + #[test] + fn wire_length_helpers_cover_boundaries_without_allocating() { + assert_eq!(authorization_wire_len(6, 5), Ok(12)); + assert!(authorization_wire_len(usize::MAX, 0).is_err()); + assert!(authorization_wire_len(usize::MAX - 1, 1).is_err()); + + assert_eq!(basic_wire_len(0, 0), Ok(10)); + assert!(basic_wire_len(usize::MAX, 0).is_err()); + assert!(basic_wire_len(usize::MAX - 1, 1).is_err()); + assert!(basic_wire_len(usize::MAX / 4 * 3, 0).is_err()); + } +} diff --git a/crates/http_headers/src/headers/cache_control.rs b/crates/http_headers/src/headers/cache_control.rs new file mode 100644 index 000000000..999fd382e --- /dev/null +++ b/crates/http_headers/src/headers/cache_control.rs @@ -0,0 +1,1445 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! `Cache-Control` parsing, summaries, and construction. + +use std::fmt::Write as _; +use std::time::Duration; +use std::{fmt, str}; + +use compact_str::CompactString; +use smallvec::SmallVec; + +use super::ExtensionValue; +use crate::sink::{ + EncodedValues, FieldEncodeOutput, FieldEncoder, FieldSensitivity, FieldSink, FieldValueWriter, InsertError, InsertErrorKind, +}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Defines the `Cache-Control` header. +/// +/// # Specification +/// +/// Defined by [RFC 9111 section 5.2](https://www.rfc-editor.org/rfc/rfc9111#section-5.2). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{CacheControl, CacheControlOwned}; +/// +/// let mut map = HeaderMap::new(); +/// CacheControl::insert( +/// &mut map, +/// CacheControlOwned::try_from("max-age=60, no-cache")?, +/// )?; +/// assert!(CacheControl::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct CacheControl { + _private: (), +} + +/// Owned value for the `Cache-Control` header. +/// +/// # Specification +/// +/// Defined by [RFC 9111 section 5.2]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::CacheControlOwned::try_from("max-age=60, no-cache")?; +/// assert_eq!(value.max_age(), Some(std::time::Duration::from_secs(60))); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Cache-Control: no-cache` prevents reuse without validation. +/// `Cache-Control: public, max-age=86400, stale-while-revalidate=60` combines +/// standard and extension directives. +/// +/// [RFC 9111 section 5.2]: https://www.rfc-editor.org/rfc/rfc9111#section-5.2 +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct CacheControlOwned { + values: SmallVec<[FieldValue; 1]>, + summary: CacheSummary, +} + +/// Borrowed value for the `Cache-Control` header. +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), http_headers::DecodeError> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{CacheControl, CacheControlView}; +/// +/// let mut map = HeaderMap::new(); +/// map.insert( +/// http::header::CACHE_CONTROL, +/// http::HeaderValue::from_static("public, max-age=31536000, immutable"), +/// ); +/// let view: CacheControlView<'_> = CacheControl::view(&map)?.expect("header present"); +/// assert_eq!( +/// view.max_age(), +/// Some(std::time::Duration::from_secs(31536000)) +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +pub struct CacheControlView<'a> { + values: FieldLines<'a>, + summary: CacheSummary, +} + +#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)] +struct CacheSummary { + no_cache: bool, + max_age: Option, +} + +impl fmt::Debug for CacheControlOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("CacheControlOwned") + .field("value_count", &self.values.len()) + .field("summary", &self.summary) + .finish() + } +} + +impl fmt::Debug for CacheControlView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("CacheControlView") + .field("value_count", &self.values.len()) + .field("summary", &self.summary) + .finish() + } +} + +/// One borrowed cache directive. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{CacheControlOwned, CacheDirectiveView}; +/// +/// let value = CacheControlOwned::try_from("s-maxage=120")?; +/// let directive: CacheDirectiveView<'_> = value.directives().next().expect("one directive"); +/// assert_eq!(directive.name(), "s-maxage"); +/// assert_eq!(directive.value(), Some(b"120".as_slice())); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct CacheDirectiveView<'a> { + name: &'a str, + value: Option<&'a [u8]>, + raw: &'a [u8], +} + +/// Builder for a canonical `Cache-Control` field value. +#[derive(Clone, Debug, Default)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{CacheControlBuilder, CacheControlOwned}; +/// +/// let builder: CacheControlBuilder = CacheControlOwned::builder() +/// .public() +/// .max_age(std::time::Duration::from_secs(31536000)) +/// .immutable(); +/// let value = builder.build()?; +/// assert_eq!( +/// value.max_age(), +/// Some(std::time::Duration::from_secs(31536000)) +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct CacheControlBuilder { + directives: SmallVec<[BuilderDirective; 4]>, +} + +#[derive(Clone, Debug)] +enum BuilderDirective { + Static(&'static str), + MaxAge(Duration), + Extension { name: CompactString, value: Option }, +} + +impl CacheControlOwned { + #[cfg(all(feature = "serde", feature = "headers-cache-control"))] + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.iter().map(FieldValue::as_field_value_ref) + } + + pub(crate) fn encoded_value(&self) -> Option { + super::normalized_comma_value(&FieldName::CacheControl, self.directives().map(CacheDirectiveView::as_bytes)) + .expect("validated cache directives remain valid when comma-joined") + } + + /// Creates a builder. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::CacheControlOwned; + /// + /// let value = CacheControlOwned::builder() + /// .public() + /// .max_age(std::time::Duration::from_secs(60)) + /// .build()?; + /// assert_eq!(value.max_age(), Some(std::time::Duration::from_secs(60))); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn builder() -> CacheControlBuilder { + CacheControlBuilder { + directives: SmallVec::new_const(), + } + } + + /// Iterates all directives in wire order. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::CacheControlOwned; + /// + /// let value = CacheControlOwned::try_from("public, max-age=31536000, immutable")?; + /// let names: Vec<_> = value + /// .directives() + /// .map(|directive| directive.name()) + /// .collect(); + /// assert_eq!(names, ["public", "max-age", "immutable"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn directives(&self) -> impl Iterator> { + self.values + .iter() + .flat_map(|value| DirectiveItems::new(value.as_bytes(), true)) + .filter_map(Result::ok) + .filter_map(|item| parse_directive(item).ok()) + } + + /// Returns whether a `no-cache` directive is present. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::CacheControlOwned; + /// + /// let value = CacheControlOwned::try_from("no-cache")?; + /// assert!(value.no_cache()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn no_cache(&self) -> bool { + self.summary.no_cache + } + + /// Returns the first valid `max-age` value. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::CacheControlOwned::try_from("max-age=60")?; + /// assert_eq!(value.max_age(), Some(std::time::Duration::from_secs(60))); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn max_age(&self) -> Option { + self.summary.max_age + } +} + +impl<'a> CacheControlView<'a> { + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.repeated() + } + + /// Iterates all directives in wire order. + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::CacheControl; + /// + /// let mut map = HeaderMap::new(); + /// map.insert( + /// http::header::CACHE_CONTROL, + /// http::HeaderValue::from_static("s-maxage=120, no-store"), + /// ); + /// let view = CacheControl::view(&map)?.expect("header present"); + /// let names: Vec<_> = view + /// .directives() + /// .map(|directive| directive.name()) + /// .collect(); + /// assert_eq!(names, ["s-maxage", "no-store"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn directives(&self) -> impl Iterator> + '_ { + self.values + .comma_items() + .filter_map(Result::ok) + .filter_map(|item| parse_directive(item).ok()) + } + + /// Returns whether a `no-cache` directive is present. + #[must_use] + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::CacheControl; + /// + /// let mut map = HeaderMap::new(); + /// map.insert( + /// http::header::CACHE_CONTROL, + /// http::HeaderValue::from_static("no-cache, no-store"), + /// ); + /// let view = CacheControl::view(&map)?.expect("header present"); + /// assert!(view.no_cache()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn no_cache(&self) -> bool { + self.summary.no_cache + } + + /// Returns the first valid `max-age` value. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::CacheControlOwned::try_from("max-age=60")?; + /// assert_eq!(value.max_age(), Some(std::time::Duration::from_secs(60))); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn max_age(&self) -> Option { + self.summary.max_age + } +} + +impl<'a> CacheDirectiveView<'a> { + /// Returns the directive name. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::CacheControlOwned; + /// + /// let value = CacheControlOwned::try_from("private, max-age=60")?; + /// let directive = value.directives().next().expect("one directive"); + /// assert_eq!(directive.name(), "private"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn name(self) -> &'a str { + self.name + } + + /// Returns the optional raw token or quoted-string value. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::CacheControlOwned::try_from("max-age=60")?; + /// assert_eq!(value.max_age(), Some(std::time::Duration::from_secs(60))); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn value(self) -> Option<&'a [u8]> { + self.value + } + + /// Returns the optional value as UTF-8. + /// + /// # Errors + /// + /// Returns an error when a quoted string contains non-UTF-8 `obs-text`. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::CacheControlOwned; + /// + /// let value = CacheControlOwned::try_from("s-maxage=120")?; + /// let directive = value.directives().next().expect("one directive"); + /// assert_eq!(directive.value_str()?, Some("120")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn value_str(self) -> Result, DecodeError> { + self.value + .map(str::from_utf8) + .transpose() + .map_err(|_invalid| DecodeError::new(&FieldName::CacheControl, DecodeErrorKind::InvalidUtf8)) + } + + /// Returns the complete directive bytes after surrounding OWS trimming. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::CacheControlOwned; + /// + /// let value = CacheControlOwned::try_from("public, max-age=31536000, immutable")?; + /// let directive = value.directives().nth(1).expect("max-age directive"); + /// assert_eq!(directive.as_bytes(), b"max-age=31536000"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_bytes(self) -> &'a [u8] { + self.raw + } +} + +impl Field for CacheControl { + type View<'a> = CacheControlView<'a>; + type Owned = CacheControlOwned; + + fn name() -> &'static FieldName { + &FieldName::CacheControl + } + + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_list_item_limit(b',', true)?; + let mut summary = CacheSummary::default(); + for (value_index, value) in lines.repeated().enumerate() { + observe_directives(value.as_bytes(), value_index, &mut summary)?; + } + Ok(Some(CacheControlView { values: lines, summary })) + } + + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_list_item_limit(b',', true)?; + let mut summary = CacheSummary::default(); + let mut copied = SmallVec::new(); + for (value_index, (value, owned)) in lines.repeated_owned()?.enumerate() { + observe_directives(value.as_bytes(), value_index, &mut summary)?; + copied.push(owned); + } + Ok(Some(CacheControlOwned { values: copied, summary })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + let encoded = value.encoded_value().map_or_else(EncodedValues::new, EncodedValues::single); + sink.set_values(Self::name(), encoded) + } +} + +fn observe_directives(bytes: &[u8], value_index: usize, summary: &mut CacheSummary) -> Result<(), DecodeError> { + // Quoted directive values are rare, and the prescan for them lowers to a + // vectorised byte search, so it costs far less than watching every byte of + // the split for a quote would. + if bytes.contains(&b'"') { + return observe_quoted_directives(bytes, value_index, summary); + } + for item in bytes.split(|byte| *byte == b',') { + let item = super::trim_ows(item); + if !item.is_empty() { + summary.observe(parse_directive_parts(item)?); + } + } + Ok(()) +} + +/// Observes a line that carries quoted directive values. +#[cold] +#[inline(never)] +fn observe_quoted_directives(bytes: &[u8], value_index: usize, summary: &mut CacheSummary) -> Result<(), DecodeError> { + for item in DirectiveItems::new(bytes, true) { + let item = item.map_err(|error| error.at_value(value_index))?; + summary.observe(parse_directive_parts(item)?); + } + Ok(()) +} + +impl TryFrom<&str> for CacheControlOwned { + type Error = DecodeError; + + fn try_from(value: &str) -> Result { + let header = FieldValue::from_str(value).map_err(|_invalid| super::invalid_syntax(&FieldName::CacheControl))?; + let summary = validate_outgoing(header.as_field_value_ref())?; + Ok(Self { + values: SmallVec::from_buf([header]), + summary, + }) + } +} + +impl TryFrom for CacheControlOwned { + type Error = DecodeError; + + fn try_from(value: String) -> Result { + let header = FieldValue::try_from(value).map_err(|_invalid| super::invalid_syntax(&FieldName::CacheControl))?; + let summary = validate_outgoing(header.as_field_value_ref())?; + Ok(Self { + values: SmallVec::from_buf([header]), + summary, + }) + } +} + +impl TryFrom for CacheControlOwned { + type Error = DecodeError; + + fn try_from(header: FieldValue) -> Result { + let summary = validate_outgoing(header.as_field_value_ref())?; + Ok(Self { + values: SmallVec::from_buf([header]), + summary, + }) + } +} + +fn validate_outgoing(header: FieldValueRef<'_>) -> Result { + let mut summary = CacheSummary::default(); + for item in DirectiveItems::new(header.as_bytes(), false) { + let item = item?; + if item.is_empty() { + return Err(super::invalid_syntax(&FieldName::CacheControl)); + } + summary.observe(parse_directive_parts(item)?); + } + Ok(summary) +} + +impl CacheSummary { + fn observe(&mut self, directive: DirectiveParts<'_>) { + match directive.kind { + DirectiveKind::NoCache => self.no_cache = true, + DirectiveKind::MaxAge if self.max_age.is_none() => { + self.max_age = directive.delta_seconds.map(Duration::from_secs); + } + DirectiveKind::MaxAge | DirectiveKind::Other => {} + } + } +} + +impl CacheControlBuilder { + /// Adds `public`. + #[must_use] + pub fn public(mut self) -> Self { + self.directives.push(BuilderDirective::Static("public")); + self + } + + /// Adds `no-cache`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::CacheControlOwned; + /// + /// let value = CacheControlOwned::builder().no_cache().build()?; + /// assert!(value.no_cache()); + /// assert_eq!( + /// value.directives().next().expect("one directive").name(), + /// "no-cache" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn no_cache(mut self) -> Self { + self.directives.push(BuilderDirective::Static("no-cache")); + self + } + + /// Adds `private`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::CacheControlOwned; + /// + /// let value = CacheControlOwned::builder().private().build()?; + /// let directive = value.directives().next().expect("one directive"); + /// assert_eq!(directive.name(), "private"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn private(mut self) -> Self { + self.directives.push(BuilderDirective::Static("private")); + self + } + + /// Adds `no-store`. + #[must_use] + pub fn no_store(mut self) -> Self { + self.directives.push(BuilderDirective::Static("no-store")); + self + } + + /// Adds `must-revalidate`. + #[must_use] + pub fn must_revalidate(mut self) -> Self { + self.directives.push(BuilderDirective::Static("must-revalidate")); + self + } + + /// Adds `immutable`. + #[must_use] + pub fn immutable(mut self) -> Self { + self.directives.push(BuilderDirective::Static("immutable")); + self + } + + /// Adds `max-age`. + /// + /// Fractional seconds are rejected when building or encoding the header. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::CacheControlOwned::try_from("max-age=60")?; + /// assert_eq!(value.max_age(), Some(std::time::Duration::from_secs(60))); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn max_age(mut self, duration: Duration) -> Self { + self.directives.push(BuilderDirective::MaxAge(duration)); + self + } + + /// Adds an extension directive for validation by [`Self::build`]. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{CacheControlOwned, ExtensionValue}; + /// + /// let value = CacheControlOwned::builder() + /// .extension("s-maxage", ExtensionValue::Value("120")) + /// .build()?; + /// let directive = value.directives().next().expect("one directive"); + /// assert_eq!(directive.name(), "s-maxage"); + /// assert_eq!(directive.value_str()?, Some("120")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[must_use] + pub fn extension(mut self, name: impl AsRef, value: ExtensionValue<'_>) -> Self { + let name = name.as_ref(); + let value = match value { + ExtensionValue::Flag => None, + ExtensionValue::Value(value) => Some(CompactString::from(value)), + }; + self.directives.push(BuilderDirective::Extension { + name: CompactString::from(name), + value, + }); + self + } + + /// Adds a flag-style extension directive. + #[must_use] + pub fn extension_flag(self, name: impl AsRef) -> Self { + self.extension(name, ExtensionValue::Flag) + } + + /// Adds an extension directive carrying a token or quoted-string value. + #[must_use] + pub fn extension_value(self, name: impl AsRef, value: impl AsRef) -> Self { + self.extension(name, ExtensionValue::Value(value.as_ref())) + } + + /// Builds a nonempty header. + /// + /// # Errors + /// + /// Returns an error when no directives were added, any pending extension + /// is malformed, or a `max-age` duration contains fractional seconds. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::CacheControlOwned; + /// + /// let value = CacheControlOwned::builder() + /// .public() + /// .max_age(std::time::Duration::from_secs(31536000)) + /// .immutable() + /// .build()?; + /// let names: Vec<_> = value + /// .directives() + /// .map(|directive| directive.name()) + /// .collect(); + /// assert_eq!(names, ["public", "max-age", "immutable"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn build(self) -> Result { + if self.directives.is_empty() { + return Err(DecodeError::new(&FieldName::CacheControl, DecodeErrorKind::InvalidSyntax)); + } + self.directives.iter().try_for_each(BuilderDirective::validate)?; + let directives = self.directives; + builder_wire_len(directives.iter().map(BuilderDirective::wire_len), directives.len()).and_then(|capacity| { + let mut wire = String::with_capacity(capacity); + for (index, directive) in directives.into_iter().enumerate() { + if index != 0 { + wire.push_str(", "); + } + directive.write_to(&mut wire); + } + CacheControlOwned::try_from(wire) + }) + } +} + +impl CacheControl { + /// Starts an allocation-free response plan with `public`. + #[must_use] + pub fn public() -> CacheControlBuilder { + CacheControlOwned::builder().public() + } + + /// Starts an allocation-free response plan with `private`. + #[must_use] + pub fn private() -> CacheControlBuilder { + CacheControlOwned::builder().private() + } + + /// Starts an allocation-free response plan with `no-cache`. + #[must_use] + pub fn no_cache() -> CacheControlBuilder { + CacheControlOwned::builder().no_cache() + } +} + +impl FieldEncoder for CacheControlBuilder { + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + if self.directives.is_empty() { + return Err(InsertError::new(InsertErrorKind::InvalidValue)); + } + self.directives + .iter() + .try_for_each(BuilderDirective::validate) + .map_err(|_invalid| InsertError::new(InsertErrorKind::InvalidValue))?; + let length = self.directives.iter().map(BuilderDirective::wire_len).sum::() + self.directives.len().saturating_sub(1) * 2; + let mut writer = output.begin_value(length, FieldSensitivity::NonSensitive)?; + for (index, directive) in self.directives.into_iter().enumerate() { + if index != 0 { + writer.write_bytes(b", ")?; + } + match directive { + BuilderDirective::Static(value) => writer.write_bytes(value.as_bytes())?, + BuilderDirective::MaxAge(duration) => { + writer.write_bytes(b"max-age=")?; + write_decimal(&mut writer, duration.as_secs())?; + } + BuilderDirective::Extension { name, value } => { + writer.write_bytes(name.as_bytes())?; + if let Some(value) = value { + writer.write_bytes(b"=")?; + writer.write_bytes(value.as_bytes())?; + } + } + } + } + writer.finish() + } +} + +fn write_decimal(writer: &mut W, mut value: u64) -> Result<(), InsertError> +where + W: FieldValueWriter, +{ + let mut storage = [0_u8; 20]; + let mut start = storage.len(); + loop { + start -= 1; + storage[start] = b'0' + (value % 10) as u8; + value /= 10; + if value == 0 { + break; + } + } + writer.write_bytes(&storage[start..]) +} + +impl BuilderDirective { + fn validate(&self) -> Result<(), DecodeError> { + match self { + Self::MaxAge(duration) if duration.subsec_nanos() != 0 => { + Err(DecodeError::new(&FieldName::CacheControl, DecodeErrorKind::InvalidNumber)) + } + Self::Extension { name, value } + if !validate::token(name.as_bytes()) || value.as_ref().is_some_and(|value| !valid_directive_value(value.as_bytes())) => + { + Err(super::invalid_syntax(&FieldName::CacheControl)) + } + _ => Ok(()), + } + } + + fn wire_len(&self) -> usize { + match self { + Self::Static(value) => value.len(), + Self::Extension { name, value } => name.len() + value.as_ref().map_or(0, |value| 1 + value.len()), + Self::MaxAge(value) => "max-age=".len() + decimal_len(value.as_secs()), + } + } + + fn write_to(self, wire: &mut String) { + match self { + Self::Static(value) => wire.push_str(value), + Self::Extension { name, value } => { + wire.push_str(&name); + if let Some(value) = value { + wire.push('='); + wire.push_str(&value); + } + } + Self::MaxAge(value) => { + write!(wire, "max-age={}", value.as_secs()).expect("formatting into a String cannot fail"); + } + } + } +} + +#[inline] +fn builder_wire_len(lengths: impl IntoIterator, directive_count: usize) -> Result { + let content_len = lengths + .into_iter() + .try_fold(0_usize, usize::checked_add) + .ok_or_else(cache_size_error)?; + let separators = directive_count.saturating_sub(1).checked_mul(2).ok_or_else(cache_size_error)?; + content_len.checked_add(separators).ok_or_else(cache_size_error) +} + +#[cold] +fn cache_size_error() -> DecodeError { + DecodeError::new(&FieldName::CacheControl, DecodeErrorKind::InvalidNumber) +} + +const fn decimal_len(mut value: u64) -> usize { + let mut length = 1; + while value >= 10 { + value /= 10; + length += 1; + } + length +} + +fn parse_directive(bytes: &[u8]) -> Result, DecodeError> { + let directive = project_directive(bytes)?; + let name = str::from_utf8(directive.name).expect("validated directive names contain only ASCII"); + Ok(CacheDirectiveView { + name, + value: directive.value, + raw: bytes, + }) +} + +fn project_directive(bytes: &[u8]) -> Result, DecodeError> { + if bytes.is_empty() { + return Err(super::invalid_syntax(&FieldName::CacheControl)); + } + if bytes.get(7) == Some(&b'=') && bytes.get(..7).is_some_and(|name| eq_ignore_ascii_case_scalar(name, b"max-age")) { + let value = &bytes[8..]; + validate_directive_value(value, false)?; + return Ok(DirectiveParts { + name: &bytes[..7], + value: Some(value), + kind: DirectiveKind::MaxAge, + delta_seconds: None, + }); + } + let bare_kind = match bytes.len() { + 6 if eq_ignore_ascii_case_scalar(bytes, b"public") => Some(DirectiveKind::Other), + 7 if eq_ignore_ascii_case_scalar(bytes, b"private") => Some(DirectiveKind::Other), + 8 if eq_ignore_ascii_case_scalar(bytes, b"no-cache") => Some(DirectiveKind::NoCache), + 8 if eq_ignore_ascii_case_scalar(bytes, b"no-store") => Some(DirectiveKind::Other), + 9 if eq_ignore_ascii_case_scalar(bytes, b"immutable") => Some(DirectiveKind::Other), + 12 if eq_ignore_ascii_case_scalar(bytes, b"no-transform") => Some(DirectiveKind::Other), + 14 if eq_ignore_ascii_case_scalar(bytes, b"only-if-cached") => Some(DirectiveKind::Other), + 15 if eq_ignore_ascii_case_scalar(bytes, b"must-revalidate") => Some(DirectiveKind::Other), + 16 if eq_ignore_ascii_case_scalar(bytes, b"proxy-revalidate") => Some(DirectiveKind::Other), + _ => None, + }; + if let Some(kind) = bare_kind { + return Ok(DirectiveParts { + name: bytes, + value: None, + kind, + delta_seconds: None, + }); + } + let mut equals = None; + for (index, byte) in bytes.iter().copied().enumerate() { + if byte == b'=' { + equals = Some(index); + break; + } + if !is_token_byte(byte) { + return Err(DecodeError::new(&FieldName::CacheControl, DecodeErrorKind::InvalidToken)); + } + } + let (name, value) = if let Some(equals) = equals { + let name = &bytes[..equals]; + if name.is_empty() { + return Err(DecodeError::new(&FieldName::CacheControl, DecodeErrorKind::InvalidToken)); + } + let value = &bytes[equals + 1..]; + validate_directive_value(value, false)?; + (name, Some(value)) + } else { + (bytes, None) + }; + Ok(DirectiveParts { + name, + value, + kind: classify_directive(name), + delta_seconds: None, + }) +} + +#[derive(Clone, Copy)] +struct DirectiveParts<'a> { + name: &'a [u8], + value: Option<&'a [u8]>, + kind: DirectiveKind, + delta_seconds: Option, +} + +#[derive(Clone, Copy)] +enum DirectiveKind { + NoCache, + MaxAge, + Other, +} + +fn parse_directive_parts(bytes: &[u8]) -> Result, DecodeError> { + if bytes.is_empty() { + return Err(super::invalid_syntax(&FieldName::CacheControl)); + } + if bytes.get(7) == Some(&b'=') && bytes.get(..7).is_some_and(|name| eq_ignore_ascii_case_scalar(name, b"max-age")) { + let value = &bytes[8..]; + return Ok(DirectiveParts { + name: &bytes[..7], + value: Some(value), + kind: DirectiveKind::MaxAge, + delta_seconds: validate_directive_value(value, true)?, + }); + } + // Bare directives are a closed set in practice, and grouping the literals + // by length turns recognition into one length test plus one comparison, + // which is cheaper than the byte-at-a-time token scan below. + let bare_kind = match bytes.len() { + 6 if eq_ignore_ascii_case_scalar(bytes, b"public") => Some(DirectiveKind::Other), + 7 if eq_ignore_ascii_case_scalar(bytes, b"private") => Some(DirectiveKind::Other), + 8 if eq_ignore_ascii_case_scalar(bytes, b"no-cache") => Some(DirectiveKind::NoCache), + 8 if eq_ignore_ascii_case_scalar(bytes, b"no-store") => Some(DirectiveKind::Other), + 9 if eq_ignore_ascii_case_scalar(bytes, b"immutable") => Some(DirectiveKind::Other), + 12 if eq_ignore_ascii_case_scalar(bytes, b"no-transform") => Some(DirectiveKind::Other), + 14 if eq_ignore_ascii_case_scalar(bytes, b"only-if-cached") => Some(DirectiveKind::Other), + 15 if eq_ignore_ascii_case_scalar(bytes, b"must-revalidate") => Some(DirectiveKind::Other), + 16 if eq_ignore_ascii_case_scalar(bytes, b"proxy-revalidate") => Some(DirectiveKind::Other), + _ => None, + }; + if let Some(kind) = bare_kind { + return Ok(DirectiveParts { + name: bytes, + value: None, + kind, + delta_seconds: None, + }); + } + let mut equals = None; + for (index, byte) in bytes.iter().copied().enumerate() { + if byte == b'=' { + equals = Some(index); + break; + } + if !is_token_byte(byte) { + return Err(DecodeError::new(&FieldName::CacheControl, DecodeErrorKind::InvalidToken)); + } + } + let (name, value) = if let Some(equals) = equals { + let name = &bytes[..equals]; + if name.is_empty() { + return Err(DecodeError::new(&FieldName::CacheControl, DecodeErrorKind::InvalidToken)); + } + (name, Some(&bytes[equals + 1..])) + } else { + (bytes, None) + }; + let kind = classify_directive(name); + let delta_seconds = value + .map(|value| validate_directive_value(value, matches!(kind, DirectiveKind::MaxAge))) + .transpose()? + .flatten(); + Ok(DirectiveParts { + name, + value, + kind, + delta_seconds, + }) +} + +fn classify_directive(name: &[u8]) -> DirectiveKind { + if eq_ignore_ascii_case_scalar(name, b"no-cache") { + DirectiveKind::NoCache + } else if eq_ignore_ascii_case_scalar(name, b"max-age") { + DirectiveKind::MaxAge + } else { + DirectiveKind::Other + } +} + +fn validate_directive_value(bytes: &[u8], parse_seconds: bool) -> Result, DecodeError> { + if bytes.first() == Some(&b'"') { + validate_quoted_value(bytes, parse_seconds) + } else { + validate_token_value(bytes, parse_seconds) + } +} + +fn validate_token_value(bytes: &[u8], parse_seconds: bool) -> Result, DecodeError> { + if bytes.is_empty() { + return Err(super::invalid_syntax(&FieldName::CacheControl)); + } + // Failed decimal parsing still needs token validation, but a second + // decimal pass cannot produce a value. + if parse_seconds && let Some(seconds) = validate::decimal_u64(bytes) { + return Ok(Some(seconds)); + } + for byte in bytes.iter().copied() { + if !is_token_byte(byte) { + return Err(super::invalid_syntax(&FieldName::CacheControl)); + } + } + Ok(None) +} + +fn validate_quoted_value(bytes: &[u8], parse_seconds: bool) -> Result, DecodeError> { + if bytes.len() < 2 || bytes.first() != Some(&b'"') || bytes.last() != Some(&b'"') { + return Err(super::invalid_syntax(&FieldName::CacheControl)); + } + + let mut escaped = false; + let mut seconds = parse_seconds.then_some(0_u64); + let mut has_digit = false; + for byte in bytes[1..bytes.len() - 1].iter().copied() { + if escaped { + if !matches!(byte, b'\t' | b' '..=b'~' | 0x80..=0xff) { + return Err(super::invalid_syntax(&FieldName::CacheControl)); + } + escaped = false; + seconds = None; + } else if byte == b'\\' { + escaped = true; + seconds = None; + } else if !matches!(byte, b'\t' | b' ' | b'!' | b'#'..=b'[' | b']'..=b'~' | 0x80..=0xff) { + return Err(super::invalid_syntax(&FieldName::CacheControl)); + } else { + has_digit |= byte.is_ascii_digit(); + seconds = seconds.and_then(|value| { + byte.is_ascii_digit() + .then(|| value.checked_mul(10)?.checked_add(u64::from(byte - b'0'))) + .flatten() + }); + } + } + if escaped { + Err(super::invalid_syntax(&FieldName::CacheControl)) + } else { + Ok(seconds.filter(|_value| has_digit)) + } +} + +fn valid_directive_value(bytes: &[u8]) -> bool { + validate_directive_value(bytes, false).is_ok() +} + +fn is_token_byte(byte: u8) -> bool { + byte.is_ascii_alphanumeric() || b"!#$%&'*+-.^_`|~".contains(&byte) +} + +fn eq_ignore_ascii_case_scalar(left: &[u8], right: &[u8]) -> bool { + left.len() == right.len() && left.iter().zip(right).all(|(left, right)| left.eq_ignore_ascii_case(right)) +} + +struct DirectiveItems<'a> { + bytes: &'a [u8], + start: usize, + position: usize, + finished: bool, + skip_empty: bool, +} + +impl<'a> DirectiveItems<'a> { + const fn new(bytes: &'a [u8], skip_empty: bool) -> Self { + Self { + bytes, + start: 0, + position: 0, + finished: false, + skip_empty, + } + } +} + +impl<'a> Iterator for DirectiveItems<'a> { + type Item = Result<&'a [u8], DecodeError>; + + fn next(&mut self) -> Option { + if self.finished { + return None; + } + let mut quoted = false; + let mut escaped = false; + while let Some(byte) = self.bytes.get(self.position).copied() { + if escaped { + escaped = false; + } else if quoted && byte == b'\\' { + escaped = true; + } else if byte == b'"' { + quoted = !quoted; + } else if !quoted && byte == b',' { + let item = super::trim_ows(&self.bytes[self.start..self.position]); + self.position += 1; + self.start = self.position; + if self.skip_empty && item.is_empty() { + continue; + } + return Some(Ok(item)); + } + self.position += 1; + } + self.finished = true; + if quoted || escaped { + Some(Err(DecodeError::new(&FieldName::CacheControl, DecodeErrorKind::UnterminatedQuote))) + } else { + let item = super::trim_ows(&self.bytes[self.start..]); + if self.skip_empty && item.is_empty() { None } else { Some(Ok(item)) } + } + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #![expect( + clippy::assertions_on_result_states, + reason = "tests classify parser outcomes without needing successful values" + )] + + use std::time::Duration; + + use super::{ + CacheControl, CacheControlOwned, CacheSummary, DirectiveItems, DirectiveKind, builder_wire_len, classify_directive, decimal_len, + observe_directives, observe_quoted_directives, parse_directive, parse_directive_parts, project_directive, validate_directive_value, + }; + use crate::DecodeError; + + /// The quote-free split is a shortcut around the general item walk, so + /// both have to reach the same summary and the same diagnostic. + #[test] + fn the_unquoted_directive_split_agrees_with_the_general_walk() { + let alphabet = b"m,\"=a1 \\"; + let mut line = Vec::with_capacity(4); + for first in alphabet { + for second in alphabet { + for third in alphabet { + for fourth in alphabet { + line.clear(); + line.extend_from_slice(&[*first, *second, *third, *fourth]); + if line.contains(&b'"') { + continue; + } + + let mut split = CacheSummary::default(); + let split = observe_directives(&line, 0, &mut split).map(|()| split); + + let mut walk = CacheSummary::default(); + let walk = observe_quoted_directives(&line, 0, &mut walk).map(|()| walk); + + assert_eq!( + split.as_ref().map_err(DecodeError::kind), + walk.as_ref().map_err(DecodeError::kind), + "{:?}", + String::from_utf8_lossy(&line) + ); + } + } + } + } + } + + use crate::sink::{EncodedValues, FieldSink}; + use crate::source::FieldSource; + use crate::{DecodeErrorKind, FieldName, FieldValue, TestSink}; + + #[test] + fn directive_projection_covers_bare_and_malformed_inputs() { + assert!(project_directive(b"").is_err()); + let bare = project_directive(b"no-cache").expect("known bare directive"); + assert_eq!(bare.name, b"no-cache"); + assert_eq!(bare.value, None); + let bare = project_directive(b"no-store").unwrap(); + assert_eq!(bare.name, b"no-store"); + assert_eq!(bare.value, None); + assert!(matches!(bare.kind, DirectiveKind::Other)); + assert_eq!(project_directive(b"=value").err().unwrap().kind(), DecodeErrorKind::InvalidToken); + assert!(project_directive(b"bad value").is_err()); + } + + #[test] + fn builder_keeps_short_extensions_inline_and_spills_long_ones() { + let short = CacheControlOwned::builder().extension_value("stale-if-error", "30"); + let long = CacheControlOwned::builder().extension_value("extension-directive-with-a-value-that-exceeds-inline-storage", "enabled"); + assert!(matches!( + short.directives.as_slice(), + [super::BuilderDirective::Extension { name, value: Some(value) }] + if !name.is_heap_allocated() && !value.is_heap_allocated() + )); + assert!(matches!( + long.directives.as_slice(), + [super::BuilderDirective::Extension { name, value: Some(value) }] + if name.is_heap_allocated() && !value.is_heap_allocated() + )); + } + + #[test] + fn builder_wire_length_checks_boundaries_without_allocating() { + assert_eq!(builder_wire_len([3, 4], 2), Ok(9)); + assert_eq!(builder_wire_len([], 0), Ok(0)); + assert!(builder_wire_len([usize::MAX, 1], 2).is_err()); + assert!(builder_wire_len([0], usize::MAX).is_err()); + assert!(builder_wire_len([usize::MAX], 2).is_err()); + } + + #[test] + fn directives_views_summaries_and_debug_cover_repeated_lines() { + let mut table = TestSink::new(); + table + .set_values( + &FieldName::CacheControl, + EncodedValues::from_vec(vec![ + FieldValue::from_static("no-cache, x-mode=\"fast, safe\""), + FieldValue::from_static("max-age=30, public"), + ]), + ) + .expect("table accepts cache control"); + + let view = CacheControl::view(&table).expect("valid cache control").expect("present"); + assert!(view.no_cache()); + assert_eq!(view.max_age(), Some(Duration::from_secs(30))); + assert_eq!(view.field_values().count(), 2); + assert!(format!("{view:?}").contains("value_count")); + let directive = view + .directives() + .find(|directive| directive.name() == "x-mode") + .expect("extension present"); + assert_eq!(directive.value(), Some(b"\"fast, safe\"".as_slice())); + assert_eq!(directive.value_str(), Ok(Some("\"fast, safe\""))); + assert_eq!(directive.as_bytes(), b"x-mode=\"fast, safe\""); + + let owned = CacheControl::owned(&table).expect("valid owned cache control").expect("present"); + assert!(owned.no_cache()); + assert_eq!(owned.max_age(), Some(Duration::from_secs(30))); + assert_eq!(owned.directives().count(), 4); + assert!(format!("{owned:?}").contains("summary")); + + let mut encoded = TestSink::new(); + CacheControl::insert(&mut encoded, owned).expect("table accepts normalized value"); + assert_eq!( + encoded + .lines(&FieldName::CacheControl) + .expect("stored") + .repeated() + .next() + .expect("one line") + .as_bytes(), + b"no-cache, x-mode=\"fast, safe\", max-age=30, public" + ); + table.remove_values(&FieldName::CacheControl); + assert!(CacheControl::view(&table).expect("absence is valid").is_none()); + assert!(CacheControl::owned(&table).expect("absence is valid").is_none()); + } + + #[test] + fn builders_cover_static_numeric_extension_and_direct_encoding() { + assert!(CacheControlOwned::builder().build().is_err()); + for builder in [ + CacheControl::public(), + CacheControl::private(), + CacheControl::no_cache(), + CacheControlOwned::builder().no_store(), + CacheControlOwned::builder().must_revalidate(), + CacheControlOwned::builder().immutable(), + ] { + assert!(builder.build().is_ok()); + } + + let builder = CacheControlOwned::builder() + .public() + .no_cache() + .private() + .no_store() + .must_revalidate() + .immutable() + .max_age(Duration::from_secs(u64::MAX)) + .extension_value("stale-if-error", "\"30\"") + .extension_flag("custom"); + let built = builder.clone().build().expect("nonempty builder"); + assert_eq!(built.max_age(), Some(Duration::from_secs(u64::MAX))); + + let mut table = TestSink::new(); + table + .set_encoded(&FieldName::CacheControl, builder) + .expect("direct encoder succeeds"); + let wire = table + .lines(&FieldName::CacheControl) + .expect("stored") + .repeated() + .next() + .expect("one line"); + assert!(wire.as_bytes().starts_with(b"public, no-cache, private")); + assert!(wire.as_bytes().ends_with(b"stale-if-error=\"30\", custom")); + + assert!(table.set_encoded(&FieldName::CacheControl, CacheControlOwned::builder()).is_err()); + assert!(CacheControlOwned::builder().extension_flag("").build().is_err()); + assert!(CacheControlOwned::builder().extension_value("ok", "bad value").build().is_err()); + assert_eq!(decimal_len(0), 1); + assert_eq!(decimal_len(9), 1); + assert_eq!(decimal_len(10), 2); + assert_eq!(decimal_len(u64::MAX), 20); + } + + #[test] + fn decimal_fallback_preserves_token_validation() { + for (value, seconds) in [ + (b"0".as_slice(), Some(0)), + (b"00000000000000000000000000000000000000060", Some(60)), + (b"18446744073709551615", Some(u64::MAX)), + (b"18446744073709551616", None), + (b"18446744073709551616!", None), + (b"123x", None), + (b"invalid", None), + ] { + assert_eq!(validate_directive_value(value, true), Ok(seconds), "{value:?}"); + assert_eq!(validate_directive_value(value, false), Ok(None), "{value:?}"); + } + for value in [b"".as_slice(), b"123 x", b"18446744073709551616/", b"invalid="] { + assert_eq!( + validate_directive_value(value, true).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax, + "{value:?}" + ); + } + } + + #[test] + fn parser_covers_quotes_escapes_tokens_and_error_kinds() { + assert!(CacheControlOwned::try_from("\n").is_err()); + assert!(CacheControlOwned::try_from(String::from("\n")).is_err()); + assert!(matches!( + classify_directive(std::hint::black_box(b"MAX-AGE")), + DirectiveKind::MaxAge + )); + let borrowed = CacheControlOwned::try_from("max-age=15").expect("valid borrowed input string"); + assert_eq!(borrowed.max_age(), Some(Duration::from_secs(15))); + let owned = CacheControlOwned::try_from(String::from("MAX-AGE=\"45\", no-cache")).expect("valid owned string"); + assert_eq!(owned.max_age(), Some(Duration::from_secs(45))); + assert!(CacheControlOwned::try_from(FieldValue::from_static("public")).is_ok()); + assert_eq!( + "public".parse::().expect("shared FromStr").directives().count(), + 1 + ); + + let obs = parse_directive(b"x-note=\"\xff\"").expect("valid obs-text quoted value"); + assert_eq!( + obs.value_str().expect_err("obs-text is not UTF-8").kind(), + DecodeErrorKind::InvalidUtf8 + ); + assert!(parse_directive(&[0xff]).is_err()); + + for valid in [ + b"public".as_slice(), + b"must-revalidate", + b"no-cache", + b"max-age=12", + b"max-age=\"12\"", + b"extension=value", + b"extension=\"a,b\"", + b"extension=\"a\\\\b\"", + ] { + assert!(parse_directive_parts(valid).is_ok(), "{valid:?}"); + } + for invalid in [ + b"".as_slice(), + b"=value", + b"bad name", + b"max-age=", + b"x=\"unterminated", + b"x=\"bad\\\n\"", + b"x=\"trailing\\\"", + ] { + assert!(parse_directive_parts(invalid).is_err(), "{invalid:?}"); + } + assert!( + parse_directive_parts(b"max-age=12x") + .expect("token value remains syntactically valid") + .delta_seconds + .is_none() + ); + assert_eq!(validate_directive_value(b"123", true), Ok(Some(123))); + assert_eq!(validate_directive_value(b"abc", true), Ok(None)); + assert_eq!(validate_directive_value(b"\"\"", true), Ok(None)); + assert_eq!(validate_directive_value(b"\"123\"", true), Ok(Some(123))); + assert!(parse_directive_parts(b"no-cache=value").is_ok()); + + let mut items = DirectiveItems::new(b", one, \"two,three\",", true); + assert_eq!(items.next(), Some(Ok(b"one".as_slice()))); + assert_eq!(items.next(), Some(Ok(b"\"two,three\"".as_slice()))); + assert_eq!(items.next(), None); + assert_eq!(items.next(), None); + assert!( + DirectiveItems::new(b"one, \"unterminated", false) + .last() + .expect("error item") + .is_err() + ); + + assert!(CacheControlOwned::try_from("public,,private").is_err()); + let mut table = TestSink::new(); + table + .set_values( + &FieldName::CacheControl, + EncodedValues::single(FieldValue::from_static("x=\"unterminated")), + ) + .expect("table accepts raw field value"); + assert_eq!( + CacheControl::view(&table).expect_err("unterminated quote must fail").value_index(), + Some(0) + ); + assert!(CacheControlOwned::try_from("").is_err()); + assert!(parse_directive_parts(b"x=\"bad\n\"").is_err()); + } +} diff --git a/crates/http_headers/src/headers/conditional/if_match.rs b/crates/http_headers/src/headers/conditional/if_match.rs new file mode 100644 index 000000000..f2d7fa0b8 --- /dev/null +++ b/crates/http_headers/src/headers/conditional/if_match.rs @@ -0,0 +1,24 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::shared::{ + ConditionalTagView, TagIter, TagListState, entity_tag_list_header, validate_tag_line_outlined_with, validate_tag_line_with, + validate_tag_slice, +}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef}; + +entity_tag_list_header!( + IfMatch, + IfMatchOwned, + IfMatchView, + "If-Match", + &FieldName::IfMatch, + "Owned value for the `If-Match` header.", + "Borrowed value for the `If-Match` header.", + "Defined by [RFC 9110 section 13.1.1](https://www.rfc-editor.org/rfc/rfc9110#section-13.1.1).", + "`If-Match: *` matches any current representation. `If-Match: \"xyzzy\", \"revision-42\"` lists strong entity tags." +); diff --git a/crates/http_headers/src/headers/conditional/if_modified_since.rs b/crates/http_headers/src/headers/conditional/if_modified_since.rs new file mode 100644 index 000000000..6eb6c3683 --- /dev/null +++ b/crates/http_headers/src/headers/conditional/if_modified_since.rs @@ -0,0 +1,20 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::time::SystemTime; + +use super::shared::{date_header, parse_http_date}; +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +date_header!( + IfModifiedSince, + IfModifiedSinceOwned, + IfModifiedSinceView, + "If-Modified-Since", + &FieldName::IfModifiedSince, + "Owned value for the `If-Modified-Since` header.", + "Borrowed value for the `If-Modified-Since` header.", + "Defined by [RFC 9110 section 13.1.3](https://www.rfc-editor.org/rfc/rfc9110#section-13.1.3).", + "`If-Modified-Since: Tue, 15 Nov 1994 08:12:31 GMT` supplies an HTTP date." +); diff --git a/crates/http_headers/src/headers/conditional/if_none_match.rs b/crates/http_headers/src/headers/conditional/if_none_match.rs new file mode 100644 index 000000000..2b14fbb87 --- /dev/null +++ b/crates/http_headers/src/headers/conditional/if_none_match.rs @@ -0,0 +1,24 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::shared::{ + ConditionalTagView, TagIter, TagListState, entity_tag_list_header, validate_tag_line_outlined_with, validate_tag_line_with, + validate_tag_slice, +}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef}; + +entity_tag_list_header!( + IfNoneMatch, + IfNoneMatchOwned, + IfNoneMatchView, + "If-None-Match", + &FieldName::IfNoneMatch, + "Owned value for the `If-None-Match` header.", + "Borrowed value for the `If-None-Match` header.", + "Defined by [RFC 9110 section 13.1.2](https://www.rfc-editor.org/rfc/rfc9110#section-13.1.2).", + "`If-None-Match: *` matches any current representation. `If-None-Match: \"strong\", W/\"weak\"` demonstrates strong and weak tags." +); diff --git a/crates/http_headers/src/headers/conditional/if_range.rs b/crates/http_headers/src/headers/conditional/if_range.rs new file mode 100644 index 000000000..c69302b88 --- /dev/null +++ b/crates/http_headers/src/headers/conditional/if_range.rs @@ -0,0 +1,497 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::time::SystemTime; + +use super::shared::{ConditionalTagView, format_http_date, parse_http_date_with}; +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +/// The parsed alternative represented by `If-Range`. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::IfRangeOwned::try_from("\"revision\"")?; +/// assert!(matches!( +/// value.value()?, +/// http_headers::headers::IfRangeValueView::EntityTag(_) +/// )); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub enum IfRangeValueView<'a> { + /// A strong entity tag. + EntityTag(ConditionalTagView<'a>), + /// An HTTP date. + Date(SystemTime), +} + +/// Defines the `If-Range` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 13.1.5](https://www.rfc-editor.org/rfc/rfc9110#section-13.1.5). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{IfRange, IfRangeOwned}; +/// +/// let mut map = HeaderMap::new(); +/// IfRange::insert(&mut map, IfRangeOwned::try_from("\"revision\"")?)?; +/// assert!(IfRange::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct IfRange { + _private: (), +} + +/// Owned value for the `If-Range` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 13.1.5]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::IfRangeOwned::try_from("\"revision\"")?; +/// assert!(matches!( +/// value.value()?, +/// http_headers::headers::IfRangeValueView::EntityTag(_) +/// )); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `If-Range: "revision-42"` carries a strong entity tag, while +/// `If-Range: Wed, 21 Oct 2015 07:28:00 GMT` carries an HTTP date. +/// +/// [RFC 9110 section 13.1.5]: https://www.rfc-editor.org/rfc/rfc9110#section-13.1.5 +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct IfRangeOwned { + value: FieldValue, + parsed: IfRangeMetadata, +} + +/// Borrowed value for the `If-Range` header. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{IfRange, IfRangeValueView, IfRangeView}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let value: IfRangeView<'_> = IfRange::decode_view(FieldValueRef::new(b"\"revision\""))?; +/// assert!(matches!( +/// value.value(), +/// IfRangeValueView::EntityTag(tag) if tag.opaque_tag() == b"revision" +/// )); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct IfRangeView<'a> { + value: FieldValueRef<'a>, + parsed: IfRangeValueView<'a>, +} + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +enum IfRangeMetadata { + EntityTag, + Date(SystemTime), +} + +impl IfRangeMetadata { + const fn from_view(value: IfRangeValueView<'_>) -> Self { + match value { + IfRangeValueView::EntityTag(_) => Self::EntityTag, + IfRangeValueView::Date(date) => Self::Date(date), + } + } +} + +impl fmt::Debug for IfRangeOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("IfRangeOwned") + .field("value", &self.parsed_value()) + .finish_non_exhaustive() + } +} + +impl IfRangeOwned { + /// Constructs an `If-Range` value from a strong entity tag. + /// + /// # Errors + /// + /// Returns an error for a weak tag. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ETagOwned, IfRangeOwned, IfRangeValueView}; + /// + /// let tag = ETagOwned::try_from("\"revision\"")?; + /// let value = IfRangeOwned::entity_tag(tag)?; + /// assert!(matches!( + /// value.value()?, + /// IfRangeValueView::EntityTag(tag) if tag.opaque_tag() == b"revision" + /// )); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn entity_tag(tag: crate::headers::ETagOwned) -> Result { + if tag.is_weak() { + return Err(crate::headers::invalid_syntax(&FieldName::IfRange)); + } + Ok(Self { + value: tag.into_field_value(), + parsed: IfRangeMetadata::EntityTag, + }) + } + + /// Constructs a canonical IMF-fixdate alternative. + /// + /// Fractional seconds are discarded to match the whole-second wire precision. + /// + /// # Errors + /// + /// Returns an error when the date cannot be represented by `httpdate`. + /// # Examples + /// + /// ```rust + /// use std::time::{Duration, UNIX_EPOCH}; + /// + /// use http_headers::headers::{IfRangeOwned, IfRangeValueView}; + /// + /// let instant = UNIX_EPOCH + Duration::from_secs(784_111_777); + /// let value = IfRangeOwned::date(instant)?; + /// assert!(matches!(value.value()?, IfRangeValueView::Date(date) if date == instant)); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn date(date: SystemTime) -> Result { + format_http_date(&FieldName::IfRange, date).and_then(Self::try_from) + } + + /// Returns the parsed alternative. + /// + /// # Errors + /// + /// Returns an error if internal preserved offsets no longer agree with + /// the field value. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::IfRangeOwned::try_from("\"revision\"")?; + /// assert!(matches!( + /// value.value()?, + /// http_headers::headers::IfRangeValueView::EntityTag(_) + /// )); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::unnecessary_wraps, + reason = "the fallible signature is retained for compatibility with the pre-1.0 accessor" + )] + pub fn value(&self) -> Result, DecodeError> { + Ok(self.parsed_value()) + } + + fn parsed_value(&self) -> IfRangeValueView<'_> { + match self.parsed { + IfRangeMetadata::EntityTag => IfRangeValueView::EntityTag(ConditionalTagView { + wire: self.value.as_bytes(), + }), + IfRangeMetadata::Date(date) => IfRangeValueView::Date(date), + } + } + + /// Returns the preserved field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::IfRangeOwned; + /// + /// let value = IfRangeOwned::try_from("\"revision\"")?; + /// assert_eq!(value.as_field_value().as_bytes(), b"\"revision\""); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Consumes the header and returns its field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::IfRangeOwned; + /// + /// let value = IfRangeOwned::try_from("Sat, 29 Oct 1994 19:43:31 GMT")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"Sat, 29 Oct 1994 19:43:31 GMT"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(IfRangeOwned, |value| value.value); + +impl<'a> IfRangeView<'a> { + /// Returns the parsed alternative. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::IfRangeOwned::try_from("\"revision\"")?; + /// assert!(matches!( + /// value.value()?, + /// http_headers::headers::IfRangeValueView::EntityTag(_) + /// )); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn value(self) -> IfRangeValueView<'a> { + self.parsed + } + + /// Returns the borrowed field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::IfRange; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let value = IfRange::decode_view(FieldValueRef::new(b"\"revision\""))?; + /// assert_eq!(value.as_field_value().as_bytes(), b"\"revision\""); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } +} + +impl SingleValueField for IfRange { + type View<'a> = IfRangeView<'a>; + type Owned = IfRangeOwned; + + fn name() -> &'static FieldName { + &FieldName::IfRange + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + Ok(IfRangeView { + value, + parsed: parse_if_range(value)?, + }) + } + + fn decode_owned(value: FieldValue) -> Result { + IfRangeOwned::try_from(value) + } + + fn decode_view_with(value: FieldValueRef<'_>, mode: crate::DecodeMode) -> Result, DecodeError> { + Ok(IfRangeView { + value, + parsed: parse_if_range_with(value, mode)?, + }) + } + + fn decode_owned_with(value: FieldValue, mode: crate::DecodeMode) -> Result { + let parsed = IfRangeMetadata::from_view(parse_if_range_with(value.as_field_value_ref(), mode)?); + Ok(IfRangeOwned { value, parsed }) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +super::super::shared::impl_string_conversions!(IfRangeOwned, &FieldName::IfRange, crate::headers::invalid_syntax, wire); + +impl TryFrom for IfRangeOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + let parsed = IfRangeMetadata::from_view(parse_if_range(value.as_field_value_ref())?); + Ok(Self { value, parsed }) + } +} + +/// Returns whether every byte of `opaque` is a legal `etagc`. +/// +/// Kept out of line so that the `If-Range` decoders stay small enough for the +/// caller to inline them, which measures better than inlining the scan itself. +#[inline(never)] +fn opaque_is_valid(opaque: &[u8]) -> bool { + crate::headers::etag::opaque_is_valid(opaque) +} + +fn parse_if_range(value: FieldValueRef<'_>) -> Result, DecodeError> { + parse_if_range_with(value, crate::DecodeMode::Strict) +} + +fn parse_if_range_with(value: FieldValueRef<'_>, mode: crate::DecodeMode) -> Result, DecodeError> { + let wire = value.as_bytes(); + if let [b'"', opaque @ .., b'"'] = wire { + return if opaque_is_valid(opaque) { + Ok(IfRangeValueView::EntityTag(ConditionalTagView { wire })) + } else { + Err(crate::headers::invalid_syntax(&FieldName::IfRange)) + }; + } + parse_if_range_date_with(value, mode) +} + +/// Handles every `If-Range` alternative that is not a complete strong tag. +/// +/// A weak validator is never a legal alternative, and it can never be an +/// HTTP-date either. Anything else, including a truncated quoted tag, is left +/// to the date parser, which rejects it with the same error. +fn parse_if_range_date_with(value: FieldValueRef<'_>, mode: crate::DecodeMode) -> Result, DecodeError> { + if value.as_bytes().starts_with(b"W/") { + return Err(crate::headers::invalid_syntax(&FieldName::IfRange)); + } + parse_http_date_with(&FieldName::IfRange, value, mode).map(IfRangeValueView::Date) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::collections::hash_map::DefaultHasher; + use std::hash::{Hash, Hasher}; + use std::time::{Duration, UNIX_EPOCH}; + + use super::{IfRange, IfRangeOwned, IfRangeValueView}; + use crate::headers::ETagOwned; + use crate::{DecodeErrorKind, DecodeMode, FieldName, FieldValue, SingleValueField}; + + fn hash(value: &impl Hash) -> u64 { + let mut hasher = DefaultHasher::new(); + value.hash(&mut hasher); + hasher.finish() + } + + #[test] + fn date_construction_matches_wire_precision_and_round_trip_identity() { + for seconds in [0, 1, 784_111_776, 784_111_777, 253_402_300_799] { + let whole = UNIX_EPOCH + Duration::from_secs(seconds); + let canonical = IfRangeOwned::date(whole).unwrap(); + for nanos in [0, 1, 500_000_000, 999_999_999] { + let constructed = IfRangeOwned::date(whole + Duration::from_nanos(nanos)).unwrap(); + assert_eq!(constructed.value().unwrap(), IfRangeValueView::Date(whole)); + assert_eq!(constructed, canonical); + assert_eq!(hash(&constructed), hash(&canonical)); + + let wire = constructed.as_field_value(); + let parsed = IfRangeOwned::try_from(wire.try_as_str().unwrap()).unwrap(); + assert_eq!(constructed, parsed); + assert_eq!(hash(&constructed), hash(&parsed)); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let view = IfRange::decode_view_with(wire.as_field_value_ref(), mode).unwrap(); + let owned = IfRange::decode_owned_with(wire.clone(), mode).unwrap(); + assert_eq!(view.value(), constructed.value().unwrap()); + assert_eq!(hash(&view.value()), hash(&constructed.value().unwrap())); + assert_eq!(owned, constructed); + assert_eq!(hash(&owned), hash(&constructed)); + } + + #[cfg(all(feature = "serde", feature = "headers-conditional"))] + { + let serialized = serde_json::to_string(&constructed).unwrap(); + let decoded: IfRangeOwned = serde_json::from_str(&serialized).unwrap(); + assert_eq!(decoded, constructed); + assert_eq!(decoded.value().unwrap(), IfRangeValueView::Date(whole)); + assert_eq!(hash(&decoded), hash(&constructed)); + } + } + } + } + + #[test] + fn date_constructor_retains_representable_bounds() { + // Windows SystemTime uses 100 ns intervals. + let before_epoch = UNIX_EPOCH - Duration::from_micros(1); + assert!(before_epoch < UNIX_EPOCH); + for instant in [ + before_epoch, + UNIX_EPOCH - Duration::from_secs(1), + UNIX_EPOCH + Duration::from_hours(70_389_528), + ] { + assert_eq!(IfRangeOwned::date(instant).unwrap_err().kind(), DecodeErrorKind::InvalidNumber); + } + } + + #[test] + fn constructors_accessors_and_trait_paths_preserve_the_selected_alternative() { + let tag = ETagOwned::try_from("\"revision\"").expect("strong tag"); + let tagged = IfRangeOwned::entity_tag(tag).expect("If-Range strong tag"); + let parsed = tagged.value().expect("tag value"); + let opaque = |value| match value { + IfRangeValueView::EntityTag(view) => Some(view.opaque_tag()), + IfRangeValueView::Date(_) => None, + }; + assert_eq!(opaque(parsed), Some(b"revision".as_slice())); + assert_eq!(tagged.as_field_value().as_bytes(), b"\"revision\""); + assert!(format!("{tagged:?}").contains("EntityTag")); + assert_eq!( + IfRangeOwned::try_from(String::from("\"owned\"")) + .expect("owned tag string") + .as_field_value(), + "\"owned\"" + ); + + let instant = UNIX_EPOCH + Duration::from_secs(784_111_777); + let dated = IfRangeOwned::date(instant).expect("representable date"); + assert_eq!(dated.value().expect("date value"), IfRangeValueView::Date(instant)); + assert_eq!(opaque(dated.value().expect("date value")), None); + + let view = ::decode_view(tagged.as_field_value().as_field_value_ref()).expect("tag view"); + assert_eq!(view.as_field_value(), tagged.as_field_value().as_field_value_ref()); + assert!(matches!(view.value(), IfRangeValueView::EntityTag(_))); + assert_eq!(::name(), &FieldName::IfRange); + + let owned = ::decode_owned(tagged.clone().into_field_value()).expect("owned tag"); + assert_eq!(::as_field_value(&owned), tagged.as_field_value()); + assert_eq!(::into_field_value(owned), tagged.into_field_value()); + } + + #[test] + fn weak_malformed_and_relaxed_date_inputs_are_distinguished() { + let weak = ETagOwned::try_from("W/\"weak\"").expect("weak tag"); + let error = IfRangeOwned::entity_tag(weak).expect_err("weak If-Range tag"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + + for wire in ["W/\"weak\"", "\"bad tag\"", "\"unterminated", "not a date"] { + let error = IfRangeOwned::try_from(wire).expect_err("invalid If-Range"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax, "{wire:?}"); + } + let error = IfRangeOwned::try_from(String::from("bad\nvalue")).expect_err("invalid field value"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!( + IfRangeOwned::try_from("bad\nvalue") + .expect_err("invalid borrowed field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let relaxed = FieldValue::from_static(" Tue, 8 Nov 1994 8:49:37 UTC "); + ::decode_view(relaxed.as_field_value_ref()).expect_err("strict date"); + let view = + ::decode_view_with(relaxed.as_field_value_ref(), DecodeMode::Relaxed).expect("relaxed date"); + assert!(matches!(view.value(), IfRangeValueView::Date(_))); + let owned = ::decode_owned_with(relaxed, DecodeMode::Relaxed).expect("relaxed owned date"); + assert_eq!(owned.as_field_value().as_bytes(), b" Tue, 8 Nov 1994 8:49:37 UTC "); + assert!(matches!(owned.value(), Ok(IfRangeValueView::Date(_)))); + } +} diff --git a/crates/http_headers/src/headers/conditional/if_unmodified_since.rs b/crates/http_headers/src/headers/conditional/if_unmodified_since.rs new file mode 100644 index 000000000..2ae5bc987 --- /dev/null +++ b/crates/http_headers/src/headers/conditional/if_unmodified_since.rs @@ -0,0 +1,20 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::time::SystemTime; + +use super::shared::{date_header, parse_http_date}; +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +date_header!( + IfUnmodifiedSince, + IfUnmodifiedSinceOwned, + IfUnmodifiedSinceView, + "If-Unmodified-Since", + &FieldName::IfUnmodifiedSince, + "Owned value for the `If-Unmodified-Since` header.", + "Borrowed value for the `If-Unmodified-Since` header.", + "Defined by [RFC 9110 section 13.1.4](https://www.rfc-editor.org/rfc/rfc9110#section-13.1.4).", + "`If-Unmodified-Since: Sat, 29 Oct 2022 19:43:31 GMT` supplies an HTTP date." +); diff --git a/crates/http_headers/src/headers/conditional/last_modified.rs b/crates/http_headers/src/headers/conditional/last_modified.rs new file mode 100644 index 000000000..c4b75cb73 --- /dev/null +++ b/crates/http_headers/src/headers/conditional/last_modified.rs @@ -0,0 +1,20 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::time::SystemTime; + +use super::shared::{date_header, parse_http_date}; +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +date_header!( + LastModified, + LastModifiedOwned, + LastModifiedView, + "Last-Modified", + &FieldName::LastModified, + "Owned value for the `Last-Modified` header.", + "Borrowed value for the `Last-Modified` header.", + "Defined by [RFC 9110 section 8.8.2](https://www.rfc-editor.org/rfc/rfc9110#section-8.8.2).", + "`Last-Modified: Thu, 01 Jan 2015 12:00:00 GMT` supplies the selected representation's modification time." +); diff --git a/crates/http_headers/src/headers/conditional/mod.rs b/crates/http_headers/src/headers/conditional/mod.rs new file mode 100644 index 000000000..d042dbd76 --- /dev/null +++ b/crates/http_headers/src/headers/conditional/mod.rs @@ -0,0 +1,27 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Conditional request and HTTP-date header types. + +mod if_match; +mod if_modified_since; +mod if_none_match; +mod if_range; +mod if_unmodified_since; +mod last_modified; +mod shared; + +#[doc(inline)] +pub use if_match::{IfMatch, IfMatchOwned, IfMatchView}; +#[doc(inline)] +pub use if_modified_since::{IfModifiedSince, IfModifiedSinceOwned, IfModifiedSinceView}; +#[doc(inline)] +pub use if_none_match::{IfNoneMatch, IfNoneMatchOwned, IfNoneMatchView}; +#[doc(inline)] +pub use if_range::{IfRange, IfRangeOwned, IfRangeValueView, IfRangeView}; +#[doc(inline)] +pub use if_unmodified_since::{IfUnmodifiedSince, IfUnmodifiedSinceOwned, IfUnmodifiedSinceView}; +#[doc(inline)] +pub use last_modified::{LastModified, LastModifiedOwned, LastModifiedView}; +#[doc(inline)] +pub use shared::ConditionalTagView; diff --git a/crates/http_headers/src/headers/conditional/shared.rs b/crates/http_headers/src/headers/conditional/shared.rs new file mode 100644 index 000000000..d734b7257 --- /dev/null +++ b/crates/http_headers/src/headers/conditional/shared.rs @@ -0,0 +1,1893 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::slice; +use std::time::{Duration, SystemTime, UNIX_EPOCH}; + +use http_headers_simd::ascii_str; + +use crate::sink::EncodedValues; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef}; + +/// Last second representable by IMF-fixdate: 9999-12-31 23:59:59 UTC. +const MAX_HTTP_DATE_SECONDS: u64 = 253_402_300_799; + +/// A borrowed entity tag from a conditional request field. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ``` +/// use http_headers::headers::{ConditionalTagView, IfMatchOwned}; +/// +/// let value = IfMatchOwned::try_from("\"revision\"")?; +/// let tag: ConditionalTagView<'_> = value.tags().next().expect("tag"); +/// assert_eq!(tag.as_bytes(), b"\"revision\""); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct ConditionalTagView<'a> { + pub(super) wire: &'a [u8], +} + +#[derive(Clone, Eq, Hash, PartialEq)] +pub(super) enum TagValues { + One(FieldValue), + Many(Vec), +} + +pub(super) trait OwnedTagSource { + fn next_tag(&mut self) -> Option; + fn remaining_hint(&self) -> usize; +} + +pub(super) struct OwnedTagSourceAdapter(pub(super) I); + +impl OwnedTagSource for OwnedTagSourceAdapter +where + I: Iterator, +{ + // LLVM emits an uncallable polymorphized instance for this type adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + fn next_tag(&mut self) -> Option { + self.0.next() + } + + fn remaining_hint(&self) -> usize { + self.0.size_hint().0 + } +} + +pub(super) fn tag_values_from_source(name: &'static FieldName, tags: &mut dyn OwnedTagSource) -> Result { + let Some(first) = tags.next_tag() else { + return Err(DecodeError::new(name, DecodeErrorKind::MissingValue)); + }; + let first = first.into_field_value(); + let Some(second) = tags.next_tag() else { + return Ok(TagValues::One(first)); + }; + let mut values = Vec::with_capacity(2_usize.saturating_add(tags.remaining_hint()).min(128)); + values.push(first); + values.push(second.into_field_value()); + while let Some(tag) = tags.next_tag() { + values.push(tag.into_field_value()); + } + Ok(TagValues::Many(values)) +} + +impl TagValues { + pub(super) fn len(&self) -> usize { + match self { + Self::One(_) => 1, + Self::Many(values) => values.len(), + } + } + + pub(super) fn iter(&self) -> slice::Iter<'_, FieldValue> { + match self { + Self::One(value) => slice::from_ref(value).iter(), + Self::Many(values) => values.iter(), + } + } + + pub(super) fn into_encoded(self) -> EncodedValues { + match self { + Self::One(value) => EncodedValues::single(value), + Self::Many(values) => EncodedValues::from_vec(values), + } + } +} + +impl<'a> ConditionalTagView<'a> { + /// Returns the complete wire-format entity tag. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfMatchOwned; + /// + /// let value = IfMatchOwned::try_from("\"revision\"")?; + /// let tag = value.tags().next().expect("tag"); + /// assert_eq!(tag.as_bytes(), b"\"revision\""); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_bytes(self) -> &'a [u8] { + self.wire + } + + /// Returns the unquoted opaque tag. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfMatchOwned; + /// + /// let value = IfMatchOwned::try_from("W/\"revision\"")?; + /// let tag = value.tags().next().expect("tag"); + /// assert_eq!(tag.opaque_tag(), b"revision"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn opaque_tag(self) -> &'a [u8] { + let start = if self.is_weak() { 3 } else { 1 }; + let (_, opaque_and_quote) = self.wire.split_at(start); + let (opaque, _) = opaque_and_quote.split_at(opaque_and_quote.len() - 1); + opaque + } + + /// Returns whether the tag uses the weak validator prefix. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfMatchOwned; + /// + /// let value = IfMatchOwned::try_from("W/\"revision\"")?; + /// let tag = value.tags().next().expect("tag"); + /// assert!(tag.is_weak()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn is_weak(self) -> bool { + matches!(self.wire.first(), Some(b'W' | b'w')) + } +} + +macro_rules! entity_tag_list_header { + ( + $descriptor:ident, + $owned:ident, + $borrowed:ident, + $header_name:literal, + $constant:expr, + $doc:literal, + $view_doc:literal, + $specification:literal, + $examples:literal + ) => { + #[doc = concat!("Defines the `", $header_name, "` header.")] + #[doc = ""] + #[doc = "# Specification"] + #[doc = ""] + #[doc = $specification] + #[derive(Debug)] + pub struct $descriptor { + _private: (), + } + + #[doc = $doc] + #[doc = ""] + #[doc = "# Specification"] + #[doc = ""] + #[doc = $specification] + #[doc = ""] + #[doc = "# Examples"] + #[doc = ""] + #[doc = $examples] + #[derive(Clone, Eq, Hash, PartialEq)] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfMatchOwned; + /// + /// let value = IfMatchOwned::try_from("\"revision\"")?; + /// let rendered = format!("{value:?}"); + /// assert!(rendered.contains("IfMatchOwned")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub struct $owned { + values: super::shared::TagValues, + wildcard: bool, + } + + #[doc = $view_doc] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::IfMatch; + /// + /// let mut map = HeaderMap::new(); + /// map.insert("if-match", http::HeaderValue::from_static("\"revision\"")); + /// let value = IfMatch::view(&map)?.expect("present"); + /// let rendered = format!("{value:?}"); + /// assert!(rendered.contains("IfMatchView")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub struct $borrowed<'a> { + values: FieldLines<'a>, + wildcard: bool, + } + + impl fmt::Debug for $owned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct(stringify!($owned)) + .field("value_count", &self.values.len()) + .field("wildcard", &self.wildcard) + .finish() + } + } + + impl fmt::Debug for $borrowed<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct(stringify!($borrowed)) + .field("value_count", &self.values.len()) + .field("wildcard", &self.wildcard) + .finish() + } + } + + impl $owned { + #[cfg(all(feature = "serde", feature = "headers-conditional"))] + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.iter().map(FieldValue::as_field_value_ref) + } + + /// Constructs the wildcard form. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfMatchOwned; + /// + /// let value = IfMatchOwned::wildcard(); + /// assert!(value.is_wildcard()); + /// assert_eq!(value.tags().count(), 0); + /// ``` + pub fn wildcard() -> Self { + Self { + values: super::shared::TagValues::One(FieldValue::from_static("*")), + wildcard: true, + } + } + + /// Constructs a tag list without combining its field lines. + /// + /// # Errors + /// + /// Returns an error when no tags are supplied. + /// # Examples + /// + /// ``` + /// use http_headers::headers::{ETagOwned, IfMatchOwned}; + /// + /// let value = IfMatchOwned::from_tags([ETagOwned::strong("first")?, ETagOwned::weak("second")?])?; + /// assert_eq!(value.tags().count(), 2); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + // LLVM emits an uncallable polymorphized instance for this adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + pub fn from_tags(tags: I) -> Result + where + I: IntoIterator, + { + let mut tags = super::shared::OwnedTagSourceAdapter(tags.into_iter()); + let stored = super::shared::tag_values_from_source($constant, &mut tags)?; + Ok(Self { + values: stored, + wildcard: false, + }) + } + + /// Returns whether this is the wildcard form. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfMatchOwned; + /// + /// let value = IfMatchOwned::wildcard(); + /// assert!(value.is_wildcard()); + /// ``` + pub const fn is_wildcard(&self) -> bool { + self.wildcard + } + + /// Iterates entity tags in field-line and list order. + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfMatchOwned; + /// + /// let value = IfMatchOwned::try_from("\"first\", W/\"second\"")?; + /// let tags: Vec<_> = value.tags().map(|tag| tag.as_bytes()).collect(); + /// assert_eq!( + /// tags, + /// vec![b"\"first\"".as_slice(), b"W/\"second\"".as_slice()], + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn tags(&self) -> impl Iterator> { + TagIter::new(self.values.iter().map(FieldValue::as_field_value_ref)) + } + } + + impl<'a> $borrowed<'a> { + /// Returns whether this is the wildcard form. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::IfMatch; + /// + /// let mut map = HeaderMap::new(); + /// map.insert("if-match", http::HeaderValue::from_static("*")); + /// let value = IfMatch::view(&map)?.expect("present"); + /// assert!(value.is_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub const fn is_wildcard(&self) -> bool { + self.wildcard + } + + /// Iterates entity tags in field-line and list order. + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::IfMatch; + /// + /// let mut map = HeaderMap::new(); + /// map.insert( + /// "if-match", + /// http::HeaderValue::from_static("\"first\", W/\"second\""), + /// ); + /// let value = IfMatch::view(&map)?.expect("present"); + /// let tags: Vec<_> = value.tags().map(|tag| tag.opaque_tag()).collect(); + /// assert_eq!(tags, vec![b"first".as_slice(), b"second".as_slice()]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn tags(&self) -> impl Iterator> + '_ { + TagIter::new(self.values.repeated()) + } + } + + impl Field for $descriptor { + type View<'a> = $borrowed<'a>; + type Owned = $owned; + + fn name() -> &'static FieldName { + $constant + } + + fn view_with(source: &S, mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_entity_tag_item_limit()?; + // The walk is spelled out here rather than shared with a + // helper so that the validated view is assembled in place. + let wildcard = { + let mut lines = lines.repeated(); + let first = lines.next().expect("FieldLines always contains at least one field line"); + let mut state = validate_tag_line_outlined_with(first.as_bytes(), TagListState::EMPTY, mode) + .ok_or_else(|| $crate::headers::invalid_syntax($constant).at_value(0))?; + let mut value_index = 1; + for value in lines { + state = validate_tag_line_outlined_with(value.as_bytes(), state, mode) + .ok_or_else(|| $crate::headers::invalid_syntax($constant).at_value(value_index))?; + value_index += 1; + } + state.finish($constant)? + }; + Ok(Some($borrowed { values: lines, wildcard })) + } + + fn owned_with(source: &S, mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + // Repeating this header over several field lines is rare, so + // that walk lives out of line and leaves the one-line path + // sole owner of the returned value. + #[cold] + #[inline(never)] + fn many<'a>( + first: FieldValue, + second: FieldValueRef<'a>, + second_owned: FieldValue, + state: TagListState, + lines: impl Iterator, FieldValue)>, + mode: crate::DecodeMode, + ) -> Result<$owned, DecodeError> { + let mut state = validate_tag_line_with(second.as_bytes(), state, mode) + .ok_or_else(|| $crate::headers::invalid_syntax($constant).at_value(1))?; + let mut storage = Vec::with_capacity(lines.size_hint().0.saturating_add(2)); + storage.push(first); + storage.push(second_owned); + let mut value_index = 2; + for (value, owned) in lines { + state = validate_tag_line_with(value.as_bytes(), state, mode) + .ok_or_else(|| $crate::headers::invalid_syntax($constant).at_value(value_index))?; + storage.push(owned); + value_index += 1; + } + Ok($owned { + values: super::shared::TagValues::Many(storage), + wildcard: state.finish($constant)?, + }) + } + + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_entity_tag_item_limit()?; + // One field line is the overwhelmingly common shape, so the + // first line is validated and captured before the walk decides + // whether any storage beyond the inline slot is needed. + let mut lines = lines.repeated_owned()?; + let (first, first_owned) = lines.next().expect("FieldLines always contains at least one field line"); + let state = validate_tag_line_with(first.as_bytes(), TagListState::EMPTY, mode) + .ok_or_else(|| $crate::headers::invalid_syntax($constant).at_value(0))?; + let Some((second, second_owned)) = lines.next() else { + let wildcard = state.finish($constant)?; + return Ok(Some($owned { + values: super::shared::TagValues::One(first_owned), + wildcard, + })); + }; + many(first_owned, second, second_owned, state, lines, mode).map(Some) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), value.values.into_encoded()) + } + } + + impl TryFrom<&str> for $owned { + type Error = DecodeError; + + fn try_from(wire: &str) -> Result { + let value = FieldValue::from_str(wire).map_err(|_invalid| $crate::headers::invalid_syntax($constant))?; + Self::try_from(value) + } + } + + impl TryFrom for $owned { + type Error = DecodeError; + + fn try_from(wire: String) -> Result { + let value = FieldValue::try_from(wire).map_err(|_invalid| $crate::headers::invalid_syntax($constant))?; + Self::try_from(value) + } + } + + impl TryFrom for $owned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + let wildcard = validate_tag_slice(value.as_bytes(), $constant)?; + Ok(Self { + values: super::shared::TagValues::One(value), + wildcard, + }) + } + } + }; +} + +pub(super) use entity_tag_list_header; + +macro_rules! date_header { + ( + $descriptor:ident, + $owned:ident, + $borrowed:ident, + $header_name:literal, + $constant:expr, + $doc:literal, + $view_doc:literal, + $specification:literal, + $examples:literal + ) => { + #[doc = concat!("Defines the `", $header_name, "` header.")] + #[doc = ""] + #[doc = "# Specification"] + #[doc = ""] + #[doc = $specification] + #[derive(Debug)] + pub struct $descriptor { + _private: (), + } + + #[doc = $doc] + #[doc = ""] + #[doc = "# Specification"] + #[doc = ""] + #[doc = $specification] + #[doc = ""] + #[doc = "# Examples"] + #[doc = ""] + #[doc = $examples] + #[derive(Clone, Eq, Hash, PartialEq)] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfModifiedSinceOwned; + /// + /// let value = IfModifiedSinceOwned::try_from("Sat, 29 Oct 1994 19:43:31 GMT")?; + /// let rendered = format!("{value:?}"); + /// assert!(rendered.contains("IfModifiedSinceOwned")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub struct $owned { + value: FieldValue, + date: SystemTime, + } + + #[doc = $view_doc] + #[derive(Clone, Copy, Eq, Hash, PartialEq)] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfModifiedSince; + /// use http_headers::{FieldValue, SingleValueField}; + /// + /// let field = FieldValue::from_static("Sat, 29 Oct 1994 19:43:31 GMT"); + /// let value = ::decode_view(field.as_field_value_ref())?; + /// let rendered = format!("{value:?}"); + /// assert!(rendered.contains("IfModifiedSinceView")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub struct $borrowed<'a> { + value: FieldValueRef<'a>, + date: SystemTime, + } + + impl fmt::Debug for $owned { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct(stringify!($owned)) + .field("date", &self.date) + .finish_non_exhaustive() + } + } + + impl fmt::Debug for $borrowed<'_> { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter + .debug_struct(stringify!($borrowed)) + .field("date", &self.date) + .finish_non_exhaustive() + } + } + + impl $owned { + /// Constructs a canonical IMF-fixdate value. + /// + /// Subsecond precision is omitted on the wire. + /// + /// # Errors + /// + /// Returns an error when the date cannot be represented by `httpdate`. + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfModifiedSinceOwned; + /// + /// let instant = std::time::UNIX_EPOCH + std::time::Duration::from_secs(783_459_811); + /// let value = IfModifiedSinceOwned::new(instant)?; + /// assert_eq!(value.date(), instant); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn new(date: SystemTime) -> Result { + super::shared::format_http_date($constant, date).and_then(Self::try_from) + } + + /// Returns the parsed date. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfModifiedSinceOwned; + /// + /// let instant = std::time::UNIX_EPOCH + std::time::Duration::from_secs(783_459_811); + /// let value = IfModifiedSinceOwned::try_from("Sat, 29 Oct 1994 19:43:31 GMT")?; + /// assert_eq!(value.date(), instant); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn date(&self) -> SystemTime { + self.date + } + + /// Returns the preserved field value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfModifiedSinceOwned; + /// + /// let value = IfModifiedSinceOwned::try_from("Sat, 29 Oct 1994 19:43:31 GMT")?; + /// assert_eq!(value.as_field_value(), "Sat, 29 Oct 1994 19:43:31 GMT"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Consumes the header and returns its field value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfModifiedSinceOwned; + /// + /// let value = IfModifiedSinceOwned::try_from("Sat, 29 Oct 1994 19:43:31 GMT")?; + /// let field = value.into_field_value(); + /// assert_eq!(field, "Sat, 29 Oct 1994 19:43:31 GMT"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } + } + + crate::headers::shared::impl_field_value_conversion!($owned, |value| value.value); + + impl<'a> $borrowed<'a> { + /// Returns the parsed date. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfModifiedSince; + /// use http_headers::{FieldValue, SingleValueField}; + /// + /// let instant = std::time::UNIX_EPOCH + std::time::Duration::from_secs(783_459_811); + /// let field = FieldValue::from_static("Sat, 29 Oct 1994 19:43:31 GMT"); + /// let value = ::decode_view(field.as_field_value_ref())?; + /// assert_eq!(value.date(), instant); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn date(self) -> SystemTime { + self.date + } + + /// Returns the borrowed field value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::IfModifiedSince; + /// use http_headers::{FieldValue, SingleValueField}; + /// + /// let field = FieldValue::from_static("Sat, 29 Oct 1994 19:43:31 GMT"); + /// let value = ::decode_view(field.as_field_value_ref())?; + /// assert_eq!(value.as_field_value(), field.as_field_value_ref()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } + } + + impl SingleValueField for $descriptor { + type View<'a> = $borrowed<'a>; + type Owned = $owned; + + fn name() -> &'static FieldName { + $constant + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + Ok($borrowed { + value, + date: parse_http_date($constant, value)?, + }) + } + + fn decode_owned(value: FieldValue) -> Result { + $owned::try_from(value) + } + + fn decode_view_with(value: FieldValueRef<'_>, mode: crate::DecodeMode) -> Result, DecodeError> { + Ok($borrowed { + value, + date: super::shared::parse_http_date_with($constant, value, mode)?, + }) + } + + fn decode_owned_with(value: FieldValue, mode: crate::DecodeMode) -> Result { + let date = super::shared::parse_http_date_with($constant, value.as_field_value_ref(), mode)?; + Ok($owned { value, date }) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } + } + + impl TryFrom<&str> for $owned { + type Error = DecodeError; + + fn try_from(wire: &str) -> Result { + let value = FieldValue::from_str(wire).map_err(|_invalid| $crate::headers::invalid_syntax($constant))?; + Self::try_from(value) + } + } + + impl TryFrom for $owned { + type Error = DecodeError; + + fn try_from(wire: String) -> Result { + let value = FieldValue::try_from(wire).map_err(|_invalid| $crate::headers::invalid_syntax($constant))?; + Self::try_from(value) + } + } + + impl TryFrom for $owned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + let date = parse_http_date($constant, value.as_field_value_ref())?; + Ok(Self { value, date }) + } + } + }; +} + +pub(super) use date_header; + +pub(super) struct TagIter<'a, I> { + values: I, + current: Option<&'a [u8]>, + position: usize, +} + +impl<'a, I> TagIter<'a, I> +where + I: Iterator>, +{ + pub(super) fn new(values: I) -> Self { + Self { + values, + current: None, + position: 0, + } + } +} + +impl<'a, I> Iterator for TagIter<'a, I> +where + I: Iterator>, +{ + type Item = ConditionalTagView<'a>; + + fn next(&mut self) -> Option { + loop { + if let Some(bytes) = self.current { + match next_list_item(bytes, &mut self.position) { + Some(ListItem::Tag(tag)) => return Some(tag), + Some(ListItem::Wildcard) => continue, + Some(ListItem::Malformed) => return None, + None => self.current = None, + } + } + self.current = Some(self.values.next()?.as_bytes()); + self.position = 0; + } + } +} + +/// One item recognized by [`next_list_item`]. +enum ListItem<'a> { + Tag(ConditionalTagView<'a>), + Wildcard, + Malformed, +} + +/// Scans one comma-delimited list item starting at `position`. +/// +/// Empty items are skipped as required by the list grammar, so `None` means +/// the field line held no further items. +/// Projects one item from a line previously accepted by +/// [`validate_tag_line_with`]. +fn next_list_item<'a>(bytes: &'a [u8], position: &mut usize) -> Option> { + let length = bytes.len(); + let mut index = *position; + + while index < length && matches!(bytes[index], b' ' | b'\t' | b',') { + index += 1; + } + if index == length { + *position = index; + return None; + } + + let start = index; + let item = if bytes[index] == b'*' { + index += 1; + ListItem::Wildcard + } else { + let weak = matches!(bytes[index], b'W' | b'w'); + if weak { + index += 1; + if bytes.get(index) != Some(&b'/') { + *position = length; + return Some(ListItem::Malformed); + } + index += 1; + } + if bytes.get(index) != Some(&b'"') { + *position = length; + return Some(ListItem::Malformed); + } + index += 1; + let Some(closing) = bytes.get(index..)?.iter().position(|byte| *byte == b'"') else { + *position = length; + return Some(ListItem::Malformed); + }; + index += closing; + assert_eq!( + bytes.get(index), + Some(&b'"'), + "position() found a quote in this same immutable slice" + ); + index += 1; + ListItem::Tag(ConditionalTagView { + wire: &bytes[start..index], + }) + }; + + while index < length && matches!(bytes[index], b' ' | b'\t') { + index += 1; + } + if index < length && bytes[index] != b',' { + *position = length; + return Some(ListItem::Malformed); + } + *position = index; + Some(item) +} + +pub(super) fn validate_tag_slice(bytes: &[u8], name: &'static FieldName) -> Result { + validate_tag_line(bytes, TagListState::EMPTY) + .ok_or_else(|| crate::headers::invalid_syntax(name))? + .finish(name) +} + +/// Whether a wildcard and whether any entity tag has been seen so far. +#[derive(Clone, Copy)] +pub(super) struct TagListState(u8); + +impl TagListState { + pub(super) const EMPTY: Self = Self(0); + + /// Set once the field has named the wildcard. + const WILDCARD: u8 = 0b01; + + /// Set once the field has named an entity tag. + const TAGGED: u8 = 0b10; + + /// Returns whether the field is the wildcard form, rejecting an empty list. + pub(super) fn finish(self, name: &'static FieldName) -> Result { + if self.0 == 0 { + Err(DecodeError::new(name, DecodeErrorKind::MissingValue)) + } else { + Ok(self.0 & Self::WILDCARD != 0) + } + } +} + +/// Byte classes used by the entity-tag list scanner. +/// +/// `OWS` marks SP and HTAB, `COMMA` marks the list delimiter, and `QUOTE` +/// marks DQUOTE. A zero class is exactly a byte that may appear inside an +/// opaque tag, so one table load classifies every byte the scanner meets. +/// Independent bits let one byte carry every scanner role without branches. +const TAG_OWS: u8 = 0b0_0001; +const TAG_COMMA: u8 = 0b0_0010; +const TAG_QUOTE: u8 = 0b0_0100; +const TAG_STAR: u8 = 0b0_1000; +const TAG_WEAK: u8 = 0b1_0000; +const TAG_RELAXED_WEAK: u8 = 0b10_0000; + +const TAG_CLASS: [u8; 256] = { + let mut table = [0_u8; 256]; + table[b'\t' as usize] = TAG_OWS; + table[b' ' as usize] = TAG_OWS; + table[b',' as usize] = TAG_COMMA; + table[b'"' as usize] = TAG_QUOTE; + table[b'*' as usize] = TAG_STAR; + table[b'W' as usize] = TAG_WEAK; + table[b'w' as usize] = TAG_RELAXED_WEAK; + table +}; + +/// Validates one field line of the `"*" / 1#entity-tag` grammar in a single +/// pass, accumulating the list state across the field lines of one field. +/// +/// Returns `None` when the line violates the grammar or repeats a wildcard. +#[inline] +fn validate_tag_line(bytes: &[u8], state: TagListState) -> Option { + validate_tag_line_with(bytes, state, crate::DecodeMode::Strict) +} + +#[inline] +pub(super) fn validate_tag_line_with(bytes: &[u8], state: TagListState, mode: crate::DecodeMode) -> Option { + let mut seen = state.0; + let length = bytes.len(); + let mut index = 0; + + 'line: loop { + // The class of the byte that opens the item also selects the item + // shape, so the scan never re-examines a byte it has classified. + let class = loop { + match bytes.get(index) { + None => break 'line, + Some(&byte) => { + let class = TAG_CLASS[byte as usize]; + if class & (TAG_OWS | TAG_COMMA) == 0 { + break class; + } + index += 1; + } + } + }; + + if class & TAG_QUOTE != 0 { + index += 1; + seen |= TagListState::TAGGED; + } else if class & TAG_WEAK != 0 || class & TAG_RELAXED_WEAK != 0 && mode == crate::DecodeMode::Relaxed { + if bytes.get(index + 1..index + 3) != Some(b"/\"".as_slice()) { + return None; + } + index += 3; + seen |= TagListState::TAGGED; + } else if class & TAG_STAR != 0 { + if seen & TagListState::WILDCARD != 0 { + return None; + } + seen |= TagListState::WILDCARD; + index += 1; + continue; + } else { + return None; + } + + while index < length && TAG_CLASS[bytes[index] as usize] & (TAG_OWS | TAG_QUOTE) == 0 { + index += 1; + } + if bytes.get(index) != Some(&b'"') { + return None; + } + index += 1; + + while index < length && TAG_CLASS[bytes[index] as usize] & TAG_OWS != 0 { + index += 1; + } + match bytes.get(index) { + None => break, + Some(&b',') => index += 1, + Some(_) => return None, + } + } + + // A wildcard is only ever legal on its own, which one test at the end of + // the line decides just as well as a test on every list item. + if seen == TagListState::WILDCARD | TagListState::TAGGED { + return None; + } + Some(TagListState(seen)) +} + +/// Validates one field line out of line. +/// +/// A borrowed decode keeps nothing but the field lines alive, so calling the +/// scan rather than inlining it leaves that frame with a shorter prologue than +/// the scan's own register demand would otherwise force on it. +#[inline(never)] +pub(super) fn validate_tag_line_outlined_with(bytes: &[u8], state: TagListState, mode: crate::DecodeMode) -> Option { + validate_tag_line_with(bytes, state, mode) +} + +pub(super) fn parse_http_date(name: &'static FieldName, value: FieldValueRef<'_>) -> Result { + let bytes = value.as_bytes(); + if let Some(seconds) = parse_imf_fixdate(bytes) { + return Ok(UNIX_EPOCH + Duration::from_secs(seconds)); + } + if bytes.first().is_some_and(|byte| matches!(byte, b' ' | b'\t')) || bytes.last().is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + return Err(crate::headers::invalid_syntax(name)); + } + let wire = ascii_str(bytes).ok_or_else(|| crate::headers::invalid_syntax(name))?; + httpdate::parse_http_date(wire).map_err(|_invalid| crate::headers::invalid_syntax(name)) +} + +pub(super) fn parse_http_date_with( + name: &'static FieldName, + value: FieldValueRef<'_>, + mode: crate::DecodeMode, +) -> Result { + if mode == crate::DecodeMode::Strict { + return parse_http_date(name, value); + } + if let Ok(date) = parse_http_date(name, value) { + return Ok(date); + } + let wire = ascii_str(value.as_bytes()) + .ok_or_else(|| crate::headers::invalid_syntax(name))? + .trim_matches([' ', '\t']); + parse_relaxed_http_date(wire).ok_or_else(|| crate::headers::invalid_syntax(name)) +} + +fn parse_relaxed_http_date(wire: &str) -> Option { + if let Ok(date) = httpdate::parse_http_date(wire) { + return Some(date); + } + let mut parts = wire.split(' '); + let weekday = parts.next()?.strip_suffix(',')?; + let day = parse_short_decimal(parts.next()?, 2)?; + let month = parts.next()?; + let year = parse_short_decimal(parts.next()?, 4)?; + let time = parts.next()?; + let zone = parts.next()?; + if parts.next().is_some() + || !matches!(weekday, "Sun" | "Mon" | "Tue" | "Wed" | "Thu" | "Fri" | "Sat") + || !matches!( + month, + "Jan" | "Feb" | "Mar" | "Apr" | "May" | "Jun" | "Jul" | "Aug" | "Sep" | "Oct" | "Nov" | "Dec" + ) + || !matches!(zone, "GMT" | "UTC") + || !(1970..=9999).contains(&year) + { + return None; + } + let mut clock = time.split(':'); + let hour = parse_short_decimal(clock.next()?, 2)?; + let minute = parse_short_decimal(clock.next()?, 2)?; + let second = parse_short_decimal(clock.next()?, 2)?; + if clock.next().is_some() || hour > 23 || minute > 59 || second > 59 { + return None; + } + // The validated fixed-width form needs neither heap storage nor formatting dispatch. + let mut normalized = *b"Sun, 00 Jan 0000 00:00:00 GMT"; + normalized[..3].copy_from_slice(weekday.as_bytes()); + normalized[5..7].copy_from_slice(&date_decimal_pair(day)); + normalized[8..11].copy_from_slice(month.as_bytes()); + normalized[12..14].copy_from_slice(&date_decimal_pair(year / 100)); + normalized[14..16].copy_from_slice(&date_decimal_pair(year % 100)); + normalized[17..19].copy_from_slice(&date_decimal_pair(hour)); + normalized[20..22].copy_from_slice(&date_decimal_pair(minute)); + normalized[23..25].copy_from_slice(&date_decimal_pair(second)); + let normalized = ascii_str(&normalized).expect("validated date names and decimal digits are ASCII"); + httpdate::parse_http_date(normalized).ok() +} + +fn date_decimal_pair(value: u16) -> [u8; 2] { + let value = u8::try_from(value).expect("validated two-digit date components are at most 99"); + [b'0' + value / 10, b'0' + value % 10] +} + +fn parse_short_decimal(value: &str, max_digits: usize) -> Option { + if value.is_empty() || value.len() > max_digits || !value.as_bytes().iter().all(u8::is_ascii_digit) { + return None; + } + value.parse().ok() +} + +/// Days elapsed in a non-leap year before the first of each month. +const MONTH_START_DAY: [u16; 12] = [0, 31, 59, 90, 120, 151, 181, 212, 243, 273, 304, 334]; + +/// Number of days in each month of a non-leap year. +const MONTH_LENGTH: [u8; 12] = [31, 28, 31, 30, 31, 30, 31, 31, 30, 31, 30, 31]; + +/// Parses the preferred `IMF-fixdate` form, returning seconds since the epoch. +/// +/// Returns `None` for anything this fast path does not fully recognize, which +/// leaves the general `httpdate` parser to accept or reject the value. Every +/// value accepted here is accepted by `httpdate` with the same instant. +/// Reads the four bytes at `offset` as one little-endian word. +const fn word_at(bytes: &[u8; 29], offset: usize) -> u32 { + u32::from_le_bytes([bytes[offset], bytes[offset + 1], bytes[offset + 2], bytes[offset + 3]]) +} + +/// Builds the word a four-byte name occupies in an `IMF-fixdate`. +#[cfg_attr(coverage_nightly, coverage(off))] +const fn name_word(name: [u8; 4]) -> u32 { + u32::from_le_bytes(name) +} + +const SUN: u32 = name_word(*b"Sun,"); +const MON: u32 = name_word(*b"Mon,"); +const TUE: u32 = name_word(*b"Tue,"); +const WED: u32 = name_word(*b"Wed,"); +const THU: u32 = name_word(*b"Thu,"); +const FRI: u32 = name_word(*b"Fri,"); +const SAT: u32 = name_word(*b"Sat,"); + +const JAN: u32 = name_word(*b"Jan "); +const FEB: u32 = name_word(*b"Feb "); +const MAR: u32 = name_word(*b"Mar "); +const APR: u32 = name_word(*b"Apr "); +const MAY: u32 = name_word(*b"May "); +const JUN: u32 = name_word(*b"Jun "); +const JUL: u32 = name_word(*b"Jul "); +const AUG: u32 = name_word(*b"Aug "); +const SEP: u32 = name_word(*b"Sep "); +const OCT: u32 = name_word(*b"Oct "); +const NOV: u32 = name_word(*b"Nov "); +const DEC: u32 = name_word(*b"Dec "); + +fn parse_imf_fixdate(bytes: &[u8]) -> Option { + // `Sun, 06 Nov 1994 08:49:37 GMT` + let bytes: &[u8; 29] = bytes.try_into().ok()?; + if bytes[4] != b' ' || bytes[16] != b' ' || bytes[19] != b':' { + return None; + } + if bytes[22] != b':' || word_at(bytes, 25) != u32::from_le_bytes(*b" GMT") { + return None; + } + + // Matching four bytes at a time turns each name table into one load and a + // switch over integers, where a slice pattern would compare byte runs arm + // by arm. The trailing delimiter rides along inside the word, so the + // comma after the weekday and the spaces around the month need no + // separate test. + let weekday = match word_at(bytes, 0) { + SUN => 0, + MON => 1, + TUE => 2, + WED => 3, + THU => 4, + FRI => 5, + SAT => 6, + _ => return None, + }; + if bytes[7] != b' ' { + return None; + } + let month = match word_at(bytes, 8) { + JAN => 1_u16, + FEB => 2, + MAR => 3, + APR => 4, + MAY => 5, + JUN => 6, + JUL => 7, + AUG => 8, + SEP => 9, + OCT => 10, + NOV => 11, + DEC => 12, + _ => return None, + }; + + let day = two_digits(bytes[5], bytes[6])?; + let hour = two_digits(bytes[17], bytes[18])?; + let minute = two_digits(bytes[20], bytes[21])?; + let second = two_digits(bytes[23], bytes[24])?; + let year = u16::from(two_digits(bytes[12], bytes[13])?) * 100 + u16::from(two_digits(bytes[14], bytes[15])?); + + if hour > 23 || minute > 59 || second > 59 || !(1970..=9999).contains(&year) { + return None; + } + + let month_index = usize::from(month - 1); + let leap_year = year % 4 == 0 && (year % 100 != 0 || year % 400 == 0); + let month_length = MONTH_LENGTH[month_index] + u8::from(leap_year && month == 2); + if day == 0 || day > month_length { + return None; + } + + let previous_year = year - 1; + let leap_days = (previous_year - 1968) / 4 - (previous_year - 1900) / 100 + (previous_year - 1600) / 400; + let year_day = MONTH_START_DAY[month_index] + u16::from(day) + u16::from(leap_year && month > 2) - 1; + let days = u64::from(year - 1970) * 365 + u64::from(leap_days) + u64::from(year_day); + + // The epoch fell on a Thursday, which is index 4 in the Sunday-based table. + if (days + 4) % 7 != weekday { + return None; + } + + Some(u64::from(second) + u64::from(minute) * 60 + u64::from(hour) * 3600 + days * 86_400) +} + +fn two_digits(high: u8, low: u8) -> Option { + let high = high.wrapping_sub(b'0'); + let low = low.wrapping_sub(b'0'); + (high < 10 && low < 10).then(|| high * 10 + low) +} + +pub(super) fn format_http_date(name: &'static FieldName, date: SystemTime) -> Result { + let elapsed = date + .duration_since(UNIX_EPOCH) + .map_err(|_before_epoch| DecodeError::new(name, DecodeErrorKind::InvalidNumber))?; + if elapsed.as_secs() > MAX_HTTP_DATE_SECONDS { + return Err(DecodeError::new(name, DecodeErrorKind::InvalidNumber)); + } + Ok(FieldValue::try_from(httpdate::fmt_http_date(date)).expect("an HTTP date contains only valid field-value bytes")) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::time::{Duration, UNIX_EPOCH}; + use std::{iter, str}; + + use super::{ + ListItem, MAX_HTTP_DATE_SECONDS, TagIter, date_decimal_pair, next_list_item, parse_http_date, parse_http_date_with, + parse_imf_fixdate, parse_relaxed_http_date, parse_short_decimal, two_digits, validate_tag_line_with, + }; + use crate::headers::{ + ETagOwned, IfMatch, IfMatchOwned, IfModifiedSince, IfModifiedSinceOwned, IfNoneMatch, IfNoneMatchOwned, IfUnmodifiedSince, + IfUnmodifiedSinceOwned, LastModified, LastModifiedOwned, + }; + use crate::sink::{EncodedValues, FieldSink, InsertError}; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef, SingleValueField}; + + struct Source { + name: &'static FieldName, + values: Vec, + } + + impl FieldSource for Source { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::from_slice(name, &self.values)).flatten() + } + } + + impl FieldSink for Source { + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.name = name; + self.values = values.into_iter().collect(); + Ok(()) + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + if self.name == name { + self.values.extend(values); + } else { + self.name = name; + self.values = values.into_iter().collect(); + } + Ok(()) + } + + fn remove_values(&mut self, name: &'static FieldName) { + if self.name == name { + self.values.clear(); + } + } + } + + #[test] + fn tag_iterator_rejects_an_unterminated_opaque_tag() { + let mut position = 0; + assert!(matches!( + super::next_list_item(b"\"unterminated", &mut position), + Some(super::ListItem::Malformed) + )); + } + + #[test] + fn tag_projection_preserves_quotes_and_advances_past_valid_separators() { + for bytes in [b"\"\"".as_slice(), b"\"strong\"", b"W/\"weak\"", b"w/\"weak\""] { + let mut position = 0; + let Some(ListItem::Tag(tag)) = next_list_item(bytes, &mut position) else { + panic!("valid quoted tag was not retained"); + }; + assert_eq!(tag.as_bytes(), bytes); + assert_eq!(position, bytes.len()); + assert!(next_list_item(bytes, &mut position).is_none()); + } + + let bytes = b" \t\"one\" \t, w/\"two\""; + let mut position = 0; + for expected in [b"\"one\"".as_slice(), b"w/\"two\""] { + let Some(ListItem::Tag(tag)) = next_list_item(bytes, &mut position) else { + panic!("comma-separated tag was not retained"); + }; + assert_eq!(tag.as_bytes(), expected); + } + assert_eq!(position, bytes.len()); + assert!(next_list_item(bytes, &mut position).is_none()); + } + + #[test] + fn tag_projection_rejects_missing_quotes_weak_slashes_and_bad_separators() { + for bytes in [ + b"bare".as_slice(), + b"Wbad", + b"W/'wrong'", + b"\"open", + b"W/\"open", + b"\"tag'x", + b"\"tag\" suffix", + b"\"tag\";\"next\"", + ] { + let mut position = 0; + assert!( + matches!(next_list_item(bytes, &mut position), Some(ListItem::Malformed)), + "{bytes:?}" + ); + assert_eq!(position, bytes.len()); + assert!(next_list_item(bytes, &mut position).is_none()); + } + } + + #[test] + fn date_headers_reject_malformed_duplicate_and_out_of_range_values() { + let error = LastModifiedOwned::try_from("not a date").expect_err("malformed date"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + + let source = Source { + name: &FieldName::IfModifiedSince, + values: vec![ + FieldValue::from_static("Sun, 06 Nov 1994 08:49:37 GMT"), + FieldValue::from_static("Mon, 07 Nov 1994 08:49:37 GMT"), + ], + }; + let error = IfModifiedSince::view(&source).expect_err("singleton duplication must fail"); + assert_eq!(error.kind(), DecodeErrorKind::UnexpectedMultipleValues); + + let too_late = UNIX_EPOCH + .checked_add(Duration::from_secs(MAX_HTTP_DATE_SECONDS + 1)) + .expect("platform represents test date"); + let error = LastModifiedOwned::new(too_late).expect_err("year 10000 must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber); + let error = IfModifiedSinceOwned::new(too_late).expect_err("year 10000 must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber); + let error = IfUnmodifiedSinceOwned::new(too_late).expect_err("year 10000 must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber); + let error = LastModifiedOwned::new(UNIX_EPOCH - Duration::from_secs(1)).expect_err("pre-epoch date must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber); + } + + #[test] + fn entity_tag_headers_cover_construction_decoding_and_storage() { + let strong = ETagOwned::try_from("\"strong\"").expect("strong tag"); + let weak = ETagOwned::try_from("W/\"weak\"").expect("weak tag"); + + let missing = IfMatchOwned::from_tags(Vec::::new()).expect_err("a tag list must not be empty"); + assert_eq!(missing.kind(), DecodeErrorKind::MissingValue); + + let one = IfMatchOwned::from_tags([strong.clone()]).expect("one tag"); + assert!(!one.is_wildcard()); + assert_eq!(one.tags().next().expect("tag").opaque_tag(), b"strong"); + assert!(format!("{one:?}").contains("value_count")); + + let many = IfNoneMatchOwned::from_tags([strong, weak, ETagOwned::try_from("\"last\"").expect("tag")]).expect("several tags"); + assert!(format!("{many:?}").contains("value_count")); + let tags: Vec<_> = many.tags().collect(); + assert_eq!(tags.len(), 3); + assert_eq!(tags[0].as_bytes(), b"\"strong\""); + assert!(tags[1].is_weak()); + assert_eq!(tags[1].opaque_tag(), b"weak"); + + let wildcard = IfMatchOwned::wildcard(); + assert!(wildcard.is_wildcard()); + assert_eq!(wildcard.tags().count(), 0); + let mut wildcard_sink = Source { + name: &FieldName::Accept, + values: Vec::new(), + }; + IfMatch::insert(&mut wildcard_sink, wildcard).expect("insert wildcard"); + assert_eq!(wildcard_sink.values, [FieldValue::from_static("*")]); + + let from_string = IfMatchOwned::try_from(String::from("\"owned\"")).expect("owned tag string"); + assert_eq!(from_string.tags().next().expect("tag").opaque_tag(), b"owned"); + let from_value = IfMatchOwned::try_from(FieldValue::from_static("\"field\"")).expect("tag field value"); + assert_eq!(from_value.tags().next().expect("tag").opaque_tag(), b"field"); + + let mut sink = Source { + name: &FieldName::Accept, + values: Vec::new(), + }; + IfNoneMatch::insert(&mut sink, many).expect("insert tag fields"); + assert_eq!(sink.name, &FieldName::IfNoneMatch); + assert_eq!(sink.values.len(), 3); + + let view = IfNoneMatch::view(&sink).expect("decode view").expect("present"); + assert!(!view.is_wildcard()); + assert_eq!(view.tags().count(), 3); + assert!(format!("{view:?}").contains("value_count")); + + let owned = IfNoneMatch::owned(&sink).expect("decode owned").expect("present"); + assert_eq!(owned.tags().count(), 3); + let mut round_trip = Source { + name: &FieldName::Accept, + values: Vec::new(), + }; + IfNoneMatch::insert(&mut round_trip, owned).expect("reinsert owned tags"); + assert_eq!(round_trip.values, sink.values, "wire field lines are retained"); + + let absent = Source { + name: &FieldName::Accept, + values: Vec::new(), + }; + assert!(IfMatch::view(&absent).expect("absent header").is_none()); + assert!(IfMatch::owned(&absent).expect("absent header").is_none()); + } + + #[test] + fn entity_tag_headers_report_line_and_mode_errors() { + let malformed_first = Source { + name: &FieldName::IfMatch, + values: vec![FieldValue::from_static("not-a-tag")], + }; + assert_eq!( + IfMatch::view(&malformed_first).expect_err("malformed first line").value_index(), + Some(0) + ); + assert_eq!( + IfMatch::owned(&malformed_first).expect_err("malformed first line").value_index(), + Some(0) + ); + + let malformed_later = Source { + name: &FieldName::IfMatch, + values: vec![ + FieldValue::from_static("\"one\""), + FieldValue::from_static("\"two\""), + FieldValue::from_static("bad"), + ], + }; + assert_eq!( + IfMatch::view(&malformed_later).expect_err("malformed third line").value_index(), + Some(2) + ); + assert_eq!( + IfMatch::owned(&malformed_later).expect_err("malformed third line").value_index(), + Some(2) + ); + + for (wire, kind) in [ + ("", DecodeErrorKind::MissingValue), + ("*, \"tag\"", DecodeErrorKind::InvalidSyntax), + ("*, *", DecodeErrorKind::InvalidSyntax), + ("Wtag", DecodeErrorKind::InvalidSyntax), + ("\"unterminated", DecodeErrorKind::InvalidSyntax), + ] { + let error = IfMatchOwned::try_from(wire).expect_err("invalid tag list"); + assert_eq!(error.kind(), kind, "{wire:?}"); + } + let error = IfMatchOwned::try_from(String::from("bad\nvalue")).expect_err("invalid field-value bytes"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + + let relaxed = Source { + name: &FieldName::IfNoneMatch, + values: vec![FieldValue::from_static("w/\"weak\"")], + }; + IfNoneMatch::view(&relaxed).expect_err("strict lowercase weak tag"); + let view = IfNoneMatch::view_with(&relaxed, DecodeMode::Relaxed) + .expect("relaxed weak tag") + .expect("present"); + assert!(view.tags().next().expect("tag").is_weak()); + assert!( + IfNoneMatch::owned_with(&relaxed, DecodeMode::Relaxed) + .expect("relaxed weak tag") + .expect("present") + .tags() + .next() + .expect("tag") + .is_weak() + ); + + let tagged = validate_tag_line_with(b"\"tag\"", super::TagListState::EMPTY, DecodeMode::Strict).expect("valid tag"); + assert!(validate_tag_line_with(b"*", tagged, DecodeMode::Strict).is_none()); + assert!(validate_tag_line_with(b"\"one\" \t, \"two\"", super::TagListState::EMPTY, DecodeMode::Strict,).is_some()); + assert!(validate_tag_line_with(b"\"tag\" suffix", super::TagListState::EMPTY, DecodeMode::Strict,).is_none()); + + assert_eq!( + IfMatchOwned::try_from("bad\nvalue") + .expect_err("invalid borrowed field bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + IfNoneMatchOwned::try_from("bad\nvalue") + .expect_err("invalid borrowed field bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn both_entity_tag_header_variants_exercise_generated_apis() { + let source = Source { + name: &FieldName::IfMatch, + values: vec![FieldValue::from_static("\"one\", W/\"two\"")], + }; + let view = IfMatch::view(&source).expect("If-Match view").expect("present"); + assert!(!view.is_wildcard()); + assert_eq!(view.tags().count(), 2); + assert!(format!("{view:?}").contains("value_count")); + let owned = IfMatch::owned(&source).expect("If-Match owned").expect("present"); + assert_eq!(owned.tags().count(), 2); + + let wildcard = IfNoneMatchOwned::wildcard(); + assert!(wildcard.is_wildcard()); + assert_eq!(wildcard.tags().count(), 0); + let mut sink = Source { + name: &FieldName::Accept, + values: Vec::new(), + }; + IfNoneMatch::insert(&mut sink, wildcard).expect("insert If-None-Match wildcard"); + let view = IfNoneMatch::view(&sink).expect("wildcard view").expect("present"); + assert!(view.is_wildcard()); + assert_eq!(view.tags().count(), 0); + + assert_eq!( + IfNoneMatchOwned::try_from("\"borrowed\"").expect("borrowed string").tags().count(), + 1 + ); + assert_eq!( + IfNoneMatchOwned::try_from(String::from("\"owned\"")) + .expect("owned string") + .tags() + .count(), + 1 + ); + assert_eq!( + IfNoneMatchOwned::try_from(FieldValue::from_static("\"field\"")) + .expect("field value") + .tags() + .count(), + 1 + ); + } + + #[test] + fn repeated_tag_line_errors_report_exact_indices_for_both_headers() { + let if_match = Source { + name: &FieldName::IfMatch, + values: vec![FieldValue::from_static("\"valid\""), FieldValue::from_static("invalid")], + }; + assert_eq!( + IfMatch::owned(&if_match).expect_err("invalid second If-Match line").value_index(), + Some(1) + ); + + let if_none_match = Source { + name: &FieldName::IfNoneMatch, + values: vec![ + FieldValue::from_static("\"first\""), + FieldValue::from_static("\"second\""), + FieldValue::from_static("invalid"), + ], + }; + assert_eq!( + IfNoneMatch::owned(&if_none_match) + .expect_err("invalid third If-None-Match line") + .value_index(), + Some(2) + ); + assert_eq!( + IfNoneMatch::view(&if_none_match) + .expect_err("invalid third If-None-Match line") + .value_index(), + Some(2) + ); + } + + #[test] + fn tag_iterator_stops_safely_on_malformed_preserved_input() { + let values = [ + FieldValue::from_static("*, \"one\""), + FieldValue::from_static("w/\"two\", \"three\""), + FieldValue::from_static("Wbad"), + FieldValue::from_static("\"unreached\""), + ]; + let tags: Vec<_> = TagIter::new(values.iter().map(FieldValue::as_field_value_ref)).collect(); + assert_eq!(tags.len(), 3); + assert_eq!(tags[0].opaque_tag(), b"one"); + assert_eq!(tags[1].opaque_tag(), b"two"); + assert_eq!(tags[2].opaque_tag(), b"three"); + + for malformed in ["bare", "\"open", "\"tag\" suffix"] { + let value = FieldValue::from_str(malformed).expect("valid field bytes"); + assert_eq!(TagIter::new(iter::once(value.as_field_value_ref())).count(), 0); + } + } + + #[test] + fn date_headers_cover_constructors_accessors_and_decode_modes() { + let instant = httpdate::parse_http_date("Tue, 08 Nov 1994 08:49:37 GMT").expect("reference date"); + let last = LastModifiedOwned::new(instant).expect("representable date"); + assert_eq!(last.date(), instant); + assert_eq!(last.as_field_value().as_bytes(), b"Tue, 08 Nov 1994 08:49:37 GMT"); + assert!(format!("{last:?}").contains("date")); + + let from_string = IfModifiedSinceOwned::try_from(String::from("Tue, 08 Nov 1994 08:49:37 GMT")).expect("date string"); + assert_eq!(from_string.date(), instant); + let from_value = IfUnmodifiedSinceOwned::try_from(FieldValue::from_static("Tue, 08 Nov 1994 08:49:37 GMT")).expect("date field"); + assert_eq!(from_value.date(), instant); + + let view = ::decode_view(last.as_field_value().as_field_value_ref()).expect("date view"); + assert_eq!(view.date(), instant); + assert_eq!(view.as_field_value(), last.as_field_value().as_field_value_ref()); + assert!(format!("{view:?}").contains("date")); + + let owned = ::decode_owned(last.clone().into_field_value()).expect("owned date"); + assert_eq!(::as_field_value(&owned), last.as_field_value()); + assert_eq!( + ::into_field_value(owned), + last.clone().into_field_value() + ); + assert_eq!(::name(), &FieldName::LastModified); + + let relaxed = FieldValue::from_static(" Tue, 8 Nov 1994 8:49:37 UTC "); + ::decode_view(relaxed.as_field_value_ref()).expect_err("strict noncanonical date"); + let view = ::decode_view_with(relaxed.as_field_value_ref(), DecodeMode::Relaxed) + .expect("relaxed date view"); + assert_eq!(view.date(), instant); + let owned = ::decode_owned_with(relaxed, DecodeMode::Relaxed).expect("relaxed owned date"); + assert_eq!(owned.date(), instant); + let strict_owned = ::decode_owned_with( + FieldValue::from_static("Tue, 08 Nov 1994 08:49:37 GMT"), + DecodeMode::Strict, + ) + .expect("strict owned date"); + assert_eq!(strict_owned.date(), instant); + } + + #[test] + fn every_date_header_variant_exercises_generated_apis() { + macro_rules! exercise { + ($descriptor:ty, $owned:ty, $name:expr) => {{ + let wire = "Tue, 08 Nov 1994 08:49:37 GMT"; + let instant = httpdate::parse_http_date(wire).expect("reference date"); + let constructed = <$owned>::new(instant).expect("constructed date"); + assert_eq!(constructed.date(), instant); + assert_eq!(constructed.as_field_value(), wire); + assert!(format!("{constructed:?}").contains("date")); + assert_eq!(constructed.clone().into_field_value(), wire); + + assert_eq!(<$owned>::try_from(wire).expect("borrowed date").date(), instant); + assert_eq!(<$owned>::try_from(String::from(wire)).expect("owned date").date(), instant); + assert_eq!( + <$owned>::try_from(FieldValue::from_static(wire)).expect("field date").date(), + instant + ); + assert_eq!( + <$owned>::try_from("bad\nvalue") + .expect_err("invalid borrowed field bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + <$owned>::try_from(String::from("bad\nvalue")) + .expect_err("invalid owned field bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + assert_eq!(<$descriptor as SingleValueField>::name(), $name); + let field = FieldValue::from_static(wire); + let view = <$descriptor as SingleValueField>::decode_view(field.as_field_value_ref()).expect("date view"); + assert_eq!(view.date(), instant); + assert_eq!(view.as_field_value(), wire); + assert!(format!("{view:?}").contains("date")); + assert_eq!( + <$descriptor as SingleValueField>::decode_view_with(field.as_field_value_ref(), DecodeMode::Relaxed,) + .expect("explicit-mode date view") + .date(), + instant + ); + + let owned = <$descriptor as SingleValueField>::decode_owned(field.clone()).expect("decoded owned date"); + assert_eq!(<$descriptor as SingleValueField>::as_field_value(&owned), &field); + assert_eq!(<$descriptor as SingleValueField>::into_field_value(owned), field); + assert_eq!( + <$descriptor as SingleValueField>::decode_owned_with(FieldValue::from_static(wire), DecodeMode::Relaxed,) + .expect("explicit-mode owned date") + .date(), + instant + ); + assert_eq!( + <$descriptor as SingleValueField>::decode_owned_with(FieldValue::from_static("not a date"), DecodeMode::Relaxed,) + .expect_err("invalid explicit-mode date") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + }}; + } + + exercise!(LastModified, LastModifiedOwned, &FieldName::LastModified); + exercise!(IfModifiedSince, IfModifiedSinceOwned, &FieldName::IfModifiedSince); + exercise!(IfUnmodifiedSince, IfUnmodifiedSinceOwned, &FieldName::IfUnmodifiedSince); + } + + #[test] + fn obsolete_dates_match_the_reference_for_every_byte_substitution() { + for wire in [b"Sun Nov 6 08:49:37 1994".as_slice(), b"Sunday, 06-Nov-94 08:49:37 GMT"] { + for offset in 0..wire.len() { + for byte in crate::test_support::substitution_bytes(wire[offset], offset, wire.len()) { + let mut input = wire.to_vec(); + input[offset] = byte; + let expected = str::from_utf8(&input) + .ok() + .filter(|text| !text.starts_with([' ', '\t']) && !text.ends_with([' ', '\t'])) + .and_then(|text| httpdate::parse_http_date(text).ok()) + .ok_or_else(|| crate::headers::invalid_syntax(&FieldName::IfModifiedSince)); + assert_eq!( + parse_http_date(&FieldName::IfModifiedSince, FieldValueRef::new(&input)), + expected, + "{input:?}" + ); + } + } + } + } + + #[test] + fn date_normalization_matches_canonical_reference_dates() { + for value in 0..100_u16 { + assert_eq!(date_decimal_pair(value).as_slice(), format!("{value:02}").as_bytes()); + } + for (short, canonical) in [ + ("Thu, 1 Jan 1970 0:0:0 UTC", "Thu, 01 Jan 1970 00:00:00 GMT"), + ("Tue, 8 Nov 1994 8:9:7 UTC", "Tue, 08 Nov 1994 08:09:07 GMT"), + ("Tue, 29 Feb 2000 9:8:7 UTC", "Tue, 29 Feb 2000 09:08:07 GMT"), + ("Thu, 29 Feb 2024 0:0:0 UTC", "Thu, 29 Feb 2024 00:00:00 GMT"), + ("Fri, 31 Dec 9999 23:59:59 UTC", "Fri, 31 Dec 9999 23:59:59 GMT"), + ("Tue, 8 Nov 1994 8:9:7 GMT", "Tue, 08 Nov 1994 08:09:07 GMT"), + ("Tue, 0 Nov 1994 8:9:7 UTC", "Tue, 00 Nov 1994 08:09:07 GMT"), + ("Mon, 29 Feb 2100 0:0:0 UTC", "Mon, 29 Feb 2100 00:00:00 GMT"), + ("Fri, 31 Apr 2020 0:0:0 UTC", "Fri, 31 Apr 2020 00:00:00 GMT"), + ("Tue, 8 Nov 1994 24:0:0 UTC", "Tue, 08 Nov 1994 24:00:00 GMT"), + ] { + assert_eq!(parse_relaxed_http_date(short), httpdate::parse_http_date(canonical).ok(), "{short}"); + } + } + + #[test] + fn imf_date_checks_every_gmt_tail_byte() { + let wire = *b"Sun, 06 Nov 1994 08:49:37 GMT"; + let seconds = httpdate::parse_http_date("Sun, 06 Nov 1994 08:49:37 GMT") + .unwrap() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_secs(); + for (offset, &original) in wire.iter().enumerate().skip(25) { + for replacement in u8::MIN..=u8::MAX { + let mut modified = wire; + modified[offset] = replacement; + assert_eq!( + parse_imf_fixdate(&modified), + (replacement == original).then_some(seconds), + "offset {offset}, replacement {replacement}" + ); + } + } + } + + #[test] + fn private_date_parsers_cover_fast_and_fallback_branches() { + let dates = [ + "Sat, 01 Jan 2000 00:00:00 GMT", + "Sun, 06 Nov 1994 08:49:37 GMT", + "Mon, 07 Feb 2000 00:00:00 GMT", + "Tue, 08 Mar 2005 00:00:00 GMT", + "Wed, 09 Apr 2008 00:00:00 GMT", + "Thu, 10 May 2012 00:00:00 GMT", + "Fri, 11 Jun 2021 00:00:00 GMT", + "Sat, 12 Jul 2014 00:00:00 GMT", + "Sun, 13 Aug 2017 00:00:00 GMT", + "Fri, 14 Sep 2018 00:00:00 GMT", + "Tue, 15 Oct 2019 00:00:00 GMT", + "Mon, 16 Nov 2020 00:00:00 GMT", + "Fri, 17 Dec 2021 00:00:00 GMT", + "Thu, 29 Feb 2024 00:00:00 GMT", + ]; + for wire in dates { + let value = FieldValue::from_str(wire).expect("date field"); + assert_eq!( + parse_http_date(&FieldName::LastModified, value.as_field_value_ref()).expect("valid date"), + httpdate::parse_http_date(wire).expect("reference date") + ); + } + + for wire in [ + "short", + " Sun, 06 Nov 1994 08:49:37 GMT", + "Sun, 06 Nov 1994 08:49:37 GMT ", + "Sun. 06 Nov 1994 08:49:37 GMT", + "Sun, 06 Nov 1994 08-49:37 GMT", + "Sun, 06 Nov 1994 08:49-37 GMT", + "Sun, 06 Nov 1994 08:49:37 UTC", + "Sun, 06XNov 1994 08:49:37 GMT", + "Bad, 06 Nov 1994 08:49:37 GMT", + "Sun, 06 Xxx 1994 08:49:37 GMT", + "Sun, 00 Nov 1994 08:49:37 GMT", + "Sun, 31 Nov 1994 08:49:37 GMT", + "Sun, 06 Nov 1969 08:49:37 GMT", + "Sun, 06 Nov 1994 24:49:37 GMT", + "Sun, 06 Nov 1994 08:60:37 GMT", + "Sun, 06 Nov 1994 08:49:60 GMT", + "Mon, 06 Nov 1994 08:49:37 GMT", + "Sun, x6 Nov 1994 08:49:37 GMT", + ] { + assert!(parse_imf_fixdate(wire.as_bytes()).is_none(), "{wire}"); + } + + assert!(parse_relaxed_http_date("Tue, 8 Nov 1994 8:49:37 UTC").is_some()); + assert!(parse_relaxed_http_date("Tue, 08 Nov 1994 08:49:37 GMT").is_some()); + for wire in [ + "Tue 8 Nov 1994 8:49:37 UTC", + "Bad, 8 Nov 1994 8:49:37 UTC", + "Tue, 8 Bad 1994 8:49:37 UTC", + "Tue, 8 Nov 1969 8:49:37 UTC", + "Tue, 8 Nov 1994 8:49:37 PST", + "Tue, 8 Nov 1994 24:49:37 UTC", + "Tue, 8 Nov 1994 8:60:37 UTC", + "Tue, 8 Nov 1994 8:49:60 UTC", + "Tue, 8 Nov 1994 8:49:37 UTC extra", + ] { + assert!(parse_relaxed_http_date(wire).is_none(), "{wire}"); + } + assert_eq!(parse_short_decimal("42", 2), Some(42)); + assert_eq!(parse_short_decimal("", 2), None); + assert_eq!(parse_short_decimal("123", 2), None); + assert_eq!(parse_short_decimal("x", 2), None); + assert_eq!(two_digits(b'4', b'2'), Some(42)); + assert_eq!(two_digits(b'x', b'2'), None); + + let canonical = FieldValue::from_static("Tue, 08 Nov 1994 08:49:37 GMT"); + assert_eq!( + parse_http_date_with(&FieldName::LastModified, canonical.as_field_value_ref(), DecodeMode::Relaxed,) + .expect("canonical relaxed date"), + httpdate::parse_http_date("Tue, 08 Nov 1994 08:49:37 GMT").expect("reference date") + ); + assert!(parse_imf_fixdate(b"Sun, 06 Nov 1994 08:49:37 GMX").is_none()); + + let invalid_utf8 = FieldValue::from_bytes([0xff]).expect("field byte is permitted"); + parse_http_date_with(&FieldName::LastModified, invalid_utf8.as_field_value_ref(), DecodeMode::Relaxed) + .expect_err("invalid UTF-8 date"); + + let mut sink = Source { + name: &FieldName::LastModified, + values: vec![FieldValue::from_static("date")], + }; + sink.remove_values(&FieldName::Accept); + assert_eq!(sink.values.len(), 1); + sink.remove_values(&FieldName::LastModified); + assert!(sink.values.is_empty()); + } +} diff --git a/crates/http_headers/src/headers/content_length.rs b/crates/http_headers/src/headers/content_length.rs new file mode 100644 index 000000000..be11fe1da --- /dev/null +++ b/crates/http_headers/src/headers/content_length.rs @@ -0,0 +1,232 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Numeric `Content-Length` parsing and encoding. + +use std::fmt; +use std::str::FromStr; + +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue}; + +/// Defines the `Content-Length` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 8.6](https://www.rfc-editor.org/rfc/rfc9110#section-8.6). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{ContentLength, ContentLengthOwned}; +/// +/// let mut map = HeaderMap::new(); +/// ContentLength::insert(&mut map, ContentLengthOwned::new(42))?; +/// assert_eq!( +/// ContentLength::view(&map)?.map(ContentLengthOwned::get), +/// Some(42) +/// ); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct ContentLength { + _private: (), +} + +/// Owned value for the `Content-Length` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 8.6]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::ContentLengthOwned::new(42); +/// assert_eq!(value.get(), 42); +/// ``` +/// +/// `Content-Length: 0` describes an empty content body, while +/// `Content-Length: 1024` describes 1,024 octets. +/// +/// [RFC 9110 section 8.6]: https://www.rfc-editor.org/rfc/rfc9110#section-8.6 +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct ContentLengthOwned(u64); + +impl ContentLengthOwned { + /// Creates a content length. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ContentLengthOwned::new(42); + /// assert_eq!(value.get(), 42); + /// ``` + pub const fn new(length: u64) -> Self { + Self(length) + } + + /// Returns the length in octets. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ContentLengthOwned::new(42); + /// assert_eq!(value.get(), 42); + /// ``` + pub const fn get(self) -> u64 { + self.0 + } +} + +impl fmt::Display for ContentLengthOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.fmt(f) + } +} + +impl FromStr for ContentLengthOwned { + type Err = DecodeError; + + fn from_str(value: &str) -> Result { + parse_decimal_ows(value.as_bytes()) + .map(Self) + .ok_or_else(|| DecodeError::new(&FieldName::ContentLength, DecodeErrorKind::InvalidNumber)) + } +} + +impl Field for ContentLength { + type View<'a> = ContentLengthOwned; + type Owned = ContentLengthOwned; + + fn name() -> &'static FieldName { + &FieldName::ContentLength + } + + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_list_item_limit(b',', false)?; + let mut repeated = lines.repeated(); + let first = repeated.next().expect("FieldLines always contains at least one value"); + let second = repeated.next(); + if second.is_none() && !first.as_bytes().contains(&b',') { + return parse_decimal_ows(first.as_bytes()) + .map(ContentLengthOwned) + .map(Some) + .ok_or_else(|| DecodeError::new(&FieldName::ContentLength, DecodeErrorKind::InvalidNumber).at_value(0)); + } + + let mut parsed = None; + parse_content_length_line(&mut parsed, first.as_bytes(), 0)?; + if let Some(second) = second { + parse_content_length_line(&mut parsed, second.as_bytes(), 1)?; + } + for (offset, value) in repeated.enumerate() { + parse_content_length_line(&mut parsed, value.as_bytes(), offset + 2)?; + } + Ok(parsed.map(ContentLengthOwned)) + } + + fn owned_with(source: &S, mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + Self::view_with(source, mode) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::single(FieldValue::from(value.0))) + } +} + +fn parse_content_length_line(parsed: &mut Option, bytes: &[u8], value_index: usize) -> Result<(), DecodeError> { + for item in bytes.split(|byte| *byte == b',') { + let number = parse_decimal_ows(item) + .ok_or_else(|| DecodeError::new(&FieldName::ContentLength, DecodeErrorKind::InvalidNumber).at_value(value_index))?; + if parsed.is_some_and(|previous| previous != number) { + return Err(DecodeError::new(&FieldName::ContentLength, DecodeErrorKind::InvalidSyntax).at_value(value_index)); + } + *parsed = Some(number); + } + Ok(()) +} + +fn parse_decimal_ows(bytes: &[u8]) -> Option { + let mut index = 0; + while bytes.get(index).is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + index += 1; + } + let digit_start = index; + let mut value = 0_u64; + while let Some(&byte) = bytes.get(index) { + if byte.is_ascii_digit() { + value = value.checked_mul(10)?.checked_add(u64::from(byte - b'0'))?; + index += 1; + } else { + break; + } + } + if index == digit_start { + return None; + } + while bytes.get(index).is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + index += 1; + } + (index == bytes.len()).then_some(value) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::str; + + use super::parse_decimal_ows; + + #[test] + fn decimal_parser_matches_independent_whitespace_and_digit_oracle() { + for wire in [ + b" \t00123\t ".as_slice(), + b"18446744073709551615", + b"18446744073709551616", + b"000000000000000000000000000001", + ] { + for offset in 0..wire.len() { + for replacement in crate::test_support::substitution_bytes(wire[offset], offset, wire.len()) { + let mut candidate = wire.to_vec(); + candidate[offset] = replacement; + let expected = str::from_utf8(&candidate) + .ok() + .map(|text| text.trim_matches([' ', '\t'])) + .filter(|text| !text.is_empty() && text.as_bytes().iter().all(u8::is_ascii_digit)) + .and_then(|text| text.parse::().ok()); + assert_eq!(parse_decimal_ows(&candidate), expected, "{candidate:?}"); + } + } + } + } + + #[test] + fn decimal_parser_rejects_every_invalid_shape() { + assert_eq!(parse_decimal_ows(b"0"), Some(0)); + assert_eq!(parse_decimal_ows(b"\t18446744073709551615 "), Some(u64::MAX)); + for value in [b"".as_slice(), b" \t", b"+1", b"-1", b"1 2", b"18446744073709551616"] { + assert_eq!(parse_decimal_ows(value), None, "{value:?}"); + } + } +} diff --git a/crates/http_headers/src/headers/content_type.rs b/crates/http_headers/src/headers/content_type.rs new file mode 100644 index 000000000..a2fffef25 --- /dev/null +++ b/crates/http_headers/src/headers/content_type.rs @@ -0,0 +1,1247 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Structured `Content-Type` parsing with lazy parameters. + +use std::hash::{Hash, Hasher}; +use std::ops::Range; +use std::str; + +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Keeps the common one- or two-parameter media type inline; increasing it enlarges every +/// parsed metadata value while reducing spills for unusually parameter-heavy values. +const INLINE_PARAMETER_CAPACITY: usize = 2; + +/// Defines the `Content-Type` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 8.3](https://www.rfc-editor.org/rfc/rfc9110#section-8.3). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{ContentType, ContentTypeOwned}; +/// +/// let mut map = HeaderMap::new(); +/// ContentType::insert(&mut map, ContentTypeOwned::try_from("text/plain")?)?; +/// assert!(ContentType::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct ContentType { + _private: (), +} + +impl ContentType { + /// Creates the canonical `application/json` response value. + #[must_use] + pub fn json() -> ContentTypeOwned { + ContentTypeOwned::json() + } +} + +/// Owned value for the `Content-Type` header. +/// +/// Equality and hashing use the original field-value bytes, ignoring +/// sensitivity and the parsed metadata cache. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 8.3]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::ContentTypeOwned::try_from("text/html; charset=utf-8")?; +/// assert_eq!(value.parameter("charset")?, Some(b"utf-8".as_slice())); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Content-Type: application/json` has no parameters. +/// `Content-Type: text/html; charset=UTF-8` carries a token parameter, while +/// `Content-Type: multipart/form-data; boundary="example boundary"` carries a +/// quoted parameter. +/// +/// [RFC 9110 section 8.3]: https://www.rfc-editor.org/rfc/rfc9110#section-8.3 +#[derive(Clone, Debug)] +pub struct ContentTypeOwned { + value: FieldValue, + metadata: ContentTypeMetadata, +} + +impl PartialEq for ContentTypeOwned { + fn eq(&self, other: &Self) -> bool { + self.value == other.value + } +} + +impl Eq for ContentTypeOwned {} + +impl Hash for ContentTypeOwned { + fn hash(&self, state: &mut H) { + self.value.hash(state); + } +} + +/// Borrowed value for the `Content-Type` header. +#[derive(Clone, Debug)] +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::header::CONTENT_TYPE; +/// use http::{HeaderMap, HeaderValue}; +/// use http_headers::Field; +/// use http_headers::headers::{ContentType, ContentTypeView}; +/// +/// let mut map = HeaderMap::new(); +/// map.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); +/// let view: ContentTypeView<'_> = ContentType::view(&map)?.expect("content type"); +/// assert_eq!(view.subtype()?, "json"); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +pub struct ContentTypeView<'a> { + value: FieldValueRef<'a>, + metadata: ContentTypeMetadata, +} + +/// One borrowed media type parameter. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ``` +/// use http_headers::headers::{ContentTypeOwned, MediaTypeParameterView}; +/// +/// let value = ContentTypeOwned::try_from("text/html; charset=utf-8")?; +/// let parameter: MediaTypeParameterView<'_> = value.parameters().next().expect("charset")?; +/// assert_eq!(parameter.name(), "charset"); +/// assert_eq!(parameter.value(), b"utf-8"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct MediaTypeParameterView<'a> { + name: &'a str, + value: &'a [u8], +} + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +struct ContentTypeHead { + type_start: u32, + type_end: u32, + subtype_start: u32, + subtype_end: u32, + parameter_start: u32, + parameter_count: u32, + inline_parameters: [InlineParameter; INLINE_PARAMETER_CAPACITY], + inline_count: u8, +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +enum ContentTypeMetadata { + ApplicationJsonUtf8, + TextHtmlUtf8, + Common { type_end: u8, subtype_end: u8 }, + Parsed(Box), +} + +impl ContentTypeMetadata { + fn head(&self) -> ContentTypeHead { + match self { + // Byte offsets in "application/json; charset=utf-8". + Self::ApplicationJsonUtf8 => common_parameter_head(11, 16, 18, 25, 26, 31), + // Byte offsets in "text/html; charset=utf-8". + Self::TextHtmlUtf8 => common_parameter_head(4, 9, 11, 18, 19, 24), + Self::Common { type_end, subtype_end } => ContentTypeHead { + type_start: 0, + type_end: u32::from(*type_end), + subtype_start: u32::from(*type_end) + 1, + subtype_end: u32::from(*subtype_end), + parameter_start: u32::from(*subtype_end), + parameter_count: 0, + inline_parameters: [InlineParameter::EMPTY; INLINE_PARAMETER_CAPACITY], + inline_count: 0, + }, + Self::Parsed(head) => **head, + } + } +} + +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +struct ParameterRange { + name: Range, + value: Range, +} + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +struct InlineParameter { + name_start: u16, + name_end: u16, + value_start: u16, + value_end: u16, +} + +impl InlineParameter { + const EMPTY: Self = Self { + name_start: 0, + name_end: 0, + value_start: 0, + value_end: 0, + }; + + fn from_range(range: &ParameterRange) -> Option { + Some(Self { + name_start: u16::try_from(range.name.start).ok()?, + name_end: u16::try_from(range.name.end).ok()?, + value_start: u16::try_from(range.value.start).ok()?, + value_end: u16::try_from(range.value.end).ok()?, + }) + } + + fn name_range(self) -> Range { + usize::from(self.name_start)..usize::from(self.name_end) + } + + fn value_range(self) -> Range { + usize::from(self.value_start)..usize::from(self.value_end) + } +} + +/// Iterator over borrowed media type parameters. +#[derive(Debug)] +/// # Examples +/// +/// ``` +/// use http_headers::headers::{ContentTypeOwned, MediaTypeParameters}; +/// +/// let value = ContentTypeOwned::try_from("multipart/form-data; boundary=----abc")?; +/// let mut parameters: MediaTypeParameters<'_> = value.parameters(); +/// let parameter = parameters.next().expect("boundary")?; +/// assert_eq!(parameter.name(), "boundary"); +/// assert_eq!(parameter.value(), b"----abc"); +/// assert!(parameters.next().is_none()); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct MediaTypeParameters<'a> { + bytes: &'a [u8], + scanner: ParameterScanner<'a>, + remaining: usize, +} + +impl ContentTypeOwned { + #[cfg(all(feature = "serde", feature = "headers-content-type"))] + pub(crate) fn field_value(&self) -> FieldValueRef<'_> { + self.value.as_field_value_ref() + } + + /// Constructs `type/subtype`. + /// + /// # Errors + /// + /// Returns an error if either component is not an HTTP token. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentTypeOwned; + /// + /// let value = ContentTypeOwned::new("multipart", "form-data")?; + /// assert_eq!(value.type_()?, "multipart"); + /// assert_eq!(value.subtype()?, "form-data"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn new(type_: impl AsRef, subtype: impl AsRef) -> Result { + let type_ = type_.as_ref(); + let subtype = subtype.as_ref(); + if !validate::token(type_.as_bytes()) || !validate::token(subtype.as_bytes()) { + return Err(DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidToken)); + } + let subtype_start = type_.len() + 1; + let total_len = subtype_start + subtype.len(); + let type_end = narrow_offset(type_.len())?; + let subtype_start_offset = narrow_offset(subtype_start)?; + let total_offset = narrow_offset(total_len)?; + let mut wire = String::with_capacity(total_len); + wire.push_str(type_); + wire.push('/'); + wire.push_str(subtype); + let value = content_type_value(wire); + Ok(Self { + value, + metadata: ContentTypeMetadata::Parsed(Box::new(ContentTypeHead { + type_start: 0, + type_end, + subtype_start: subtype_start_offset, + subtype_end: total_offset, + parameter_start: total_offset, + parameter_count: 0, + inline_parameters: [InlineParameter::EMPTY; INLINE_PARAMETER_CAPACITY], + inline_count: 0, + })), + }) + } + + /// Returns `application/json`. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentTypeOwned; + /// + /// let value = ContentTypeOwned::json(); + /// assert_eq!(value.into_field_value().as_bytes(), b"application/json"); + /// ``` + pub fn json() -> Self { + Self { + value: FieldValue::from_static("application/json"), + metadata: ContentTypeMetadata::Common { + type_end: 11, + subtype_end: 16, + }, + } + } + + /// Returns the top-level media type. + /// + /// # Errors + /// + /// Returns an error if the stored metadata and wire value disagree. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentTypeOwned; + /// + /// let value = ContentTypeOwned::try_from("application/json")?; + /// assert_eq!(value.type_()?, "application"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn type_(&self) -> Result<&str, DecodeError> { + component( + self.value.as_bytes(), + usize::try_from(self.metadata.head().type_start).unwrap_or(usize::MAX) + ..usize::try_from(self.metadata.head().type_end).unwrap_or(usize::MAX), + ) + } + + /// Returns the media subtype. + /// + /// # Errors + /// + /// Returns an error if the stored metadata and wire value disagree. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentTypeOwned; + /// + /// let value = ContentTypeOwned::try_from("text/html; charset=utf-8")?; + /// assert_eq!(value.subtype()?, "html"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn subtype(&self) -> Result<&str, DecodeError> { + component( + self.value.as_bytes(), + usize::try_from(self.metadata.head().subtype_start).unwrap_or(usize::MAX) + ..usize::try_from(self.metadata.head().subtype_end).unwrap_or(usize::MAX), + ) + } + + /// Iterates parameters in wire order without allocating. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentTypeOwned; + /// + /// let value = ContentTypeOwned::try_from("text/html; charset=utf-8")?; + /// let mut parameters = value.parameters(); + /// let parameter = parameters.next().expect("charset")?; + /// assert_eq!(parameter.name(), "charset"); + /// assert_eq!(parameter.value(), b"utf-8"); + /// assert!(parameters.next().is_none()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn parameters(&self) -> MediaTypeParameters<'_> { + MediaTypeParameters::lazy(self.value.as_bytes(), self.metadata.head()) + } + + /// Returns the first parameter matching `name` case-insensitively. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the wire value. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ContentTypeOwned::try_from("text/html; charset=utf-8")?; + /// assert_eq!(value.parameter("charset")?, Some(b"utf-8".as_slice())); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn parameter(&self, name: &str) -> Result, DecodeError> { + find_parameter(self.value.as_bytes(), self.metadata.head(), name) + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentTypeOwned; + /// + /// let value = ContentTypeOwned::try_from("multipart/form-data; boundary=----abc")?; + /// let field_value = value.into_field_value(); + /// assert_eq!( + /// field_value.as_bytes(), + /// b"multipart/form-data; boundary=----abc" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::shared::impl_field_value_conversion!(ContentTypeOwned, |value| value.value); + +impl<'a> ContentTypeView<'a> { + /// Returns the top-level media type. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the wire value. + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), Box> { + /// use http::header::CONTENT_TYPE; + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::ContentType; + /// + /// let mut map = HeaderMap::new(); + /// map.insert( + /// CONTENT_TYPE, + /// HeaderValue::from_static("multipart/form-data; boundary=----abc"), + /// ); + /// let view = ContentType::view(&map)?.expect("content type"); + /// assert_eq!(view.type_()?, "multipart"); + /// # Ok::<(), Box>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn type_(&self) -> Result<&'a str, DecodeError> { + component( + self.value.as_bytes(), + usize::try_from(self.metadata.head().type_start).unwrap_or(usize::MAX) + ..usize::try_from(self.metadata.head().type_end).unwrap_or(usize::MAX), + ) + } + + /// Returns the media subtype. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the wire value. + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), Box> { + /// use http::header::CONTENT_TYPE; + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::ContentType; + /// + /// let mut map = HeaderMap::new(); + /// map.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); + /// let view = ContentType::view(&map)?.expect("content type"); + /// assert_eq!(view.subtype()?, "json"); + /// # Ok::<(), Box>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn subtype(&self) -> Result<&'a str, DecodeError> { + component( + self.value.as_bytes(), + usize::try_from(self.metadata.head().subtype_start).unwrap_or(usize::MAX) + ..usize::try_from(self.metadata.head().subtype_end).unwrap_or(usize::MAX), + ) + } + + /// Iterates parameters in wire order. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), Box> { + /// use http::header::CONTENT_TYPE; + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::ContentType; + /// + /// let mut map = HeaderMap::new(); + /// map.insert( + /// CONTENT_TYPE, + /// HeaderValue::from_static("text/html; charset=utf-8"), + /// ); + /// let view = ContentType::view(&map)?.expect("content type"); + /// let parameter = view.parameters().next().expect("charset")?; + /// assert_eq!(parameter.name(), "charset"); + /// assert_eq!(parameter.value(), b"utf-8"); + /// # Ok::<(), Box>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn parameters(&self) -> MediaTypeParameters<'a> { + MediaTypeParameters::lazy(self.value.as_bytes(), self.metadata.head()) + } + + /// Returns the first parameter matching `name` case-insensitively. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the wire value. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ContentTypeOwned::try_from("text/html; charset=utf-8")?; + /// assert_eq!(value.parameter("charset")?, Some(b"utf-8".as_slice())); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn parameter(&self, name: &str) -> Result, DecodeError> { + find_parameter(self.value.as_bytes(), self.metadata.head(), name) + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), Box> { + /// use http::header::CONTENT_TYPE; + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::ContentType; + /// + /// let mut map = HeaderMap::new(); + /// map.insert(CONTENT_TYPE, HeaderValue::from_static("application/json")); + /// let view = ContentType::view(&map)?.expect("content type"); + /// assert_eq!(view.as_field_value().as_bytes(), b"application/json"); + /// # Ok::<(), Box>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub const fn as_field_value(&self) -> FieldValueRef<'a> { + self.value + } +} + +impl<'a> MediaTypeParameterView<'a> { + /// Returns the parameter name. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentTypeOwned; + /// + /// let value = ContentTypeOwned::try_from("text/html; charset=utf-8")?; + /// let parameter = value.parameters().next().expect("charset")?; + /// assert_eq!(parameter.name(), "charset"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn name(self) -> &'a str { + self.name + } + + /// Returns the raw token or quoted-string parameter value. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ContentTypeOwned::try_from("text/html; charset=utf-8")?; + /// assert_eq!(value.parameter("charset")?, Some(b"utf-8".as_slice())); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn value(self) -> &'a [u8] { + self.value + } +} + +impl Field for ContentType { + type View<'a> = ContentTypeView<'a>; + type Owned = ContentTypeOwned; + + fn name() -> &'static FieldName { + &FieldName::ContentType + } + + fn view_with(source: &S, mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_list_item_limit(b';', false)?; + let value = lines.exactly_one()?; + let metadata = parse_metadata_with(value.as_bytes(), mode)?; + Ok(Some(ContentTypeView { value, metadata })) + } + + fn owned_with(source: &S, mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_list_item_limit(b';', false)?; + let owned = lines.exactly_one_owned()?; + let metadata = parse_metadata_with(owned.as_bytes(), mode)?; + Ok(Some(ContentTypeOwned { value: owned, metadata })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::single(value.value)) + } +} + +super::shared::impl_string_conversions!(ContentTypeOwned, &FieldName::ContentType, super::invalid_syntax, value); + +impl TryFrom for ContentTypeOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + let metadata = parse_metadata(value.as_bytes())?; + Ok(Self { value, metadata }) + } +} + +fn content_type_value(wire: String) -> FieldValue { + FieldValue::try_from(wire).expect("validated media type components produce a field value") +} + +impl<'a> MediaTypeParameters<'a> { + fn lazy(bytes: &'a [u8], head: ContentTypeHead) -> Self { + Self { + bytes, + scanner: ParameterScanner::new(bytes, usize::try_from(head.parameter_start).unwrap_or(usize::MAX)), + remaining: usize::try_from(head.parameter_count).unwrap_or(usize::MAX), + } + } +} + +impl<'a> Iterator for MediaTypeParameters<'a> { + type Item = Result, DecodeError>; + + fn next(&mut self) -> Option { + let range = self.scanner.next()?; + self.remaining = self.remaining.saturating_sub(1); + Some(range.and_then(|range| parameter_from_range(self.bytes, &range))) + } + + fn size_hint(&self) -> (usize, Option) { + (0, Some(self.remaining)) + } +} + +fn parse_metadata(bytes: &[u8]) -> Result { + if let Some(metadata) = common_metadata(bytes) { + return Ok(metadata); + } + let mut head = parse_prefix(bytes)?; + for parameter in ParameterScanner::new(bytes, usize::try_from(head.parameter_start).unwrap_or(usize::MAX)) { + head.observe_parameter(¶meter?)?; + } + Ok(ContentTypeMetadata::Parsed(Box::new(head))) +} + +#[inline] +fn parse_metadata_with(bytes: &[u8], mode: crate::DecodeMode) -> Result { + if mode == crate::DecodeMode::Strict { + return parse_metadata(bytes); + } + if let Ok(metadata) = parse_metadata(bytes) { + return Ok(metadata); + } + let mut head = parse_prefix_relaxed(bytes)?; + for parameter in ParameterScanner::new(bytes, usize::try_from(head.parameter_start).unwrap_or(usize::MAX)) { + head.observe_parameter(¶meter?)?; + } + Ok(ContentTypeMetadata::Parsed(Box::new(head))) +} + +fn common_metadata(bytes: &[u8]) -> Option { + if bytes == b"application/json; charset=utf-8" { + return Some(ContentTypeMetadata::ApplicationJsonUtf8); + } + let (type_end, subtype_start, subtype_end) = match bytes { + b"application/json" => (11, 12, 16), + b"application/octet-stream" => (11, 12, 24), + b"text/html" => (4, 5, 9), + b"text/plain" => (4, 5, 10), + b"text/css" => (4, 5, 8), + b"application/javascript" => (11, 12, 22), + b"text/html; charset=utf-8" => { + return Some(ContentTypeMetadata::TextHtmlUtf8); + } + _ => return None, + }; + debug_assert_eq!( + subtype_start, + type_end + 1, + "known media type subtype must begin immediately after the slash" + ); + Some(ContentTypeMetadata::Common { + type_end: u8::try_from(type_end).ok()?, + subtype_end: u8::try_from(subtype_end).ok()?, + }) +} + +fn common_parameter_head( + type_end: u32, + subtype_end: u32, + name_start: u16, + name_end: u16, + value_start: u16, + value_end: u16, +) -> ContentTypeHead { + ContentTypeHead { + type_start: 0, + type_end, + subtype_start: type_end + 1, + subtype_end, + parameter_start: subtype_end, + parameter_count: 1, + inline_parameters: [ + InlineParameter { + name_start, + name_end, + value_start, + value_end, + }, + InlineParameter::EMPTY, + ], + inline_count: 1, + } +} + +fn parse_prefix(bytes: &[u8]) -> Result { + let mut position = 0; + let type_range = take_token(bytes, &mut position)?; + if bytes.get(position) != Some(&b'/') { + return Err(super::invalid_syntax(&FieldName::ContentType)); + } + position += 1; + let subtype_range = take_token(bytes, &mut position)?; + let parameter_start = position; + Ok(ContentTypeHead { + type_start: narrow_offset(type_range.start)?, + type_end: narrow_offset(type_range.end)?, + subtype_start: narrow_offset(subtype_range.start)?, + subtype_end: narrow_offset(subtype_range.end)?, + parameter_start: narrow_offset(parameter_start)?, + parameter_count: 0, + inline_parameters: [InlineParameter::EMPTY; INLINE_PARAMETER_CAPACITY], + inline_count: 0, + }) +} + +fn parse_prefix_relaxed(bytes: &[u8]) -> Result { + let mut position = 0; + let type_range = take_token(bytes, &mut position)?; + skip_ows(bytes, &mut position); + if bytes.get(position) != Some(&b'/') { + return Err(super::invalid_syntax(&FieldName::ContentType)); + } + position += 1; + skip_ows(bytes, &mut position); + let subtype_range = take_token(bytes, &mut position)?; + let parameter_start = position; + Ok(ContentTypeHead { + type_start: narrow_offset(type_range.start)?, + type_end: narrow_offset(type_range.end)?, + subtype_start: narrow_offset(subtype_range.start)?, + subtype_end: narrow_offset(subtype_range.end)?, + parameter_start: narrow_offset(parameter_start)?, + parameter_count: 0, + inline_parameters: [InlineParameter::EMPTY; INLINE_PARAMETER_CAPACITY], + inline_count: 0, + }) +} + +impl ContentTypeHead { + fn observe_parameter(&mut self, parameter: &ParameterRange) -> Result<(), DecodeError> { + if usize::try_from(self.parameter_count).unwrap_or(usize::MAX) == usize::from(self.inline_count) + && usize::from(self.inline_count) < INLINE_PARAMETER_CAPACITY + && let Some(compact) = InlineParameter::from_range(parameter) + { + self.inline_parameters[usize::from(self.inline_count)] = compact; + self.inline_count += 1; + } + self.parameter_count = self + .parameter_count + .checked_add(1) + .ok_or_else(|| DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidNumber))?; + Ok(()) + } +} + +#[derive(Debug)] +struct ParameterScanner<'a> { + bytes: &'a [u8], + position: usize, +} + +impl<'a> ParameterScanner<'a> { + const fn new(bytes: &'a [u8], position: usize) -> Self { + Self { bytes, position } + } +} + +impl Iterator for ParameterScanner<'_> { + type Item = Result; + + fn next(&mut self) -> Option { + loop { + skip_ows(self.bytes, &mut self.position); + if self.position == self.bytes.len() { + return None; + } + if self.bytes.get(self.position) != Some(&b';') { + self.position = self.bytes.len(); + return Some(Err(super::invalid_syntax(&FieldName::ContentType))); + } + self.position += 1; + skip_ows(self.bytes, &mut self.position); + if self.position == self.bytes.len() || self.bytes.get(self.position) == Some(&b';') { + continue; + } + break; + } + + let name = match take_token(self.bytes, &mut self.position) { + Ok(name) => name, + Err(error) => { + self.position = self.bytes.len(); + return Some(Err(error)); + } + }; + if self.bytes.get(self.position) != Some(&b'=') { + self.position = self.bytes.len(); + return Some(Err(super::invalid_syntax(&FieldName::ContentType))); + } + self.position += 1; + let value = match take_parameter_value(self.bytes, &mut self.position) { + Ok(value) => value, + Err(error) => { + self.position = self.bytes.len(); + return Some(Err(error)); + } + }; + Some(narrow_parameter_range(name, value)) + } +} + +fn take_token(bytes: &[u8], position: &mut usize) -> Result, DecodeError> { + let start = *position; + while bytes + .get(*position) + .is_some_and(|byte| TOKEN_BITS[usize::from(*byte >> 6)] & (1_u64 << (*byte & 63)) != 0) + { + *position += 1; + } + + if *position > start { + Ok(start..*position) + } else { + Err(DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidToken)) + } +} + +// Four words avoid branch-heavy token classification without a 256-byte lookup table. +static TOKEN_BITS: [u64; 4] = { + let mut table = [0; 4]; + let mut byte = 0_u8; + loop { + if validate::token_byte(byte) { + table[(byte >> 6) as usize] |= 1_u64 << (byte & 63); + } + if byte == u8::MAX { + break; + } + byte += 1; + } + table +}; + +fn take_parameter_value(bytes: &[u8], position: &mut usize) -> Result, DecodeError> { + if bytes.get(*position) != Some(&b'"') { + let mut end = *position; + let token = take_token(bytes, &mut end)?; + *position = end; + return Ok(token); + } + let start = *position; + *position += 1; + let mut escaped = false; + while let Some(byte) = bytes.get(*position).copied() { + *position += 1; + if escaped { + if !matches!(byte, b'\t' | b' '..=b'~' | 0x80..=0xff) { + return Err(super::invalid_syntax(&FieldName::ContentType)); + } + escaped = false; + } else if byte == b'\\' { + escaped = true; + } else if byte == b'"' { + return Ok(start..*position); + } else if !matches!(byte, b'\t' | b' ' | b'!' | b'#'..=b'[' | b']'..=b'~' | 0x80..=0xff) { + return Err(super::invalid_syntax(&FieldName::ContentType)); + } + } + Err(DecodeError::new(&FieldName::ContentType, DecodeErrorKind::UnterminatedQuote)) +} + +fn skip_ows(bytes: &[u8], position: &mut usize) { + while bytes.get(*position).is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + *position += 1; + } +} + +fn component(bytes: &[u8], range: Range) -> Result<&str, DecodeError> { + let component = bytes.get(range).ok_or_else(|| super::invalid_syntax(&FieldName::ContentType))?; + str::from_utf8(component).map_err(|_invalid| DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidUtf8)) +} + +fn parameter_from_range<'a>(bytes: &'a [u8], range: &ParameterRange) -> Result, DecodeError> { + let name = component(bytes, widen_range(range.name.clone()))?; + let value = bytes + .get(widen_range(range.value.clone())) + .ok_or_else(|| super::invalid_syntax(&FieldName::ContentType))?; + Ok(MediaTypeParameterView { name, value }) +} + +fn find_parameter<'a>(bytes: &'a [u8], head: ContentTypeHead, name: &str) -> Result, DecodeError> { + for parameter in &head.inline_parameters[..usize::from(head.inline_count)] { + let parameter = parameter_from_inline(bytes, *parameter)?; + if validate::eq_ignore_ascii_case(parameter.name.as_bytes(), name.as_bytes()) { + return Ok(Some(parameter.value)); + } + } + if usize::try_from(head.parameter_count).unwrap_or(usize::MAX) <= usize::from(head.inline_count) { + return Ok(None); + } + for parameter in + ParameterScanner::new(bytes, usize::try_from(head.parameter_start).unwrap_or(usize::MAX)).skip(usize::from(head.inline_count)) + { + let parameter = parameter_from_range(bytes, ¶meter?)?; + if validate::eq_ignore_ascii_case(parameter.name.as_bytes(), name.as_bytes()) { + return Ok(Some(parameter.value)); + } + } + Ok(None) +} + +fn narrow_offset(offset: usize) -> Result { + u32::try_from(offset).map_err(|_overflow| DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidNumber)) +} + +fn narrow_range(range: Range) -> Result, DecodeError> { + Ok(narrow_offset(range.start)?..narrow_offset(range.end)?) +} + +#[inline] +fn narrow_parameter_range(name: Range, value: Range) -> Result { + Ok(ParameterRange { + name: narrow_range(name)?, + value: narrow_range(value)?, + }) +} + +fn widen_range(range: Range) -> Range { + usize::try_from(range.start).unwrap_or(usize::MAX)..usize::try_from(range.end).unwrap_or(usize::MAX) +} + +fn parameter_from_inline(bytes: &[u8], range: InlineParameter) -> Result, DecodeError> { + let name = component(bytes, range.name_range())?; + let value = bytes + .get(range.value_range()) + .ok_or_else(|| super::invalid_syntax(&FieldName::ContentType))?; + Ok(MediaTypeParameterView { name, value }) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #![expect( + clippy::assertions_on_result_states, + reason = "tests classify parser outcomes without needing successful values" + )] + + use std::slice; + + use super::{ + ContentType, ContentTypeHead, ContentTypeMetadata, ContentTypeOwned, InlineParameter, ParameterRange, common_metadata, component, + narrow_offset, narrow_parameter_range, parameter_from_inline, parameter_from_range, parse_metadata_with, take_parameter_value, + take_token, + }; + use crate::sink::FieldSink; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeErrorKind, DecodeMode, FieldName, FieldValue, TestSink}; + + struct Source(FieldValue); + + impl FieldSource for Source { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::ContentType) + .then(|| FieldLines::from_slice(name, slice::from_ref(&self.0))) + .flatten() + } + } + + #[test] + fn ordinary_view_uses_lazy_inline_metadata() { + let source = Source(FieldValue::from_static("application/json; charset=utf-8")); + let view = ContentType::view(&source).expect("valid media type").expect("media type present"); + assert_eq!(view.parameters().count(), 1); + assert_eq!(view.type_(), Ok("application")); + } + + #[test] + fn token_prefix_scanning_checks_every_byte_and_preserves_the_delimiter() { + const TOKEN_CHARACTERS: &[u8] = b"!#$%&'*+-.^_`|~0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz"; + + for byte in u8::MIN..=u8::MAX { + let bytes = [b'/', byte, b'/']; + let mut position = 1; + let token = take_token(&bytes, &mut position); + if TOKEN_CHARACTERS.contains(&byte) { + assert_eq!(token, Ok(1..2), "{byte}"); + assert_eq!(position, 2, "{byte}"); + } else { + assert_eq!(token.unwrap_err().kind(), DecodeErrorKind::InvalidToken, "{byte}"); + assert_eq!(position, 1, "{byte}"); + } + let prefixed = [b'/', b'a', byte, b'/']; + let mut position = 1; + let end = if TOKEN_CHARACTERS.contains(&byte) { 3 } else { 2 }; + assert_eq!(take_token(&prefixed, &mut position), Ok(1..end), "{byte}"); + assert_eq!(position, end, "{byte}"); + } + } + + #[test] + fn parameter_lookup_falls_back_beyond_inline_capacity() { + let content_type = ContentTypeOwned::try_from("text/plain; a=1; b=2; c=3; d=4").expect("valid media type"); + assert_eq!(content_type.metadata.head().inline_count, 2); + assert_eq!(content_type.parameter("a"), Ok(Some(b"1".as_slice()))); + assert_eq!(content_type.parameter("b"), Ok(Some(b"2".as_slice()))); + assert_eq!(content_type.parameter("d"), Ok(Some(b"4".as_slice()))); + assert_eq!(content_type.parameter("missing"), Ok(None)); + } + + #[test] + fn common_and_constructed_types_cover_accessors_and_round_trips() { + for (wire, type_, subtype) in [ + ("application/json", "application", "json"), + ("application/octet-stream", "application", "octet-stream"), + ("text/html", "text", "html"), + ("text/plain", "text", "plain"), + ("text/css", "text", "css"), + ("application/javascript", "application", "javascript"), + ] { + let content_type = ContentTypeOwned::try_from(wire).expect("common media type"); + assert_eq!(content_type.type_(), Ok(type_)); + assert_eq!(content_type.subtype(), Ok(subtype)); + assert_eq!(content_type.parameters().size_hint(), (0, Some(0))); + assert_eq!(content_type.parameter("missing"), Ok(None)); + } + + let constructed = ContentTypeOwned::new("image", "svg+xml").expect("valid tokens"); + assert_eq!(constructed.type_(), Ok("image")); + assert_eq!(constructed.subtype(), Ok("svg+xml")); + assert_eq!(constructed.clone().into_field_value(), FieldValue::from_static("image/svg+xml")); + assert_eq!(ContentTypeOwned::json().type_(), Ok("application")); + assert_eq!(ContentType::json().subtype(), Ok("json")); + assert!(ContentTypeOwned::new("", "plain").is_err()); + assert!(ContentTypeOwned::new("text", "bad value").is_err()); + assert_eq!( + ContentTypeOwned::try_from(String::from("audio/ogg")) + .expect("valid owned media type") + .subtype(), + Ok("ogg") + ); + + let html_utf8 = ContentTypeOwned::try_from("text/html; charset=utf-8").expect("common parameter form"); + assert_eq!(html_utf8.parameters().count(), 1); + assert_eq!(html_utf8.parameter("charset"), Ok(Some(b"utf-8".as_slice()))); + + let mut table = TestSink::new(); + ContentType::insert(&mut table, constructed).expect("table accepts content type"); + let view = ContentType::view(&table).expect("valid content type").expect("present"); + assert_eq!(view.type_(), Ok("image")); + assert_eq!(view.subtype(), Ok("svg+xml")); + assert_eq!(view.parameter("missing"), Ok(None)); + assert_eq!(view.parameters().count(), 0); + assert_eq!(view.as_field_value().as_bytes(), b"image/svg+xml"); + assert!(ContentType::owned(&table).expect("valid owned content type").is_some()); + table.remove_values(&FieldName::ContentType); + assert!(ContentType::view(&table).expect("absence is valid").is_none()); + assert!(ContentType::owned(&table).expect("absence is valid").is_none()); + } + + #[test] + fn parameters_cover_inline_lazy_quoted_and_scanner_error_paths() { + let value = ContentTypeOwned::try_from("text/plain; Charset=utf-8; note=\"a\\\\b\"; third=value").expect("valid parameters"); + let parameters = value + .parameters() + .collect::, _>>() + .expect("stored parameters remain valid"); + assert_eq!(parameters.len(), 3); + assert_eq!(parameters[0].name(), "Charset"); + assert_eq!(parameters[0].value(), b"utf-8"); + assert_eq!(parameters[1].value(), b"\"a\\\\b\""); + assert_eq!(value.parameter("charset"), Ok(Some(b"utf-8".as_slice()))); + assert_eq!(value.parameter("THIRD"), Ok(Some(b"value".as_slice()))); + assert_eq!( + ContentTypeOwned::try_from("text/plain; ; charset=utf-8;") + .expect("empty parameter slots are valid") + .parameters() + .count(), + 1 + ); + + for wire in [ + "text", + "text/", + "/plain", + "text plain", + "text/plain trailing", + "text/plain; name", + "text/plain; =value", + "text/plain; name=", + "text/plain; name=\"unterminated", + "text/plain; name=\"bad\\\n\"", + ] { + assert!(ContentTypeOwned::try_from(wire).is_err(), "{wire:?}"); + } + assert!(ContentTypeOwned::try_from("text/plain; name =value").is_err()); + assert!(ContentTypeOwned::try_from("text/plain; name= value").is_err()); + + let bytes = b"\"quoted\""; + let mut position = 0; + assert_eq!(take_parameter_value(bytes, &mut position), Ok(0..8)); + let mut position = 0; + assert!(take_parameter_value(b"\"unterminated", &mut position).is_err()); + let mut position = 0; + assert!(take_parameter_value(b"\"bad\\\n\"", &mut position).is_err()); + let mut position = 0; + assert!(take_parameter_value(b"\"bad\n\"", &mut position).is_err()); + } + + #[test] + fn relaxed_parsing_and_private_metadata_errors_are_covered() { + assert!(parse_metadata_with(b"text / plain", DecodeMode::Strict).is_err()); + let relaxed = parse_metadata_with(b"text / plain", DecodeMode::Relaxed).expect("relaxed spacing"); + let head = relaxed.head(); + assert_eq!(head.type_end, 4); + assert_eq!(head.subtype_start, 7); + assert!(parse_metadata_with(b"bad", DecodeMode::Relaxed).is_err()); + assert!(parse_metadata_with(b"text/plain", DecodeMode::Relaxed).is_ok()); + assert!(parse_metadata_with(b"text / plain; charset=utf-8", DecodeMode::Relaxed).is_ok()); + + assert!(common_metadata(b"unknown/type").is_none()); + assert!(matches!( + common_metadata(b"application/json; charset=utf-8"), + Some(ContentTypeMetadata::ApplicationJsonUtf8) + )); + assert!(matches!( + common_metadata(b"text/html; charset=utf-8"), + Some(ContentTypeMetadata::TextHtmlUtf8) + )); + + assert!(component(b"abc", 4..5).is_err()); + assert_eq!( + component(&[0xff], 0..1).expect_err("component is not UTF-8").kind(), + DecodeErrorKind::InvalidUtf8 + ); + assert!(narrow_offset(u32::MAX as usize + 1).is_err()); + + let bad_range = ParameterRange { name: 0..5, value: 6..7 }; + assert!(parameter_from_range(b"a=b", &bad_range).is_err()); + assert!( + parameter_from_inline( + b"a=b", + InlineParameter { + name_start: 0, + name_end: 1, + value_start: 9, + value_end: 10, + } + ) + .is_err() + ); + assert!( + parameter_from_inline( + b"a=b", + InlineParameter { + name_start: 9, + name_end: 10, + value_start: 2, + value_end: 3, + } + ) + .is_err() + ); + assert!( + InlineParameter::from_range(&ParameterRange { + name: 0..(u32::from(u16::MAX) + 1), + value: 0..1, + }) + .is_none() + ); + + let mut overflowing = ContentTypeHead { + type_start: 0, + type_end: 1, + subtype_start: 2, + subtype_end: 3, + parameter_start: 3, + parameter_count: u32::MAX, + inline_parameters: [InlineParameter::EMPTY; 2], + inline_count: 0, + }; + assert!(overflowing.observe_parameter(&ParameterRange { name: 0..1, value: 2..3 }).is_err()); + } + + #[test] + fn private_conversion_error_paths_are_covered() { + let overflow = u32::MAX as usize + 1; + assert!(narrow_parameter_range(overflow..overflow, 0..0).is_err()); + assert!(narrow_parameter_range(0..0, overflow..overflow).is_err()); + assert!(ContentTypeOwned::try_from(String::from("\n")).is_err()); + assert!(parameter_from_range(b"", &ParameterRange { name: 0..0, value: 1..2 }).is_err()); + } +} diff --git a/crates/http_headers/src/headers/cors/access_control_allow_credentials.rs b/crates/http_headers/src/headers/cors/access_control_allow_credentials.rs new file mode 100644 index 000000000..dfba9c1e6 --- /dev/null +++ b/crates/http_headers/src/headers/cors/access_control_allow_credentials.rs @@ -0,0 +1,356 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::marker::PhantomData; + +use super::shared::{invalid_syntax, trimmed_range}; +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef}; + +/// Defines the `Access-Control-Allow-Credentials` header. +/// +/// # Specification +/// +/// Defined by the Fetch standard's +/// [CORS protocol and credentials section](https://fetch.spec.whatwg.org/#http-access-control-allow-credentials). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{ +/// AccessControlAllowCredentials, AccessControlAllowCredentialsOwned, +/// }; +/// +/// let mut map = HeaderMap::new(); +/// AccessControlAllowCredentials::insert(&mut map, AccessControlAllowCredentialsOwned::allow())?; +/// assert!(AccessControlAllowCredentials::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct AccessControlAllowCredentials { + _private: (), +} + +/// Owned value for the `Access-Control-Allow-Credentials` header. +/// +/// The grammar admits a single field value, so the type carries no storage and +/// always encodes the canonical `true`; whitespace framing a decoded value is +/// not preserved. +/// +/// # Specification +/// +/// Defined by the Fetch standard's [CORS protocol and credentials section]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::AccessControlAllowCredentialsOwned::allow(); +/// assert_eq!(value.into_field_value(), "true"); +/// ``` +/// +/// `Access-Control-Allow-Credentials: true` is the only valid value. +/// +/// [CORS protocol and credentials section]: https://fetch.spec.whatwg.org/#http-access-control-allow-credentials +#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)] +pub struct AccessControlAllowCredentialsOwned; + +/// The only field value `Access-Control-Allow-Credentials` accepts. +const ALLOW_CREDENTIALS_TRUE: &str = "true"; + +/// The canonical field value, handed out in place of decoded storage. +static ALLOW_CREDENTIALS_VALUE: FieldValue = FieldValue::from_static(ALLOW_CREDENTIALS_TRUE); + +/// Borrowed value for the `Access-Control-Allow-Credentials` header. +#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{AccessControlAllowCredentials, AccessControlAllowCredentialsView}; +/// use http_headers::source::{FieldLines, FieldSource}; +/// use http_headers::{Field, FieldName}; +/// +/// struct Source; +/// +/// impl FieldSource for Source { +/// fn lines(&self, name: &'static FieldName) -> Option> { +/// (name == &FieldName::AccessControlAllowCredentials) +/// .then(|| FieldLines::single(name, b"true")) +/// } +/// } +/// +/// let view: AccessControlAllowCredentialsView<'_> = +/// AccessControlAllowCredentials::view(&Source)?.expect("present"); +/// assert_eq!(view.as_str(), "true"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct AccessControlAllowCredentialsView<'a> { + values: PhantomData>, +} + +impl AccessControlAllowCredentialsOwned { + /// Constructs the only valid value, case-sensitive `true`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowCredentialsOwned; + /// + /// let value = AccessControlAllowCredentialsOwned::allow(); + /// assert_eq!(value.to_string(), "true"); + /// ``` + pub const fn allow() -> Self { + Self + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowCredentialsOwned; + /// + /// let value = AccessControlAllowCredentialsOwned::allow(); + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value, "true"); + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(AccessControlAllowCredentialsOwned, |_value| ALLOW_CREDENTIALS_VALUE.clone()); + +impl fmt::Display for AccessControlAllowCredentialsOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(ALLOW_CREDENTIALS_TRUE) + } +} + +impl<'a> AccessControlAllowCredentialsView<'a> { + /// Returns the semantic value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowCredentials; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source; + /// + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::AccessControlAllowCredentials) + /// .then(|| FieldLines::single(name, b"true")) + /// } + /// } + /// + /// let view = AccessControlAllowCredentials::view(&Source)?.expect("present"); + /// assert_eq!(view.as_str(), "true"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::unused_self, + reason = "the validated zero-sized view exposes the same accessor shape as value-carrying views" + )] + pub const fn as_str(self) -> &'a str { + ALLOW_CREDENTIALS_TRUE + } + + /// Returns the canonical field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowCredentials; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source; + /// + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::AccessControlAllowCredentials) + /// .then(|| FieldLines::single(name, b"true")) + /// } + /// } + /// + /// let view = AccessControlAllowCredentials::view(&Source)?.expect("present"); + /// assert_eq!(view.as_field_value(), "true"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::unused_self, + reason = "the validated zero-sized view exposes the same accessor shape as value-carrying views" + )] + pub fn as_field_value(self) -> FieldValueRef<'a> { + ALLOW_CREDENTIALS_VALUE.as_field_value_ref() + } +} + +impl Field for AccessControlAllowCredentials { + type View<'a> = AccessControlAllowCredentialsView<'a>; + type Owned = AccessControlAllowCredentialsOwned; + + fn name() -> &'static FieldName { + &FieldName::AccessControlAllowCredentials + } + + #[inline] + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_custom_source()?; + if is_credentials_true(lines.exactly_one()?.as_bytes()) { + Ok(Some(AccessControlAllowCredentialsView { values: PhantomData })) + } else { + Err(invalid_syntax(&FieldName::AccessControlAllowCredentials)) + } + } + + fn owned_with(source: &S, mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + Self::view_with(source, mode).map(|view| view.map(|_| AccessControlAllowCredentialsOwned)) + } + + fn insert(sink: &mut S, _value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::single(ALLOW_CREDENTIALS_VALUE.clone())) + } +} + +super::super::shared::impl_string_conversions!( + AccessControlAllowCredentialsOwned, + &FieldName::AccessControlAllowCredentials, + invalid_syntax, + value +); + +impl TryFrom for AccessControlAllowCredentialsOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + if is_credentials_true(value.as_bytes()) { + Ok(Self) + } else { + Err(invalid_syntax(&FieldName::AccessControlAllowCredentials)) + } + } +} +#[inline] +fn is_credentials_true(bytes: &[u8]) -> bool { + bytes == b"true" || is_trimmed_credentials_true(bytes) +} + +#[cold] +fn is_trimmed_credentials_true(bytes: &[u8]) -> bool { + bytes.get(trimmed_range(bytes)) == Some(b"true".as_slice()) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{AccessControlAllowCredentials, AccessControlAllowCredentialsOwned, is_credentials_true}; + use crate::headers::cors::test_map::TestMap; + use crate::{DecodeErrorKind, FieldName, FieldValue}; + + #[test] + fn canonical_value_construction_display_and_conversion_are_storage_free() { + let value = AccessControlAllowCredentialsOwned::allow(); + assert_eq!(value, AccessControlAllowCredentialsOwned); + assert_eq!(value.to_string(), "true"); + assert_eq!(value.into_field_value(), "true"); + + assert_eq!( + AccessControlAllowCredentialsOwned::try_from(" true ").expect("whitespace framed true"), + value + ); + assert_eq!( + AccessControlAllowCredentialsOwned::try_from(String::from("\ttrue\t")).expect("owned true string"), + value + ); + assert_eq!( + AccessControlAllowCredentialsOwned::try_from(FieldValue::from_static("true")).expect("true field"), + value + ); + + assert!(is_credentials_true(b"true")); + assert!(is_credentials_true(b" \ttrue\t ")); + assert!(!is_credentials_true(b"TRUE")); + assert!(!is_credentials_true(b"false")); + } + + #[test] + fn header_decode_insert_absence_and_errors_cover_all_paths() { + let source = TestMap::new(&FieldName::AccessControlAllowCredentials, vec![FieldValue::from_static(" true ")]); + let view = AccessControlAllowCredentials::view(&source) + .expect("valid credentials view") + .expect("present"); + assert_eq!(view.as_str(), "true"); + assert_eq!(view.as_field_value(), "true"); + assert_eq!( + AccessControlAllowCredentials::owned(&source) + .expect("valid credentials owned") + .expect("present"), + AccessControlAllowCredentialsOwned + ); + + let mut sink = TestMap::new(&FieldName::Accept, Vec::new()); + AccessControlAllowCredentials::insert(&mut sink, AccessControlAllowCredentialsOwned::allow()).expect("insert credentials"); + assert_eq!(sink.name, &FieldName::AccessControlAllowCredentials); + assert_eq!(sink.values, [FieldValue::from_static("true")]); + + let absent = TestMap::new(&FieldName::Accept, Vec::new()); + assert!(AccessControlAllowCredentials::view(&absent).expect("absent").is_none()); + assert!(AccessControlAllowCredentials::owned(&absent).expect("absent").is_none()); + + for wire in ["TRUE", "false", "true value"] { + let error = AccessControlAllowCredentialsOwned::try_from(wire).expect_err("only lowercase true is accepted"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + } + let error = AccessControlAllowCredentialsOwned::try_from(String::from("bad\nvalue")).expect_err("invalid field bytes"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!( + AccessControlAllowCredentialsOwned::try_from("bad\nvalue") + .expect_err("invalid borrowed field bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let invalid = TestMap::new(&FieldName::AccessControlAllowCredentials, vec![FieldValue::from_static("false")]); + assert_eq!( + AccessControlAllowCredentials::view(&invalid) + .expect_err("invalid credentials header") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let duplicate = TestMap::new( + &FieldName::AccessControlAllowCredentials, + vec![FieldValue::from_static("true"), FieldValue::from_static("true")], + ); + assert_eq!( + AccessControlAllowCredentials::view(&duplicate) + .expect_err("singleton header") + .kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + } +} diff --git a/crates/http_headers/src/headers/cors/access_control_allow_headers.rs b/crates/http_headers/src/headers/cors/access_control_allow_headers.rs new file mode 100644 index 000000000..9b8227f8c --- /dev/null +++ b/crates/http_headers/src/headers/cors/access_control_allow_headers.rs @@ -0,0 +1,40 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::super::FieldNameView; +use super::shared::{CorsList, CorsListView, define_header_name_list, impl_header_name_wildcard}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef, validate}; + +define_header_name_list!( + AccessControlAllowHeaders, + AccessControlAllowHeadersOwned, + AccessControlAllowHeadersView, + "Access-Control-Allow-Headers", + &FieldName::AccessControlAllowHeaders, + true, + "Defined by the Fetch standard's [CORS protocol and credentials section](https://fetch.spec.whatwg.org/#http-access-control-allow-headers).", + "`Access-Control-Allow-Headers: Content-Type, Authorization` lists field names, `Access-Control-Allow-Headers: *` uses a wildcard, and an empty field value is accepted." +); + +impl_header_name_wildcard!(AccessControlAllowHeadersOwned, AccessControlAllowHeadersView); + +impl AccessControlAllowHeadersOwned { + /// Constructs a present field with an empty field-name list. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let value = AccessControlAllowHeadersOwned::empty(); + /// assert!(value.is_empty()); + /// assert_eq!(value.len(), 0); + /// ``` + pub fn empty() -> Self { + Self(CorsList::empty()) + } +} diff --git a/crates/http_headers/src/headers/cors/access_control_allow_methods.rs b/crates/http_headers/src/headers/cors/access_control_allow_methods.rs new file mode 100644 index 000000000..12faa05fd --- /dev/null +++ b/crates/http_headers/src/headers/cors/access_control_allow_methods.rs @@ -0,0 +1,703 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::super::MethodView; +use super::shared::{CorsList, CorsListView, method_ref_validated, validate_list, validate_single_list}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef, validate}; + +macro_rules! define_method_list { + ($descriptor:ident, $owned:ident, $borrowed:ident, $name:expr) => { + /// Defines the `Access-Control-Allow-Methods` header. + /// + /// # Specification + /// + /// Defined by the Fetch standard's + /// [CORS protocol and credentials section](https://fetch.spec.whatwg.org/#http-access-control-allow-methods). + #[derive(Debug)] + pub struct $descriptor { + _private: (), + } + + /// Owned value for the `Access-Control-Allow-Methods` header. + /// + /// # Specification + /// + /// Defined by the Fetch standard's [CORS protocol and credentials section]. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// + /// let value = http_headers::headers::AccessControlAllowMethodsOwned::try_from( + /// FieldValue::from_static("GET, POST"), + /// )?; + /// assert_eq!(value.iter().count(), 2); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + /// + /// `Access-Control-Allow-Methods: GET, POST, PUT` lists methods, + /// `Access-Control-Allow-Methods: *` uses a wildcard, and an empty field + /// value is also accepted. + /// + /// [CORS protocol and credentials section]: https://fetch.spec.whatwg.org/#http-access-control-allow-methods + #[derive(Clone, Eq, Hash, PartialEq)] + pub struct $owned(CorsList); + + /// Borrowed value for the `Access-Control-Allow-Methods` header. + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowMethods; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("GET, POST"), + /// ); + /// let value = AccessControlAllowMethods::view(&headers)?.expect("present"); + /// assert_eq!(value.len(), 2); + /// let fmt_output = format!("{value:?}"); + /// assert!(fmt_output.contains("method_count")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub struct $borrowed<'a>(CorsListView<'a>); + + impl fmt::Debug for $owned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct(stringify!($owned)) + .field("value_count", &self.0.value_count()) + .field("method_count", &self.len()) + .finish() + } + } + + impl fmt::Display for $owned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + crate::headers::shared::fmt_ascii_values(self.field_values(), f) + } + } + + impl fmt::Debug for $borrowed<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct(stringify!($borrowed)) + .field("value_count", &self.0.values.len()) + .field("method_count", &self.len()) + .finish() + } + } + + impl $owned { + /// Constructs one canonical field line from validated methods. + /// + /// Empty input is valid for this response header. Duplicate methods + /// are retained in input order. + /// + /// # Errors + /// + /// Returns an error if an item is not an HTTP method token. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let value = AccessControlAllowMethodsOwned::from_methods(["GET", "POST"])?; + /// assert_eq!( + /// value + /// .iter() + /// .map(|method| method.as_str()) + /// .collect::>(), + /// ["GET", "POST"] + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + // LLVM emits an uncallable polymorphized instance for this adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + pub fn from_methods(methods: I) -> Result + where + I: IntoIterator, + S: AsRef, + { + CorsList::from_items($name, methods, true).map(Self) + } + + /// Constructs the wildcard field value. + /// + /// Whether the wildcard has wildcard semantics depends on the + /// request and is deliberately not inferred here. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let wildcard = AccessControlAllowMethodsOwned::wildcard(); + /// assert!(wildcard.contains_wildcard()); + /// assert!(wildcard.is_wildcard()); + /// + /// let explicit = AccessControlAllowMethodsOwned::from_methods(["GET"])?; + /// assert!(!explicit.contains_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn wildcard() -> Self { + Self(CorsList::wildcard()) + } + + /// Constructs a present field with an empty method list. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let value = AccessControlAllowMethodsOwned::empty(); + /// assert!(value.is_empty()); + /// assert_eq!(value.len(), 0); + /// ``` + pub fn empty() -> Self { + Self(CorsList::empty()) + } + + /// Validates and adopts complete field lines. + /// + /// # Errors + /// + /// Returns an error for no field lines or malformed list members. + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let value = AccessControlAllowMethodsOwned::from_field_values(vec![ + /// FieldValue::from_static("GET, POST"), + /// FieldValue::from_static("PATCH"), + /// ])?; + /// assert_eq!(value.len(), 3); + /// assert_eq!(value.field_values().count(), 2); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_field_values(values: Vec) -> Result { + CorsList::from_field_values($name, values, true).map(Self) + } + + /// Iterates methods in wire order without allocating. + /// + /// Duplicate methods are returned separately. + /// [`MethodView`] preserves spelling and compares case-sensitively. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let value = AccessControlAllowMethodsOwned::from_methods(["GET", "POST"])?; + /// let methods = value + /// .iter() + /// .map(|method| method.as_str()) + /// .collect::>(); + /// assert_eq!(methods, ["GET", "POST"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn iter(&self) -> super::CorsMethods<'_> { + super::CorsMethods::new(self.0.field_values()) + } + + /// Returns the number of list members, including duplicates. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let value = AccessControlAllowMethodsOwned::from_methods(["GET", "POST", "PATCH"])?; + /// assert_eq!(value.len(), 3); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn len(&self) -> usize { + self.iter().count() + } + + /// Returns whether the list has no members. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let empty = AccessControlAllowMethodsOwned::empty(); + /// assert!(empty.is_empty()); + /// + /// let value = AccessControlAllowMethodsOwned::from_methods(["GET"])?; + /// assert!(!value.is_empty()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn is_empty(&self) -> bool { + self.iter().next().is_none() + } + + /// Returns whether any member is `*`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let wildcard = AccessControlAllowMethodsOwned::wildcard(); + /// assert!(wildcard.contains_wildcard()); + /// + /// let explicit = AccessControlAllowMethodsOwned::from_methods(["GET", "POST"])?; + /// assert!(!explicit.contains_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn contains_wildcard(&self) -> bool { + self.iter().any(|method| method.as_bytes() == b"*") + } + + /// Returns whether `*` is the only list member. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let wildcard = AccessControlAllowMethodsOwned::wildcard(); + /// assert!(wildcard.is_wildcard()); + /// + /// let mixed = AccessControlAllowMethodsOwned::from_methods(["*", "GET"])?; + /// assert!(mixed.contains_wildcard()); + /// assert!(!mixed.is_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn is_wildcard(&self) -> bool { + let mut methods = self.iter(); + methods.next().is_some_and(|method| method.as_bytes() == b"*") && methods.next().is_none() + } + + /// Iterates original field lines in wire order. + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let value = AccessControlAllowMethodsOwned::try_from(FieldValue::from_static("GET, POST"))?; + /// let mut fields = value.field_values(); + /// assert_eq!( + /// fields.next().map(|field| field.as_bytes()), + /// Some(b"GET, POST".as_slice()) + /// ); + /// assert!(fields.next().is_none()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn field_values(&self) -> impl Iterator> { + self.0.field_values().map(FieldValue::as_field_value_ref) + } + + /// Returns the original field lines. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// use http_headers::headers::AccessControlAllowMethodsOwned; + /// + /// let value = AccessControlAllowMethodsOwned::try_from(FieldValue::from_static("GET, POST"))?; + /// let fields = value.into_field_values(); + /// assert_eq!(fields.len(), 1); + /// assert_eq!(fields[0].as_bytes(), b"GET, POST"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_values(self) -> Vec { + self.0.into_field_values() + } + } + + impl<'a> IntoIterator for &'a $owned { + type Item = MethodView<'a>; + type IntoIter = super::CorsMethods<'a>; + + fn into_iter(self) -> Self::IntoIter { + self.iter() + } + } + + impl<'a> $borrowed<'a> { + /// Iterates methods in wire order without allocating. + /// + /// Duplicate methods are returned separately. + /// [`MethodView`] preserves spelling and compares case-sensitively. + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowMethods; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("GET, POST"), + /// ); + /// let value = AccessControlAllowMethods::view(&headers)?.expect("present"); + /// let methods = value + /// .iter() + /// .map(|method| method.as_str()) + /// .collect::>(); + /// assert_eq!(methods, ["GET", "POST"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn iter(&self) -> impl Iterator> + '_ { + self.0 + .values + .repeated() + .flat_map(|value| value.as_bytes().split(|byte| *byte == b',')) + .map(validate::trim_ows) + .filter(|item| !item.is_empty()) + .map(method_ref_validated) + } + + /// Returns the number of list members, including duplicates. + #[must_use] + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowMethods; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("GET, POST"), + /// ); + /// headers.append( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("PATCH"), + /// ); + /// let value = AccessControlAllowMethods::view(&headers)?.expect("present"); + /// assert_eq!(value.len(), 3); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn len(&self) -> usize { + self.iter().count() + } + + /// Returns whether the list has no members. + #[must_use] + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowMethods; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static(""), + /// ); + /// let value = AccessControlAllowMethods::view(&headers)?.expect("present"); + /// assert!(value.is_empty()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn is_empty(&self) -> bool { + self.iter().next().is_none() + } + + /// Returns whether any member is `*`. + #[must_use] + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowMethods; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("*, GET"), + /// ); + /// let value = AccessControlAllowMethods::view(&headers)?.expect("present"); + /// assert!(value.contains_wildcard()); + /// + /// let mut explicit_headers = HeaderMap::new(); + /// explicit_headers.insert( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("GET"), + /// ); + /// let explicit = AccessControlAllowMethods::view(&explicit_headers)?.expect("present"); + /// assert!(!explicit.contains_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn contains_wildcard(&self) -> bool { + self.iter().any(|method| method.as_bytes() == b"*") + } + + /// Returns whether `*` is the only list member. + #[must_use] + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowMethods; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("*"), + /// ); + /// let wildcard = AccessControlAllowMethods::view(&headers)?.expect("present"); + /// assert!(wildcard.is_wildcard()); + /// + /// let mut mixed_headers = HeaderMap::new(); + /// mixed_headers.insert( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("*, GET"), + /// ); + /// let mixed = AccessControlAllowMethods::view(&mixed_headers)?.expect("present"); + /// assert!(mixed.contains_wildcard()); + /// assert!(!mixed.is_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn is_wildcard(&self) -> bool { + let mut methods = self.iter(); + methods.next().is_some_and(|method| method.as_bytes() == b"*") && methods.next().is_none() + } + + /// Iterates original field lines in wire order. + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowMethods; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("GET, POST"), + /// ); + /// headers.append( + /// "access-control-allow-methods", + /// http::HeaderValue::from_static("PATCH"), + /// ); + /// let value = AccessControlAllowMethods::view(&headers)?.expect("present"); + /// let fields = value + /// .field_values() + /// .map(|field| field.as_bytes()) + /// .collect::>(); + /// assert_eq!(fields, [b"GET, POST".as_slice(), b"PATCH".as_slice()]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn field_values(&self) -> impl Iterator> + '_ { + self.0.values.repeated() + } + } + + impl Field for $descriptor { + type View<'a> = $borrowed<'a>; + type Owned = $owned; + + fn name() -> &'static FieldName { + $name + } + + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_list_item_limit(b',', true)?; + let mut repeated = lines.repeated(); + let first = repeated.next().expect("FieldLines always contains at least one field line"); + if repeated.next().is_none() { + validate_single_list($name, first, true)?; + return Ok(Some($borrowed(CorsListView { values: lines }))); + } + validate_list($name, lines.repeated(), true)?; + Ok(Some($borrowed(CorsListView { values: lines }))) + } + + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + CorsList::decode($name, &lines, true).map($owned).map(Some) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), value.0.into_encoded()) + } + } + + impl TryFrom for $owned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + CorsList::from_field_value($name, value, true).map(Self) + } + } + + impl TryFrom> for $owned { + type Error = DecodeError; + + fn try_from(values: Vec) -> Result { + Self::from_field_values(values) + } + } + }; +} + +define_method_list!( + AccessControlAllowMethods, + AccessControlAllowMethodsOwned, + AccessControlAllowMethodsView, + &FieldName::AccessControlAllowMethods +); + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{AccessControlAllowMethods, AccessControlAllowMethodsOwned}; + use crate::headers::cors::test_map::TestMap; + use crate::{DecodeErrorKind, FieldName, FieldValue}; + + #[test] + fn constructors_and_accessors_cover_empty_wildcard_and_repeated_lists() { + let methods = AccessControlAllowMethodsOwned::from_methods(["GET", "CUSTOM", "GET"]).expect("valid method list"); + assert_eq!( + methods.iter().map(super::MethodView::as_str).collect::>(), + ["GET", "CUSTOM", "GET"] + ); + assert_eq!(methods.len(), 3); + assert!(!methods.is_empty()); + assert!(!methods.contains_wildcard()); + assert!(!methods.is_wildcard()); + assert_eq!(methods.field_values().count(), 1); + assert!(format!("{methods:?}").contains("method_count")); + + let wildcard = AccessControlAllowMethodsOwned::wildcard(); + assert!(wildcard.contains_wildcard()); + assert!(wildcard.is_wildcard()); + let mixed = AccessControlAllowMethodsOwned::from_methods(["*", "GET"]).expect("mixed wildcard list"); + assert!(mixed.contains_wildcard()); + assert!(!mixed.is_wildcard()); + assert!(AccessControlAllowMethodsOwned::empty().is_empty()); + + let repeated = + AccessControlAllowMethodsOwned::from_field_values(vec![FieldValue::from_static("GET, POST"), FieldValue::from_static("PATCH")]) + .expect("repeated list"); + assert_eq!(repeated.len(), 3); + assert_eq!(repeated.field_values().count(), 2); + assert_eq!(repeated.into_field_values().len(), 2); + + let one = AccessControlAllowMethodsOwned::try_from(FieldValue::from_static("GET")).expect("one field line"); + assert_eq!(one.len(), 1); + let many = AccessControlAllowMethodsOwned::try_from(vec![FieldValue::from_static("GET"), FieldValue::from_static("POST")]) + .expect("many field lines"); + assert_eq!(many.len(), 2); + } + + #[test] + fn borrowed_owned_insert_and_error_paths_preserve_field_lines() { + let source = TestMap::new( + &FieldName::AccessControlAllowMethods, + vec![FieldValue::from_static("GET, CUSTOM"), FieldValue::from_static("*, POST")], + ); + let view = AccessControlAllowMethods::view(&source) + .expect("valid borrowed list") + .expect("present"); + assert_eq!(view.len(), 4); + assert!(!view.is_empty()); + assert!(view.contains_wildcard()); + assert!(!view.is_wildcard()); + assert_eq!(view.field_values().count(), 2); + assert!(format!("{view:?}").contains("method_count")); + + let wildcard_source = TestMap::new(&FieldName::AccessControlAllowMethods, vec![FieldValue::from_static("*")]); + let wildcard_view = AccessControlAllowMethods::view(&wildcard_source) + .expect("wildcard view") + .expect("present"); + assert!(wildcard_view.is_wildcard()); + + let owned = AccessControlAllowMethods::owned(&source) + .expect("valid owned list") + .expect("present"); + assert_eq!(owned.len(), 4); + let mut sink = TestMap::new(&FieldName::Accept, Vec::new()); + AccessControlAllowMethods::insert(&mut sink, owned).expect("insert method list"); + assert_eq!(sink.name, &FieldName::AccessControlAllowMethods); + assert_eq!(sink.values, source.values); + + let absent = TestMap::new(&FieldName::Accept, Vec::new()); + assert!(AccessControlAllowMethods::view(&absent).expect("absent").is_none()); + assert!(AccessControlAllowMethods::owned(&absent).expect("absent").is_none()); + + let no_lines = AccessControlAllowMethodsOwned::from_field_values(Vec::new()).expect_err("no field lines"); + assert_eq!(no_lines.kind(), DecodeErrorKind::MissingValue); + let invalid = AccessControlAllowMethodsOwned::from_methods(["GET", "bad method"]).expect_err("invalid method"); + assert_eq!(invalid.kind(), DecodeErrorKind::InvalidToken); + + let bad_later = TestMap::new( + &FieldName::AccessControlAllowMethods, + vec![FieldValue::from_static("GET"), FieldValue::from_static("bad method")], + ); + let error = AccessControlAllowMethods::view(&bad_later).expect_err("bad second field line"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidToken); + assert_eq!(error.value_index(), Some(1)); + let error = AccessControlAllowMethods::owned(&bad_later).expect_err("bad second field line"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidToken); + assert_eq!(error.value_index(), Some(1)); + } +} diff --git a/crates/http_headers/src/headers/cors/access_control_allow_origin.rs b/crates/http_headers/src/headers/cors/access_control_allow_origin.rs new file mode 100644 index 000000000..60162b42b --- /dev/null +++ b/crates/http_headers/src/headers/cors/access_control_allow_origin.rs @@ -0,0 +1,1364 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::fmt::Write as _; +use std::net::{Ipv4Addr, Ipv6Addr}; +use std::ops::Range; +use std::str::{self, FromStr as _}; + +use super::shared::{ + BYTE_CLASS, CLASS_AUTHORITY, CLASS_DIGIT, CLASS_DOMAIN, CLASS_HEX, CLASS_LABEL_EDGE, CLASS_SCHEME, invalid_syntax, trimmed_range, +}; +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef}; + +mod components; +pub use components::{AccessControlAllowOriginKind, OriginDomainView, OriginHost, OriginScheme, SerializedOriginView}; + +/// Defines the `Access-Control-Allow-Origin` header. +/// +/// Serialized tuple origins are limited to the `ftp`, `http`, `https`, `ws`, +/// and `wss` schemes. Origins that serialize opaquely are represented by the +/// case-sensitive `null` value rather than a scheme-and-host value. +/// +/// # Specification +/// +/// Defined by the Fetch standard's +/// [CORS protocol and credentials section](https://fetch.spec.whatwg.org/#http-access-control-allow-origin). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{AccessControlAllowOrigin, AccessControlAllowOriginOwned}; +/// +/// let mut map = HeaderMap::new(); +/// AccessControlAllowOrigin::insert(&mut map, AccessControlAllowOriginOwned::wildcard())?; +/// assert!(AccessControlAllowOrigin::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct AccessControlAllowOrigin { + _private: (), +} + +/// Owned value for the `Access-Control-Allow-Origin` header. +/// +/// # Specification +/// +/// Defined by the Fetch standard's [CORS protocol and credentials section]. +/// +/// # Examples +/// +/// ```rust +/// let value = +/// http_headers::headers::AccessControlAllowOriginOwned::try_from("https://example.com")?; +/// assert_eq!(value.origin()?, Some("https://example.com")); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Access-Control-Allow-Origin: *` permits a wildcard origin, +/// `Access-Control-Allow-Origin: null` carries the opaque origin, and +/// `Access-Control-Allow-Origin: https://api.example.com:8443` carries a +/// serialized origin. +/// +/// [CORS protocol and credentials section]: https://fetch.spec.whatwg.org/#http-access-control-allow-origin +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct AccessControlAllowOriginOwned { + value: FieldValue, + parsed: ParsedOrigin, +} + +/// Borrowed value for the `Access-Control-Allow-Origin` header. +#[derive(Clone, Copy, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), http_headers::DecodeError> { +/// use http::{HeaderMap, HeaderValue}; +/// use http_headers::Field; +/// use http_headers::headers::{AccessControlAllowOrigin, AccessControlAllowOriginView}; +/// +/// let mut headers = HeaderMap::new(); +/// headers.insert( +/// "access-control-allow-origin", +/// HeaderValue::from_static("null"), +/// ); +/// let value: AccessControlAllowOriginView<'_> = +/// AccessControlAllowOrigin::view(&headers)?.expect("present"); +/// assert!(value.is_null()); +/// assert_eq!(value.as_str(), "null"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +pub struct AccessControlAllowOriginView<'a> { + value: FieldValueRef<'a>, + serialized: &'a str, + start: usize, + end: usize, + kind: OriginKind, +} + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +struct ParsedOrigin { + start: usize, + end: usize, + kind: OriginKind, +} + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +enum OriginKind { + Wildcard, + Null, + Origin(ParsedTuple), +} + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +struct ParsedTuple { + scheme: OriginScheme, + host_start: usize, + host_end: usize, + host: ParsedOriginHost, + port: Option, +} + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +enum ParsedOriginHost { + Domain, + Ipv4(Ipv4Addr), + Ipv6(Ipv6Addr), +} + +impl OriginKind { + fn project(self, serialized: &str) -> AccessControlAllowOriginKind<'_> { + match self { + Self::Wildcard => AccessControlAllowOriginKind::Wildcard, + Self::Null => AccessControlAllowOriginKind::Null, + Self::Origin(tuple) => AccessControlAllowOriginKind::Origin(SerializedOriginView { + serialized, + scheme: tuple.scheme, + host: match tuple.host { + ParsedOriginHost::Domain => OriginHost::Domain(OriginDomainView { + text: &serialized[tuple.host_start..tuple.host_end], + }), + ParsedOriginHost::Ipv4(address) => OriginHost::Ipv4(address), + ParsedOriginHost::Ipv6(address) => OriginHost::Ipv6(address), + }, + port: tuple.port, + }), + } + } +} + +impl fmt::Debug for AccessControlAllowOriginOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("AccessControlAllowOriginOwned") + .field("kind", &self.parsed.kind) + .finish_non_exhaustive() + } +} + +impl fmt::Debug for AccessControlAllowOriginView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("AccessControlAllowOriginView") + .field("kind", &self.kind) + .finish_non_exhaustive() + } +} + +impl AccessControlAllowOriginOwned { + /// Constructs a canonical serialized origin from validated components. + /// + /// Default ports are omitted. IPv6 uses compressed lowercase hexadecimal, + /// including for IPv4-mapped addresses; no path or trailing slash is added. + /// + /// # Errors + /// + /// Returns an error if an HTTP or HTTPS domain with an explicit non-default + /// port exceeds the existing 253-byte serialized-authority limit. + /// + /// # Examples + /// + /// ``` + /// use std::net::Ipv6Addr; + /// + /// use http_headers::headers::{ + /// AccessControlAllowOriginKind, AccessControlAllowOriginOwned, OriginHost, OriginScheme, + /// }; + /// + /// let value = AccessControlAllowOriginOwned::from_parts( + /// OriginScheme::Https, + /// OriginHost::Ipv6(Ipv6Addr::LOCALHOST), + /// Some(443), + /// )?; + /// assert_eq!(value.as_str()?, "https://[::1]"); + /// if let AccessControlAllowOriginKind::Origin(origin) = value.kind() { + /// assert_eq!(origin.port(), None); + /// assert_eq!(origin.effective_port(), 443); + /// } + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::missing_panics_doc, + reason = "validated components and String formatting cannot violate field-value invariants" + )] + pub fn from_parts(scheme: OriginScheme, host: OriginHost<'_>, port: Option) -> Result { + let host_capacity = match host { + OriginHost::Domain(domain) => domain.as_str().len(), + OriginHost::Ipv4(_) => 15, + OriginHost::Ipv6(_) => 41, + }; + let mut serialized = String::with_capacity(host_capacity.saturating_add(scheme.as_str().len() + 9)); + serialized.push_str(scheme.as_str()); + serialized.push_str("://"); + let host_start = serialized.len(); + let host = match host { + OriginHost::Domain(domain) => { + serialized.push_str(domain.as_str()); + ParsedOriginHost::Domain + } + OriginHost::Ipv4(address) => { + write!(serialized, "{address}").expect("writing an IPv4 address into a String cannot fail"); + ParsedOriginHost::Ipv4(address) + } + OriginHost::Ipv6(address) => { + let mut buffer = [0; 39]; + serialized.push('['); + serialized.push_str(serialize_ipv6(address, &mut buffer)); + serialized.push(']'); + ParsedOriginHost::Ipv6(address) + } + }; + let host_end = serialized.len(); + let port = port.filter(|port| *port != scheme.default_port()); + if let Some(port) = port { + write!(serialized, ":{port}").expect("writing a network port into a String cannot fail"); + } + let end = serialized.len(); + if matches!(scheme, OriginScheme::Http | OriginScheme::Https) && port.is_some() && end - host_start > 253 { + return Err(invalid_syntax(&FieldName::AccessControlAllowOrigin)); + } + Ok(Self { + value: FieldValue::try_from(serialized).expect("serialized validated origin components are valid field value bytes"), + parsed: ParsedOrigin { + start: 0, + end, + kind: OriginKind::Origin(ParsedTuple { + scheme, + host_start, + host_end, + host, + port, + }), + }, + }) + } + + /// Returns wildcard, null, or retained serialized-origin components. + #[must_use] + #[inline] + #[expect(clippy::missing_panics_doc, reason = "the offsets are private and validated at construction")] + pub fn kind(&self) -> AccessControlAllowOriginKind<'_> { + self.parsed.kind.project( + self.as_str() + .expect("semantic offsets were validated when constructing this origin"), + ) + } + + /// Borrows the original bytes and retained components without parsing again. + #[must_use] + #[inline] + #[expect(clippy::missing_panics_doc, reason = "the offsets are private and validated at construction")] + pub fn as_view(&self) -> AccessControlAllowOriginView<'_> { + AccessControlAllowOriginView { + value: self.value.as_field_value_ref(), + serialized: self + .as_str() + .expect("semantic offsets were validated when constructing this origin"), + start: self.parsed.start, + end: self.parsed.end, + kind: self.parsed.kind, + } + } + + /// Returns the original field value, including surrounding whitespace. + #[must_use] + #[inline] + pub const fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Constructs the wildcard value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowOriginOwned; + /// + /// let value = AccessControlAllowOriginOwned::wildcard(); + /// assert!(value.is_wildcard()); + /// assert_eq!(value.as_str()?, "*"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn wildcard() -> Self { + Self { + value: FieldValue::from_static("*"), + parsed: ParsedOrigin { + start: 0, + end: 1, + kind: OriginKind::Wildcard, + }, + } + } + + /// Constructs the case-sensitive `null` origin value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowOriginOwned; + /// + /// let value = AccessControlAllowOriginOwned::null(); + /// assert!(value.is_null()); + /// assert_eq!(value.as_str()?, "null"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn null() -> Self { + Self { + value: FieldValue::from_static("null"), + parsed: ParsedOrigin { + start: 0, + end: 4, + kind: OriginKind::Null, + }, + } + } + + /// Constructs a serialized origin. + /// + /// This accepts tuple origins with the `ftp`, `http`, `https`, `ws`, and + /// `wss` schemes. Origins that serialize opaquely must use [`Self::null`]. + /// Paths, queries, fragments, user information, uppercase schemes or + /// domains, and non-serialized IP addresses are rejected. + /// + /// # Errors + /// + /// Returns an error if `origin` is not a serialized origin. + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowOriginOwned; + /// + /// let value = AccessControlAllowOriginOwned::from_origin("https://example.com")?; + /// assert_eq!(value.origin()?, Some("https://example.com")); + /// assert!(!value.is_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_origin(origin: impl AsRef) -> Result { + let parsed = Self::try_from(origin.as_ref())?; + if matches!(parsed.parsed.kind, OriginKind::Origin(_)) { + Ok(parsed) + } else { + Err(invalid_syntax(&FieldName::AccessControlAllowOrigin)) + } + } + + #[cfg(all(feature = "serde", feature = "headers-cors"))] + pub(crate) const fn field_value(&self) -> &FieldValue { + &self.value + } + + /// Returns whether this is the wildcard value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowOriginOwned; + /// + /// let wildcard = AccessControlAllowOriginOwned::wildcard(); + /// let origin = AccessControlAllowOriginOwned::from_origin("https://example.com")?; + /// assert!(wildcard.is_wildcard()); + /// assert!(!origin.is_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn is_wildcard(&self) -> bool { + matches!(self.parsed.kind, OriginKind::Wildcard) + } + + /// Returns whether this is the `null` origin value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowOriginOwned; + /// + /// let null = AccessControlAllowOriginOwned::null(); + /// let wildcard = AccessControlAllowOriginOwned::wildcard(); + /// assert!(null.is_null()); + /// assert!(!wildcard.is_null()); + /// ``` + pub const fn is_null(&self) -> bool { + matches!(self.parsed.kind, OriginKind::Null) + } + + /// Returns a serialized origin, excluding wildcard and `null`. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the field value. + /// # Examples + /// + /// ```rust + /// let value = + /// http_headers::headers::AccessControlAllowOriginOwned::try_from("https://example.com")?; + /// assert_eq!(value.origin()?, Some("https://example.com")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn origin(&self) -> Result, DecodeError> { + if !matches!(self.parsed.kind, OriginKind::Origin(_)) { + return Ok(None); + } + semantic_str( + &FieldName::AccessControlAllowOrigin, + self.value.as_bytes(), + self.parsed.start..self.parsed.end, + ) + .map(Some) + } + + /// Returns the semantic value after surrounding optional whitespace. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the field value. + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowOriginOwned; + /// + /// let null = AccessControlAllowOriginOwned::try_from(" null ")?; + /// let origin = AccessControlAllowOriginOwned::from_origin("https://example.com")?; + /// assert_eq!(null.as_str()?, "null"); + /// assert_eq!(origin.as_str()?, "https://example.com"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_str(&self) -> Result<&str, DecodeError> { + semantic_str( + &FieldName::AccessControlAllowOrigin, + self.value.as_bytes(), + self.parsed.start..self.parsed.end, + ) + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowOriginOwned; + /// + /// let value = AccessControlAllowOriginOwned::from_origin("https://example.com")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"https://example.com"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(AccessControlAllowOriginOwned, |value| value.value); + +impl fmt::Display for AccessControlAllowOriginOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str().map_err(|_invalid| fmt::Error)?) + } +} + +impl<'a> AccessControlAllowOriginView<'a> { + /// Returns wildcard, null, or retained serialized-origin components. + #[must_use] + #[inline] + pub fn kind(self) -> AccessControlAllowOriginKind<'a> { + self.kind.project(self.serialized) + } + + /// Returns whether this is the wildcard value. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowOrigin; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert("access-control-allow-origin", HeaderValue::from_static("*")); + /// let value = AccessControlAllowOrigin::view(&headers)?.expect("present"); + /// assert!(value.is_wildcard()); + /// assert_eq!(value.origin(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub const fn is_wildcard(self) -> bool { + matches!(self.kind, OriginKind::Wildcard) + } + + /// Returns whether this is the `null` origin value. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowOrigin; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-origin", + /// HeaderValue::from_static("null"), + /// ); + /// let value = AccessControlAllowOrigin::view(&headers)?.expect("present"); + /// assert!(value.is_null()); + /// assert_eq!(value.origin(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub const fn is_null(self) -> bool { + matches!(self.kind, OriginKind::Null) + } + + /// Returns a serialized origin, excluding wildcard and `null`. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = + /// http_headers::headers::AccessControlAllowOriginOwned::try_from("https://example.com")?; + /// assert_eq!(value.origin()?, Some("https://example.com")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn origin(self) -> Option<&'a str> { + if matches!(self.kind, OriginKind::Origin(_)) { + Some(self.serialized) + } else { + None + } + } + + /// Returns the semantic value after surrounding optional whitespace. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowOrigin; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-origin", + /// HeaderValue::from_static(" https://example.com "), + /// ); + /// let value = AccessControlAllowOrigin::view(&headers)?.expect("present"); + /// assert_eq!(value.as_str(), "https://example.com"); + /// assert_eq!(value.origin(), Some("https://example.com")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub const fn as_str(self) -> &'a str { + self.serialized + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), Box> { + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowOrigin; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-origin", + /// HeaderValue::from_static(" https://example.com "), + /// ); + /// let value = AccessControlAllowOrigin::view(&headers)?.expect("present"); + /// assert_eq!(value.as_field_value().to_str()?, " https://example.com "); + /// # Ok::<(), Box>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } +} + +impl Field for AccessControlAllowOrigin { + type View<'a> = AccessControlAllowOriginView<'a>; + type Owned = AccessControlAllowOriginOwned; + + fn name() -> &'static FieldName { + &FieldName::AccessControlAllowOrigin + } + + #[inline] + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_custom_source()?; + let value = lines.exactly_one()?; + let (parsed, serialized) = parse_allow_origin_value(value)?; + Ok(Some(AccessControlAllowOriginView { + value, + serialized, + start: parsed.start, + end: parsed.end, + kind: parsed.kind, + })) + } + + /// Parses straight into owned storage, since the serialized origin a view + /// borrows is not kept. + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + let owned = lines.exactly_one_owned()?; + let parsed = parse_allow_origin_value(owned.as_field_value_ref())?.0; + Ok(Some(AccessControlAllowOriginOwned { value: owned, parsed })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::single(value.value)) + } +} + +super::super::shared::impl_string_conversions!( + AccessControlAllowOriginOwned, + &FieldName::AccessControlAllowOrigin, + invalid_syntax, + value +); + +impl TryFrom for AccessControlAllowOriginOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + let parsed = parse_allow_origin_value(value.as_field_value_ref())?.0; + Ok(Self { value, parsed }) + } +} +#[inline] +fn parse_allow_origin_value(value: FieldValueRef<'_>) -> Result<(ParsedOrigin, &str), DecodeError> { + let range = trimmed_range(value.as_bytes()); + let serialized = visible_ascii_str(value, range.clone()).map_or_else( + || semantic_str(&FieldName::AccessControlAllowOrigin, value.as_bytes(), range.clone()), + Ok, + )?; + let kind = match serialized.as_bytes() { + b"*" => OriginKind::Wildcard, + b"null" => OriginKind::Null, + origin => OriginKind::Origin(parse_serialized_origin(origin).ok_or_else(|| invalid_syntax(&FieldName::AccessControlAllowOrigin))?), + }; + Ok(( + ParsedOrigin { + start: range.start, + end: range.end, + kind, + }, + serialized, + )) +} + +fn parse_serialized_origin(origin: &[u8]) -> Option { + if let Some(valid) = common_http_origin(origin) { + return valid.then(|| { + let scheme = if origin.starts_with(b"https:") { + OriginScheme::Https + } else { + OriginScheme::Http + }; + ParsedTuple { + scheme, + host_start: scheme.as_str().len() + 3, + host_end: origin.len(), + host: ParsedOriginHost::Domain, + port: None, + } + }); + } + let (scheme, authority) = split_scheme(origin)?; + let typed_scheme = OriginScheme::parse(scheme)?; + let (host, host_len, port) = parse_serialized_authority(authority, scheme)?; + let host_start = scheme.len() + 3; + Some(ParsedTuple { + scheme: typed_scheme, + host_start, + host_end: host_start + host_len, + host, + port, + }) +} + +fn common_http_origin(origin: &[u8]) -> Option { + let host = origin.strip_prefix(b"https://").or_else(|| origin.strip_prefix(b"http://"))?; + let host = host.strip_suffix(b".").unwrap_or(host); + if host.is_empty() || host.len() > 253 { + return Some(false); + } + + let mut label_start = 0_usize; + let mut saw_non_digit = false; + for (index, &byte) in host.iter().enumerate() { + let class = BYTE_CLASS[usize::from(byte)]; + if class & CLASS_DOMAIN != 0 { + saw_non_digit |= class & CLASS_DIGIT == 0; + continue; + } + if byte == b'.' { + if !valid_domain_label(&host[label_start..index]) { + return Some(false); + } + label_start = index + 1; + continue; + } + // A port or an IP-literal needs the general authority rules; any other + // byte is outside `reg-name`, which those rules also reject. + return (byte != b':' && byte != b'[').then_some(false); + } + if !saw_non_digit { + return None; + } + Some(valid_domain_label(&host[label_start..])) +} + +#[inline] +fn valid_domain_label(label: &[u8]) -> bool { + let (Some(&first), Some(&last)) = (label.first(), label.last()) else { + return false; + }; + label.len() <= 63 && BYTE_CLASS[usize::from(first)] & CLASS_LABEL_EDGE != 0 && BYTE_CLASS[usize::from(last)] & CLASS_LABEL_EDGE != 0 +} + +/// Splits a serialized origin at its first `://` and validates the scheme. +fn split_scheme(origin: &[u8]) -> Option<(&[u8], &[u8])> { + let mut index = 0_usize; + while let Some(&byte) = origin.get(index) { + if byte == b':' { + break; + } + if BYTE_CLASS[usize::from(byte)] & CLASS_SCHEME == 0 { + return None; + } + index += 1; + } + let rest = index.checked_add(3)?; + if !origin.first().is_some_and(u8::is_ascii_lowercase) || origin.get(index..rest) != Some(b"://".as_slice()) { + return None; + } + Some((origin.get(..index)?, origin.get(rest..)?)) +} + +fn parse_serialized_authority(authority: &[u8], scheme: &[u8]) -> Option<(ParsedOriginHost, usize, Option)> { + let (&first, rest) = authority.split_first()?; + + if first == b'[' { + if rest + .iter() + .any(|byte| BYTE_CLASS[usize::from(*byte)] & CLASS_AUTHORITY == 0 && *byte != b':') + { + return None; + } + let close = rest.iter().position(|byte| *byte == b']')?; + let (host, suffix) = rest.split_at(close); + let address = parse_serialized_ipv6(host)?; + let port = parse_serialized_port_suffix(suffix.get(1..)?, scheme)?; + return Some((ParsedOriginHost::Ipv6(address), close + 2, port)); + } + + let mut colon = None; + for (index, &byte) in authority.iter().enumerate() { + if BYTE_CLASS[usize::from(byte)] & CLASS_AUTHORITY != 0 { + continue; + } + if byte != b':' || colon.is_some() { + return None; + } + colon = Some(index); + } + + let (host, port) = match colon { + Some(index) => (authority.get(..index).unwrap_or_default(), authority.get(index.saturating_add(1)..)), + None => (authority, None), + }; + let port = match port { + Some(port) => Some(parse_serialized_port(port, scheme)?), + None => None, + }; + Some((parse_serialized_host(host)?, host.len(), port)) +} + +#[expect(clippy::option_option, reason = "outer None is invalid syntax; inner None is a valid absent port")] +fn parse_serialized_port_suffix(suffix: &[u8], scheme: &[u8]) -> Option> { + match suffix.split_first() { + None => Some(None), + Some((&b':', port)) => parse_serialized_port(port, scheme).map(Some), + Some(_) => None, + } +} + +fn parse_serialized_port(port: &[u8], scheme: &[u8]) -> Option { + if port.is_empty() || port.len() > 5 || (port.len() > 1 && port.first() == Some(&b'0')) { + return None; + } + let mut value = 0_u32; + for &byte in port { + if BYTE_CLASS[usize::from(byte)] & CLASS_DIGIT == 0 { + return None; + } + value = value * 10 + u32::from(byte.wrapping_sub(b'0')); + } + let value = u16::try_from(value).ok()?; + if OriginScheme::parse(scheme).is_some_and(|scheme| value == scheme.default_port()) { + return None; + } + Some(value) +} + +fn parse_serialized_host(host: &[u8]) -> Option { + let host = match host.split_last() { + Some((&b'.', head)) => head, + _ => host, + }; + if host.is_empty() || host.len() > 253 { + return None; + } + let mut label_count = 0_usize; + let mut all_numeric = true; + let mut valid_ipv4 = true; + let mut valid_domain = true; + let mut octets = [0; 4]; + let mut start = 0_usize; + let mut common = u8::MAX; + for (index, &byte) in host.iter().enumerate() { + if byte != b'.' { + common &= BYTE_CLASS[usize::from(byte)]; + continue; + } + let label = host.get(start..index).unwrap_or_default(); + let (numeric, ipv4, domain) = parsed_label_classes(label, common); + all_numeric &= numeric; + valid_ipv4 &= ipv4.is_some(); + if let Some(octet) = octets.get_mut(label_count) { + *octet = ipv4.unwrap_or_default(); + } + valid_domain &= domain; + label_count = label_count.saturating_add(1); + start = index.saturating_add(1); + common = u8::MAX; + } + let label = host.get(start..).unwrap_or_default(); + let (numeric, ipv4, domain) = parsed_label_classes(label, common); + all_numeric &= numeric; + valid_ipv4 &= ipv4.is_some(); + if let Some(octet) = octets.get_mut(label_count) { + *octet = ipv4.unwrap_or_default(); + } + valid_domain &= domain; + label_count = label_count.saturating_add(1); + + if all_numeric { + (label_count == 4 && valid_ipv4).then_some(ParsedOriginHost::Ipv4(Ipv4Addr::from(octets))) + } else { + valid_domain.then_some(ParsedOriginHost::Domain) + } +} + +/// Classifies one host label as numeric, a valid IPv4 octet, and a valid domain label. +/// +/// `common` is the intersection of the byte classes of every byte in `label`. +fn parsed_label_classes(label: &[u8], common: u8) -> (bool, Option, bool) { + let (Some(&first), Some(&last)) = (label.first(), label.last()) else { + return (false, None, false); + }; + + let numeric = common & CLASS_DIGIT != 0; + let ipv4 = if numeric && label.len() <= 3 && (label.len() == 1 || first != b'0') { + let mut value = 0_u32; + for &byte in label { + value = value * 10 + u32::from(byte.wrapping_sub(b'0')); + } + u8::try_from(value).ok() + } else { + None + }; + let domain = common & CLASS_DOMAIN != 0 + && label.len() <= 63 + && BYTE_CLASS[usize::from(first)] & CLASS_LABEL_EDGE != 0 + && BYTE_CLASS[usize::from(last)] & CLASS_LABEL_EDGE != 0; + + (numeric, ipv4, domain) +} + +#[inline(never)] // Keep IPv6 parsing and formatting off the domain-origin stack. +fn parse_serialized_ipv6(address: &[u8]) -> Option { + // Canonical hexadecimal IPv6 has at most eight four-digit groups and seven colons. + if address.is_empty() + || address.len() > 39 + || address + .iter() + .any(|byte| BYTE_CLASS[usize::from(*byte)] & CLASS_HEX == 0 && *byte != b':') + { + return None; + } + let text = http_headers_simd::ascii_str(address).expect("serialized IPv6 bytes are ASCII"); + let parsed = Ipv6Addr::from_str(text).ok()?; + let mut buffer = [0; 39]; + (serialize_ipv6(parsed, &mut buffer).as_bytes() == address).then_some(parsed) +} + +fn serialize_ipv6(address: Ipv6Addr, serialized: &mut [u8; 39]) -> &str { + let segments = address.segments(); + let mut longest_start = None; + let mut longest_len = 1_usize; + let mut index = 0_usize; + while index < segments.len() { + if segments[index] != 0 { + index += 1; + continue; + } + let start = index; + while index < segments.len() && segments[index] == 0 { + index += 1; + } + let len = index - start; + if len > longest_len { + longest_start = Some(start); + longest_len = len; + } + } + + let mut written = 0_usize; + let mut index = 0_usize; + while index < segments.len() { + if longest_start == Some(index) { + serialized[written..written + 2].copy_from_slice(b"::"); + written += 2; + index += longest_len; + continue; + } + if written != 0 && serialized[written - 1] != b':' { + serialized[written] = b':'; + written += 1; + } + written += write_hex_segment(segments[index], &mut serialized[written..]); + index += 1; + } + http_headers_simd::ascii_str(&serialized[..written]).expect("IPv6 serialization writes only ASCII hexadecimal digits and colons") +} + +fn write_hex_segment(segment: u16, output: &mut [u8]) -> usize { + let digits = b"0123456789abcdef"; + let leading_zeros = segment.leading_zeros() as usize / 4; + let count = 4_usize.saturating_sub(leading_zeros).max(1); + for (index, slot) in output[..count].iter_mut().enumerate() { + let shift = (count - index - 1) * 4; + *slot = digits[usize::from((segment >> shift) & 0xf)]; + } + count +} + +#[cfg(test)] +fn valid_h16_sequence(sequence: &[u8]) -> Option { + if sequence.is_empty() { + return Some(0); + } + let mut count = 0_usize; + for block in sequence.split(|byte| *byte == b':') { + if block.is_empty() + || block.len() > 4 + || (block.len() > 1 && block.first() == Some(&b'0')) + || !block.iter().all(|byte| BYTE_CLASS[usize::from(*byte)] & CLASS_HEX != 0) + { + return None; + } + count = count.checked_add(1)?; + } + Some(count) +} + +/// Borrows `range` of a field value made only of visible ASCII. +/// +/// A serialized origin never leaves that alphabet, so the branch-free ASCII +/// reduction settles the conversion far more cheaply than the general UTF-8 +/// validator, which is left to handle everything else. +fn visible_ascii_str(value: FieldValueRef<'_>, range: Range) -> Option<&str> { + http_headers_simd::ascii_str(value.as_bytes().get(range)?) +} + +fn semantic_str<'a>(name: &'static FieldName, bytes: &'a [u8], range: Range) -> Result<&'a str, DecodeError> { + let bytes = bytes.get(range).ok_or_else(|| invalid_syntax(name))?; + str::from_utf8(bytes).map_err(|_invalid| DecodeError::new(name, DecodeErrorKind::InvalidUtf8)) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{ + AccessControlAllowOrigin, AccessControlAllowOriginOwned, OriginKind, ParsedOrigin, common_http_origin, semantic_str, split_scheme, + valid_domain_label, valid_h16_sequence, visible_ascii_str, + }; + use crate::headers::cors::test_map::TestMap; + use crate::{DecodeErrorKind, FieldName, FieldValue}; + + fn label_classes(label: &[u8], common: u8) -> (bool, bool, bool) { + let (numeric, octet, domain) = super::parsed_label_classes(label, common); + (numeric, octet.is_some(), domain) + } + + fn valid_serialized_origin(origin: &[u8]) -> bool { + super::parse_serialized_origin(origin).is_some() + } + + fn valid_serialized_authority(authority: &[u8], scheme: &[u8]) -> bool { + super::parse_serialized_authority(authority, scheme).is_some() + } + + fn valid_serialized_host(host: &[u8]) -> bool { + super::parse_serialized_host(host).is_some() + } + + fn valid_serialized_ipv6(address: &[u8]) -> bool { + super::parse_serialized_ipv6(address).is_some() + } + + fn valid_serialized_port(port: &[u8], scheme: &[u8]) -> bool { + super::parse_serialized_port(port, scheme).is_some() + } + + fn valid_serialized_port_suffix(suffix: &[u8], scheme: &[u8]) -> bool { + super::parse_serialized_port_suffix(suffix, scheme).is_some() + } + + #[test] + fn constructors_accessors_display_and_header_paths_preserve_wire_values() { + let wildcard = AccessControlAllowOriginOwned::wildcard(); + assert!(wildcard.is_wildcard()); + assert!(!wildcard.is_null()); + assert_eq!(wildcard.origin().expect("origin accessor"), None); + assert_eq!(wildcard.as_str().expect("wildcard text"), "*"); + + let null = AccessControlAllowOriginOwned::null(); + assert!(!null.is_wildcard()); + assert!(null.is_null()); + assert_eq!(null.origin().expect("origin accessor"), None); + assert_eq!(null.to_string(), "null"); + + let origin = AccessControlAllowOriginOwned::from_origin("https://example.com").expect("origin"); + assert_eq!(origin.origin().expect("serialized origin"), Some("https://example.com")); + assert_eq!(origin.as_str().expect("origin text"), "https://example.com"); + assert_eq!(origin.to_string(), "https://example.com"); + assert!(format!("{origin:?}").contains("Origin")); + assert_eq!(origin.into_field_value(), "https://example.com"); + + AccessControlAllowOriginOwned::from_origin("*").expect_err("wildcard is not an origin"); + AccessControlAllowOriginOwned::from_origin("null").expect_err("null is not an origin"); + let framed_null = AccessControlAllowOriginOwned::try_from(String::from(" null ")).expect("owned null"); + assert!(framed_null.is_null()); + assert_eq!(framed_null.as_str().expect("null text"), "null"); + assert_eq!( + AccessControlAllowOriginOwned::try_from(FieldValue::from_static(" https://example.com ")) + .expect("field origin") + .origin() + .expect("origin accessor"), + Some("https://example.com") + ); + + let source = TestMap::new( + &FieldName::AccessControlAllowOrigin, + vec![FieldValue::from_static(" https://example.com ")], + ); + let view = AccessControlAllowOrigin::view(&source) + .expect("valid origin view") + .expect("present"); + assert!(!view.is_wildcard()); + assert!(!view.is_null()); + assert_eq!(view.origin(), Some("https://example.com")); + assert_eq!(view.as_str(), "https://example.com"); + assert_eq!(view.as_field_value(), " https://example.com "); + assert!(format!("{view:?}").contains("Origin")); + + let owned = AccessControlAllowOrigin::owned(&source) + .expect("valid owned origin") + .expect("present"); + let mut sink = TestMap::new(&FieldName::Accept, Vec::new()); + AccessControlAllowOrigin::insert(&mut sink, owned).expect("insert origin"); + assert_eq!(sink.name, &FieldName::AccessControlAllowOrigin); + assert_eq!(sink.values, source.values); + + let wildcard_source = TestMap::new(&FieldName::AccessControlAllowOrigin, vec![FieldValue::from_static("*")]); + assert!( + AccessControlAllowOrigin::view(&wildcard_source) + .expect("wildcard view") + .expect("present") + .is_wildcard() + ); + assert_eq!( + AccessControlAllowOrigin::view(&wildcard_source) + .expect("wildcard view") + .expect("present") + .origin(), + None + ); + let null_source = TestMap::new(&FieldName::AccessControlAllowOrigin, vec![FieldValue::from_static("null")]); + assert!( + AccessControlAllowOrigin::view(&null_source) + .expect("null view") + .expect("present") + .is_null() + ); + + let absent = TestMap::new(&FieldName::Accept, Vec::new()); + assert!(AccessControlAllowOrigin::view(&absent).expect("absent").is_none()); + assert!(AccessControlAllowOrigin::owned(&absent).expect("absent").is_none()); + } + + #[test] + fn serialized_origin_validation_accepts_canonical_host_forms_and_ports() { + for origin in [ + "http://example.com", + "https://example.com.", + "https://sub-domain.example", + "https://127.0.0.1", + "https://[2001:db8::1]", + "https://[::1]:8443", + "https://[2001::1:0:0:1:1]", + "https://[2001:db8:0:1:2:3:4:5]", + "https://[::ffff:c000:280]", + "ftp://example.com:22", + "ws://example.com:8080", + "wss://example.com:8443", + ] { + assert!(valid_serialized_origin(origin.as_bytes()), "{origin}"); + assert!(AccessControlAllowOriginOwned::from_origin(origin).is_ok(), "{origin}"); + } + + // The `http`/`https` recognizer shortcuts the general authority rules, + // so every answer it commits to must match what those rules decide. + let alphabet = b"a0-.:[]_A%"; + for first in alphabet { + for second in alphabet { + for third in alphabet { + let mut origin = b"https://".to_vec(); + origin.extend_from_slice(&[*first, *second, *third]); + let Some(recognized) = common_http_origin(&origin) else { + continue; + }; + let general = split_scheme(&origin).is_some_and(|(scheme, authority)| valid_serialized_authority(authority, scheme)); + assert_eq!( + recognized, + general, + "{:?}", + str::from_utf8(&origin).expect("the alphabet is ASCII only") + ); + } + } + } + + assert_eq!(common_http_origin(b"https://example.com"), Some(true)); + assert_eq!(common_http_origin(b"https://127.0.0.1"), None); + assert_eq!(common_http_origin(b"custom://example.com"), None); + assert_eq!( + split_scheme(b"custom+v1://host"), + Some((b"custom+v1".as_slice(), b"host".as_slice())) + ); + assert!(valid_domain_label(b"example")); + assert!(!valid_domain_label(b"")); + assert!(valid_serialized_authority(b"example.com", b"https")); + assert!(valid_serialized_port_suffix(b"", b"https")); + assert!(valid_serialized_port_suffix(b":8443", b"https")); + assert!(valid_serialized_port(b"8443", b"https")); + assert!(valid_serialized_port(b"80", b"custom")); + assert!(valid_serialized_host(b"127.0.0.1")); + assert!(valid_serialized_ipv6(b"2001:db8::1")); + assert!(valid_serialized_ipv6(b"2001:db8:1:2:3:4:5:6")); + assert_eq!(valid_h16_sequence(b"2001:db8"), Some(2)); + assert_eq!(valid_h16_sequence(b""), Some(0)); + + let classes = label_classes(b"255", super::CLASS_DIGIT | super::CLASS_DOMAIN); + assert_eq!(classes, (true, true, true)); + } + + #[test] + fn serialized_origin_validation_rejects_noncanonical_and_malformed_forms() { + for origin in [ + "", + "*", + "null", + "HTTPS://example.com", + "1http://example.com", + "http:/example.com", + "http://", + "http://example.com/path", + "http://example.com?query", + "http://user@example.com", + "http://-example.com", + "http://example-.com", + "http://example..com", + "http://999.1.1.1", + "http://01.2.3.4", + "http://1.2.3", + "http://example.com:80", + "https://example.com:443", + "ftp://example.com:21", + "custom+v1://host:80", + "https://example.com:", + "https://example.com:0001", + "https://example.com:65536", + "https://example.com:port", + "https://example.com:1:2", + "https://[2001:0db8::1]", + "https://[2001:db8:0:0:1:2:3:4]", + "https://[2001:0:0:1::1:1]", + "https://[2001:db8:0::1]", + "https://[2001::db8::1]", + "https://[2001:db8:1:2:3:4:5:6:7]", + "https://[2001:db8::1", + "https://[2001:db8::1]suffix", + ] { + assert!(!valid_serialized_origin(origin.as_bytes()), "{origin}"); + } + + assert_eq!(common_http_origin(b"https://"), Some(false)); + assert_eq!(common_http_origin(b"https://bad_host"), Some(false)); + assert_eq!(split_scheme(b"Nope://host"), None); + assert_eq!(split_scheme(b"http:/host"), None); + assert!(!valid_serialized_authority(b"", b"https")); + assert!(!valid_serialized_authority(b"host/path", b"https")); + assert!(!valid_serialized_authority(b"[::1]/", b"https")); + assert!(!valid_serialized_port_suffix(b"suffix", b"https")); + assert!(!valid_serialized_port(b"", b"https")); + assert!(!valid_serialized_port(b"01", b"https")); + assert!(!valid_serialized_port(b"65536", b"https")); + assert!(!valid_serialized_port(b"x", b"https")); + assert!(!valid_serialized_host(b"")); + assert!(valid_serialized_host(b"example.com.")); + assert!(!valid_serialized_host(b"1.2.3")); + assert!(!valid_serialized_ipv6(b"")); + assert!(!valid_serialized_ipv6(b"xyz")); + assert!(!valid_serialized_ipv6(b"1::2::3")); + assert_eq!(valid_h16_sequence(b"0001"), None); + assert_eq!(valid_h16_sequence(b"12345"), None); + assert_eq!(valid_h16_sequence(b"1::2"), None); + assert!(!valid_serialized_ipv6(b"2001::0db8")); + assert_eq!(label_classes(b"", super::CLASS_DIGIT | super::CLASS_DOMAIN), (false, false, false)); + + let error = AccessControlAllowOriginOwned::try_from(String::from("bad\norigin")).expect_err("invalid field bytes"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!( + AccessControlAllowOriginOwned::try_from("bad\norigin") + .expect_err("invalid borrowed field bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + AccessControlAllowOriginOwned::try_from("not-an-origin") + .expect_err("invalid visible origin") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + let duplicate = TestMap::new( + &FieldName::AccessControlAllowOrigin, + vec![ + FieldValue::from_static("https://one.example"), + FieldValue::from_static("https://two.example"), + ], + ); + assert_eq!( + AccessControlAllowOrigin::view(&duplicate).expect_err("singleton header").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + } + + #[test] + fn semantic_accessors_revalidate_private_metadata() { + let out_of_bounds = AccessControlAllowOriginOwned { + value: FieldValue::from_static("x"), + parsed: ParsedOrigin { + start: 2, + end: 3, + kind: OriginKind::Origin(super::parse_serialized_origin(b"https://example.com").unwrap()), + }, + }; + assert_eq!( + out_of_bounds.origin().expect_err("metadata range").kind(), + DecodeErrorKind::InvalidSyntax + ); + + let invalid_utf8 = FieldValue::from_bytes([0xff]).expect("non-ASCII field value byte is permitted"); + assert_eq!( + super::parse_allow_origin_value(invalid_utf8.as_field_value_ref()) + .expect_err("invalid UTF-8 origin") + .kind(), + DecodeErrorKind::InvalidUtf8 + ); + let invalid_utf8 = AccessControlAllowOriginOwned { + value: invalid_utf8, + parsed: ParsedOrigin { + start: 0, + end: 1, + kind: OriginKind::Origin(super::parse_serialized_origin(b"https://example.com").unwrap()), + }, + }; + assert_eq!( + invalid_utf8.as_str().expect_err("invalid UTF-8 metadata").kind(), + DecodeErrorKind::InvalidUtf8 + ); + + let value = FieldValue::from_static("visible"); + assert_eq!(visible_ascii_str(value.as_field_value_ref(), 0..7), Some("visible")); + assert_eq!(visible_ascii_str(value.as_field_value_ref(), 9..10), None); + assert_eq!( + semantic_str(&FieldName::AccessControlAllowOrigin, b"value", 0..5).expect("semantic text"), + "value" + ); + } +} diff --git a/crates/http_headers/src/headers/cors/access_control_allow_origin/components.rs b/crates/http_headers/src/headers/cors/access_control_allow_origin/components.rs new file mode 100644 index 000000000..adebe0223 --- /dev/null +++ b/crates/http_headers/src/headers/cors/access_control_allow_origin/components.rs @@ -0,0 +1,177 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::net::{Ipv4Addr, Ipv6Addr}; + +/// The semantic value of `Access-Control-Allow-Origin`. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub enum AccessControlAllowOriginKind<'a> { + /// The wildcard token. + Wildcard, + /// The opaque-origin serialization `null`, not an opaque-origin identity. + Null, + /// A validated serialized tuple origin. + Origin(SerializedOriginView<'a>), +} + +/// A scheme accepted in serialized tuple origins. +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub enum OriginScheme { + /// File Transfer Protocol. + Ftp, + /// HTTP. + Http, + /// HTTP over TLS. + Https, + /// WebSocket. + Ws, + /// WebSocket over TLS. + Wss, +} + +impl OriginScheme { + /// Returns the lowercase scheme without `://`. + #[must_use] + #[inline] + pub const fn as_str(self) -> &'static str { + match self { + Self::Ftp => "ftp", + Self::Http => "http", + Self::Https => "https", + Self::Ws => "ws", + Self::Wss => "wss", + } + } + + /// Returns the default network port for this scheme. + #[must_use] + #[inline] + pub const fn default_port(self) -> u16 { + match self { + Self::Ftp => 21, + Self::Http | Self::Ws => 80, + Self::Https | Self::Wss => 443, + } + } + + pub(super) fn parse(scheme: &[u8]) -> Option { + match scheme { + b"ftp" => Some(Self::Ftp), + b"http" => Some(Self::Http), + b"https" => Some(Self::Https), + b"ws" => Some(Self::Ws), + b"wss" => Some(Self::Wss), + _ => None, + } + } +} + +impl fmt::Display for OriginScheme { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } +} + +/// A validated origin host, distinct from the broader URI `Host` grammar. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub enum OriginHost<'a> { + /// A lowercase ASCII domain, possibly with a trailing dot. + Domain(OriginDomainView<'a>), + /// A dotted-decimal IPv4 address. + Ipv4(Ipv4Addr), + /// An IPv6 address. + Ipv6(Ipv6Addr), +} + +/// A validated, serialized ASCII domain in an origin. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct OriginDomainView<'a> { + pub(super) text: &'a str, +} + +impl<'a> OriginDomainView<'a> { + /// Validates a domain for typed origin construction. + /// + /// # Errors + /// + /// Rejects uppercase, Unicode, IP addresses, empty labels, and other text + /// outside the existing serialized-origin domain grammar. + pub fn new(text: &'a str) -> Result { + if super::parse_serialized_host(text.as_bytes()) == Some(super::ParsedOriginHost::Domain) { + Ok(Self { text }) + } else { + Err(super::invalid_syntax(&crate::FieldName::AccessControlAllowOrigin)) + } + } + + /// Returns the validated domain, preserving any trailing dot. + #[must_use] + #[inline] + pub const fn as_str(self) -> &'a str { + self.text + } +} + +impl fmt::Display for OriginDomainView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.text) + } +} + +/// The retained components of a serialized tuple origin. +/// +/// Reading these components does not parse or normalize the original value. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct SerializedOriginView<'a> { + pub(super) serialized: &'a str, + pub(super) scheme: OriginScheme, + pub(super) host: OriginHost<'a>, + pub(super) port: Option, +} + +impl<'a> SerializedOriginView<'a> { + /// Returns the complete serialized origin, excluding surrounding whitespace. + #[must_use] + #[inline] + pub const fn as_str(self) -> &'a str { + self.serialized + } + + /// Returns the validated scheme. + #[must_use] + #[inline] + pub const fn scheme(self) -> OriginScheme { + self.scheme + } + + /// Returns the domain or retained IP address. + #[must_use] + #[inline] + pub const fn host(self) -> OriginHost<'a> { + self.host + } + + /// Returns the explicitly serialized port, excluding an omitted default. + #[must_use] + #[inline] + pub const fn port(self) -> Option { + self.port + } + + /// Returns the explicit port or the scheme's default. + #[must_use] + #[inline] + pub const fn effective_port(self) -> u16 { + match self.port { + Some(port) => port, + None => self.scheme.default_port(), + } + } +} + +impl fmt::Display for SerializedOriginView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.serialized) + } +} diff --git a/crates/http_headers/src/headers/cors/access_control_expose_headers.rs b/crates/http_headers/src/headers/cors/access_control_expose_headers.rs new file mode 100644 index 000000000..0f31ac917 --- /dev/null +++ b/crates/http_headers/src/headers/cors/access_control_expose_headers.rs @@ -0,0 +1,40 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::super::FieldNameView; +use super::shared::{CorsList, CorsListView, define_header_name_list, impl_header_name_wildcard}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef, validate}; + +define_header_name_list!( + AccessControlExposeHeaders, + AccessControlExposeHeadersOwned, + AccessControlExposeHeadersView, + "Access-Control-Expose-Headers", + &FieldName::AccessControlExposeHeaders, + true, + "Defined by the Fetch standard's [CORS protocol and credentials section](https://fetch.spec.whatwg.org/#http-access-control-expose-headers).", + "`Access-Control-Expose-Headers: X-Request-Id, Content-Length` lists exposed fields, `Access-Control-Expose-Headers: *` uses a wildcard, and an empty field value is accepted." +); + +impl_header_name_wildcard!(AccessControlExposeHeadersOwned, AccessControlExposeHeadersView); + +impl AccessControlExposeHeadersOwned { + /// Constructs a present field with an empty field-name list. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlExposeHeadersOwned; + /// + /// let value = AccessControlExposeHeadersOwned::empty(); + /// assert!(value.is_empty()); + /// assert_eq!(value.len(), 0); + /// ``` + pub fn empty() -> Self { + Self(CorsList::empty()) + } +} diff --git a/crates/http_headers/src/headers/cors/access_control_max_age.rs b/crates/http_headers/src/headers/cors/access_control_max_age.rs new file mode 100644 index 000000000..b94042a33 --- /dev/null +++ b/crates/http_headers/src/headers/cors/access_control_max_age.rs @@ -0,0 +1,334 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::time::Duration; + +use super::shared::{invalid_number, invalid_syntax, trimmed_range, untrimmed_range}; +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Defines the `Access-Control-Max-Age` header. +/// +/// # Specification +/// +/// Defined by the Fetch standard's +/// [CORS protocol and credentials section](https://fetch.spec.whatwg.org/#http-access-control-max-age). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{AccessControlMaxAge, AccessControlMaxAgeOwned}; +/// +/// let mut map = HeaderMap::new(); +/// AccessControlMaxAge::insert(&mut map, AccessControlMaxAgeOwned::new(600))?; +/// assert!(AccessControlMaxAge::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct AccessControlMaxAge { + _private: (), +} + +/// Owned value for the `Access-Control-Max-Age` header. +/// +/// The field value carries nothing but a delta-seconds count, so that count is +/// all this type keeps: leading zeroes and surrounding whitespace are accepted +/// when decoding and dropped, and encoding renders the canonical decimal. +/// +/// # Specification +/// +/// Defined by the Fetch standard's [CORS protocol and credentials section]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::AccessControlMaxAgeOwned::new(600); +/// assert_eq!(value.seconds(), 600); +/// ``` +/// +/// `Access-Control-Max-Age: 600` permits caching for ten minutes, while +/// `Access-Control-Max-Age: 0` disables reuse. Decoding ` 00600 ` yields the +/// same header as decoding `600`. +/// +/// [CORS protocol and credentials section]: https://fetch.spec.whatwg.org/#http-access-control-max-age +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct AccessControlMaxAgeOwned { + seconds: u64, +} + +impl AccessControlMaxAgeOwned { + /// Constructs a max age in seconds. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlMaxAgeOwned; + /// + /// let value = AccessControlMaxAgeOwned::new(600); + /// assert_eq!(value.seconds(), 600); + /// assert_eq!(value.to_string(), "600"); + /// ``` + pub const fn new(seconds: u64) -> Self { + Self { seconds } + } + + /// Constructs a max age from a whole-second duration. + /// + /// # Errors + /// + /// Returns [`DecodeErrorKind::InvalidNumber`] if the duration contains + /// fractional seconds. + /// + /// # Examples + /// + /// ``` + /// use std::time::Duration; + /// + /// use http_headers::headers::AccessControlMaxAgeOwned; + /// + /// let value = AccessControlMaxAgeOwned::from_duration(Duration::from_secs(2))?; + /// assert_eq!(value.seconds(), 2); + /// assert_eq!(value.duration(), Duration::from_secs(2)); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn from_duration(duration: Duration) -> Result { + if duration.subsec_nanos() != 0 { + return Err(DecodeError::new(&FieldName::AccessControlMaxAge, DecodeErrorKind::InvalidNumber)); + } + Ok(Self { + seconds: duration.as_secs(), + }) + } + + /// Returns the max age in seconds. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlMaxAgeOwned; + /// + /// let value = AccessControlMaxAgeOwned::new(600); + /// assert_eq!(value.seconds(), 600); + /// + /// let disabled = AccessControlMaxAgeOwned::new(0); + /// assert_eq!(disabled.seconds(), 0); + /// ``` + #[expect(clippy::trivially_copy_pass_by_ref, reason = "accessors consistently borrow owned header values")] + pub const fn seconds(&self) -> u64 { + self.seconds + } + + /// Returns the max age as a duration. + #[must_use] + /// # Examples + /// + /// ``` + /// use std::time::Duration; + /// + /// use http_headers::headers::AccessControlMaxAgeOwned; + /// + /// let value = AccessControlMaxAgeOwned::new(600); + /// assert_eq!(value.duration(), Duration::from_secs(600)); + /// ``` + #[expect(clippy::trivially_copy_pass_by_ref, reason = "accessors consistently borrow owned header values")] + pub const fn duration(&self) -> Duration { + Duration::from_secs(self.seconds) + } + + /// Renders the canonical field value. + /// + /// Decoding does not keep the bytes it read, so a value that carried + /// leading zeroes or surrounding whitespace comes back without them. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlMaxAgeOwned; + /// + /// let value = AccessControlMaxAgeOwned::try_from(" 00600 ")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"600"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(AccessControlMaxAgeOwned, |value| FieldValue::from(value.seconds)); + +impl fmt::Display for AccessControlMaxAgeOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.seconds.fmt(f) + } +} +/// Reads the delta-seconds a `Access-Control-Max-Age` field value holds. +fn max_age_seconds(value: FieldValueRef<'_>) -> Result { + let all = value.as_bytes(); + let range = untrimmed_range(all).unwrap_or_else(|| trimmed_range(all)); + let bytes = &all[range]; + validate::decimal_u64(bytes).ok_or_else(|| invalid_number(&FieldName::AccessControlMaxAge)) +} + +impl Field for AccessControlMaxAge { + type View<'a> = AccessControlMaxAgeOwned; + type Owned = AccessControlMaxAgeOwned; + + fn name() -> &'static FieldName { + &FieldName::AccessControlMaxAge + } + + /// Reads the count directly, since the field value carries nothing a + /// borrowed form could retain that the count does not already capture. + #[inline] + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_custom_source()?; + let seconds = max_age_seconds(lines.exactly_one()?)?; + Ok(Some(AccessControlMaxAgeOwned { seconds })) + } + + /// Reads the count without building a view, since the owned form keeps + /// nothing else. + #[inline] + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_custom_source()?; + let seconds = max_age_seconds(lines.exactly_one()?)?; + Ok(Some(AccessControlMaxAgeOwned { seconds })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::single(FieldValue::from(value.seconds))) + } +} + +impl TryFrom for AccessControlMaxAgeOwned { + type Error = DecodeError; + + fn try_from(duration: Duration) -> Result { + Self::from_duration(duration) + } +} + +impl From for Duration { + fn from(value: AccessControlMaxAgeOwned) -> Self { + value.duration() + } +} + +super::super::shared::impl_string_conversions!(AccessControlMaxAgeOwned, &FieldName::AccessControlMaxAge, invalid_syntax, value); + +impl TryFrom for AccessControlMaxAgeOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + let seconds = max_age_seconds(value.as_field_value_ref())?; + Ok(Self { seconds }) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::time::Duration; + + use super::{AccessControlMaxAge, AccessControlMaxAgeOwned}; + use crate::headers::cors::test_map::TestMap; + use crate::{DecodeErrorKind, FieldName, FieldValue}; + + #[test] + fn constructors_accessors_display_and_conversions_are_canonical() { + let value = AccessControlMaxAgeOwned::new(600); + assert_eq!(value.seconds(), 600); + assert_eq!(value.duration(), Duration::from_mins(10)); + assert_eq!(value.to_string(), "600"); + assert_eq!(value.into_field_value(), "600"); + + let from_duration = AccessControlMaxAgeOwned::from_duration(Duration::from_secs(2)).unwrap(); + assert_eq!(from_duration.seconds(), 2); + assert_eq!( + AccessControlMaxAgeOwned::try_from(Duration::from_secs(9)).unwrap(), + AccessControlMaxAgeOwned::new(9) + ); + assert_eq!(Duration::from(AccessControlMaxAgeOwned::new(7)), Duration::from_secs(7)); + + assert_eq!( + AccessControlMaxAgeOwned::try_from(" 00600 ").expect("whitespace and leading zeroes"), + value + ); + assert_eq!( + AccessControlMaxAgeOwned::try_from(String::from("600")).expect("owned decimal"), + value + ); + assert_eq!( + AccessControlMaxAgeOwned::try_from(FieldValue::from_static("600")).expect("decimal field"), + value + ); + } + + #[test] + fn header_paths_and_numeric_errors_cover_borrowed_and_owned_decoding() { + let source = TestMap::new(&FieldName::AccessControlMaxAge, vec![FieldValue::from_static(" 00600 ")]); + let view = AccessControlMaxAge::view(&source).expect("valid max age").expect("present"); + assert_eq!(view.seconds(), 600); + assert_eq!(view.duration(), Duration::from_mins(10)); + + let owned = AccessControlMaxAge::owned(&source).expect("valid owned max age").expect("present"); + assert_eq!(owned.seconds(), 600); + assert_eq!(view, owned); + + let mut sink = TestMap::new(&FieldName::Accept, Vec::new()); + AccessControlMaxAge::insert(&mut sink, owned).expect("insert max age"); + assert_eq!(sink.name, &FieldName::AccessControlMaxAge); + assert_eq!(sink.values, [FieldValue::from_static("600")]); + + let absent = TestMap::new(&FieldName::Accept, Vec::new()); + assert!(AccessControlMaxAge::view(&absent).expect("absent").is_none()); + assert!(AccessControlMaxAge::owned(&absent).expect("absent").is_none()); + + for wire in ["", " ", "-1", "+1", "1.0", "18446744073709551616"] { + let error = AccessControlMaxAgeOwned::try_from(wire).expect_err("invalid delta-seconds"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber, "{wire:?}"); + } + let error = AccessControlMaxAgeOwned::try_from(String::from("bad\nvalue")).expect_err("invalid field bytes"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!( + AccessControlMaxAgeOwned::try_from("bad\nvalue") + .expect_err("invalid borrowed field bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let duplicate = TestMap::new( + &FieldName::AccessControlMaxAge, + vec![FieldValue::from_static("1"), FieldValue::from_static("2")], + ); + assert_eq!( + AccessControlMaxAge::view(&duplicate).expect_err("singleton header").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + } +} diff --git a/crates/http_headers/src/headers/cors/access_control_request_headers.rs b/crates/http_headers/src/headers/cors/access_control_request_headers.rs new file mode 100644 index 000000000..9a95d1306 --- /dev/null +++ b/crates/http_headers/src/headers/cors/access_control_request_headers.rs @@ -0,0 +1,21 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::super::FieldNameView; +use super::shared::{CorsList, CorsListView, define_header_name_list}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef, validate}; + +define_header_name_list!( + AccessControlRequestHeaders, + AccessControlRequestHeadersOwned, + AccessControlRequestHeadersView, + "Access-Control-Request-Headers", + &FieldName::AccessControlRequestHeaders, + false, + "Defined by the Fetch standard's [CORS-preflight fetch section](https://fetch.spec.whatwg.org/#http-access-control-request-headers).", + "`Access-Control-Request-Headers: Content-Type` names one field; `Access-Control-Request-Headers: X-Custom-Field, Authorization` names several." +); diff --git a/crates/http_headers/src/headers/cors/access_control_request_method.rs b/crates/http_headers/src/headers/cors/access_control_request_method.rs new file mode 100644 index 000000000..daf46178c --- /dev/null +++ b/crates/http_headers/src/headers/cors/access_control_request_method.rs @@ -0,0 +1,537 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::ops::Range; + +use super::super::MethodView; +use super::shared::{common_method, invalid_syntax, invalid_token, method_ref, trimmed_range, untrimmed_range}; +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef}; + +/// Defines the `Access-Control-Request-Method` header. +/// +/// # Specification +/// +/// Defined by the Fetch standard's +/// [CORS-preflight fetch section](https://fetch.spec.whatwg.org/#http-access-control-request-method). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{AccessControlRequestMethod, AccessControlRequestMethodOwned}; +/// +/// let mut map = HeaderMap::new(); +/// let method = AccessControlRequestMethodOwned::from_method("PATCH")?; +/// AccessControlRequestMethod::insert(&mut map, method)?; +/// assert!(AccessControlRequestMethod::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct AccessControlRequestMethod { + _private: (), +} + +/// Owned value for the `Access-Control-Request-Method` header. +/// +/// # Specification +/// +/// Defined by the Fetch standard's [CORS-preflight fetch section]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::AccessControlRequestMethodOwned::try_from("PATCH")?; +/// assert_eq!(value.method()?.as_str(), "PATCH"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Access-Control-Request-Method: PATCH` requests permission to use `PATCH`; +/// extension method tokens are also preserved. +/// +/// [CORS-preflight fetch section]: https://fetch.spec.whatwg.org/#http-access-control-request-method +#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct AccessControlRequestMethodOwned { + method: RequestMethodStore, +} + +/// Storage for an owned request method. +/// +/// Registered methods are known constants, so they are held by name and the +/// wire storage is kept only for extension tokens. +#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +enum RequestMethodStore { + Registered(&'static str), + Extension(FieldValue), +} + +impl RequestMethodStore { + #[expect(clippy::inline_always, reason = "measured: fusing the store into the caller saves 29 Ir")] + #[inline(always)] + fn of(values: &FieldLines<'_>) -> Result { + let mut repeated = values.repeated(); + let value = repeated.next().expect("FieldLines always contains at least one field line"); + if repeated.next().is_some() { + return Err(DecodeError::new(values.name(), DecodeErrorKind::UnexpectedMultipleValues)); + } + if let Some(method) = registered_method(value.as_bytes()) { + return Ok(Self::Registered(method)); + } + extension_method(value.as_bytes())?; + let owned = values.exactly_one_owned()?; + Ok(Self::Extension(owned)) + } + + fn into_field_value(self) -> FieldValue { + match self { + Self::Registered(method) => FieldValue::from_static(method), + Self::Extension(value) => value, + } + } + + fn as_method(&self) -> Result, DecodeError> { + match self { + Self::Registered(method) => Ok(MethodView::from_validated(method.as_bytes())), + Self::Extension(value) => extension_method(value.as_bytes()), + } + } +} + +/// Recognizes a registered method for owned storage. +/// +/// Kept apart from [`common_method`] so that inlining it into the owned +/// decoding path cannot change how the borrowed path is compiled. +#[expect(clippy::inline_always, reason = "measured: keeping the owned matcher separate saves 29 Ir")] +#[inline(always)] +fn registered_method(token: &[u8]) -> Option<&'static str> { + match token { + b"GET" => Some("GET"), + b"PUT" => Some("PUT"), + b"HEAD" => Some("HEAD"), + b"POST" => Some("POST"), + b"PATCH" => Some("PATCH"), + b"TRACE" => Some("TRACE"), + b"DELETE" => Some("DELETE"), + b"CONNECT" => Some("CONNECT"), + b"OPTIONS" => Some("OPTIONS"), + _ => None, + } +} + +/// Borrowed value for the `Access-Control-Request-Method` header. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{AccessControlRequestMethod, AccessControlRequestMethodView}; +/// use http_headers::source::{FieldLines, FieldSource}; +/// use http_headers::{Field, FieldName}; +/// +/// struct Source; +/// +/// impl FieldSource for Source { +/// fn lines(&self, name: &'static FieldName) -> Option> { +/// (name == &FieldName::AccessControlRequestMethod) +/// .then(|| FieldLines::single(name, b"POST")) +/// } +/// } +/// +/// let view: AccessControlRequestMethodView<'_> = +/// AccessControlRequestMethod::view(&Source)?.expect("request method"); +/// assert_eq!(view.method().as_str(), "POST"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct AccessControlRequestMethodView<'a> { + value: FieldValueRef<'a>, + method: MethodView<'a>, +} + +impl AccessControlRequestMethodOwned { + /// Constructs a request method, including extension methods. + /// + /// # Errors + /// + /// Returns an error if `method` is not an HTTP method token. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlRequestMethodOwned; + /// + /// let value = AccessControlRequestMethodOwned::from_method("PATCH")?; + /// assert_eq!(value.method()?.as_str(), "PATCH"); + /// assert!(AccessControlRequestMethodOwned::from_method("bad method").is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_method(method: impl AsRef) -> Result { + Self::try_from(method.as_ref()) + } + + /// Returns the method without allocating. + /// + /// [`MethodView`] preserves spelling and compares case-sensitively. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the field value. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlRequestMethodOwned; + /// + /// let registered = AccessControlRequestMethodOwned::from_method("POST")?; + /// assert_eq!(registered.method()?.as_str(), "POST"); + /// + /// let extension = AccessControlRequestMethodOwned::try_from(" CUSTOM ")?; + /// assert_eq!(extension.method()?.as_str(), "CUSTOM"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn method(&self) -> Result, DecodeError> { + self.method.as_method() + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlRequestMethodOwned; + /// + /// let registered = AccessControlRequestMethodOwned::from_method("DELETE")?; + /// assert_eq!(registered.into_field_value(), "DELETE"); + /// + /// let extension = AccessControlRequestMethodOwned::try_from(" CUSTOM ")?; + /// assert_eq!(extension.into_field_value(), " CUSTOM "); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } + + #[cfg(all(feature = "serde", feature = "headers-cors"))] + pub(crate) fn field_value(&self) -> FieldValueRef<'_> { + match &self.method { + RequestMethodStore::Registered(method) => FieldValueRef::new(method.as_bytes()), + RequestMethodStore::Extension(value) => value.as_field_value_ref(), + } + } +} + +super::super::shared::impl_field_value_conversion!(AccessControlRequestMethodOwned, |value| value.method.into_field_value()); + +impl fmt::Display for AccessControlRequestMethodOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.method().map_err(|_invalid| fmt::Error)?.as_str()) + } +} + +impl<'a> AccessControlRequestMethodView<'a> { + /// Returns the method without allocating. + /// + /// [`MethodView`] preserves spelling and compares case-sensitively. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlRequestMethod; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source; + /// + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::AccessControlRequestMethod) + /// .then(|| FieldLines::single(name, b"PATCH")) + /// } + /// } + /// + /// let view = AccessControlRequestMethod::view(&Source)?.expect("request method"); + /// assert_eq!(view.method().as_str(), "PATCH"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn method(self) -> MethodView<'a> { + self.method + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AccessControlRequestMethod; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source; + /// + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::AccessControlRequestMethod) + /// .then(|| FieldLines::single(name, b" CUSTOM ")) + /// } + /// } + /// + /// let view = AccessControlRequestMethod::view(&Source)?.expect("request method"); + /// assert_eq!(view.method().as_str(), "CUSTOM"); + /// assert_eq!(view.as_field_value(), " CUSTOM "); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } +} + +/// Validates one `Access-Control-Request-Method` field value. +/// +/// Registered methods are recognized whole, so the common case never reaches +/// the token scan behind it. +#[inline] +fn request_method_of(bytes: &[u8]) -> Result, DecodeError> { + match common_method(bytes) { + Some(method) => Ok(MethodView::from_validated(method.as_bytes())), + None => extension_method(bytes), + } +} + +/// Validates a field value that no registered method matched. +fn extension_method(bytes: &[u8]) -> Result, DecodeError> { + let token = &bytes[request_method_range(bytes)]; + method_ref(token).ok_or_else(|| invalid_token(&FieldName::AccessControlRequestMethod)) +} + +/// Locates the method token inside a `Access-Control-Request-Method` value. +fn request_method_range(bytes: &[u8]) -> Range { + untrimmed_range(bytes).unwrap_or_else(|| trimmed_range(bytes)) +} + +/// Reads the sole field line, with the cardinality check folded into the caller. +/// +/// [`FieldLines::exactly_one`] carries only an inlining hint, which the +/// optimizer declines on this path, leaving a call in front of what is +/// otherwise a token match. Forcing the fusion here mirrors what +/// [`RequestMethodStore::of`] already does for the owned form, so the two arms +/// reach the same shape as well as the same answer. +#[expect( + clippy::inline_always, + reason = "fusing the cardinality check into the caller keeps the single-line view a straight run" +)] +#[inline(always)] +fn sole_line<'a>(values: &FieldLines<'a>) -> Result, DecodeError> { + let mut repeated = values.repeated(); + let value = repeated.next().expect("FieldLines always contains at least one field line"); + if repeated.next().is_some() { + return Err(DecodeError::new(values.name(), DecodeErrorKind::UnexpectedMultipleValues)); + } + Ok(value) +} + +impl Field for AccessControlRequestMethod { + type View<'a> = AccessControlRequestMethodView<'a>; + type Owned = AccessControlRequestMethodOwned; + + fn name() -> &'static FieldName { + &FieldName::AccessControlRequestMethod + } + + #[inline] + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_custom_source()?; + let value = sole_line(&lines)?; + let method = request_method_of(value.as_bytes())?; + Ok(Some(AccessControlRequestMethodView { value, method })) + } + + /// Validates without building a view, since the owned form keeps only the + /// field value. + #[inline] + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_custom_source()?; + let method = RequestMethodStore::of(&lines)?; + Ok(Some(AccessControlRequestMethodOwned { method })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::single(value.method.into_field_value())) + } +} + +#[cfg(feature = "http")] +impl TryFrom for AccessControlRequestMethodOwned { + type Error = DecodeError; + + fn try_from(method: http::Method) -> Result { + Self::from_method(method) + } +} + +super::super::shared::impl_string_conversions!( + AccessControlRequestMethodOwned, + &FieldName::AccessControlRequestMethod, + invalid_syntax, + value +); + +impl TryFrom for AccessControlRequestMethodOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + if let Some(method) = registered_method(value.as_bytes()) { + return Ok(Self { + method: RequestMethodStore::Registered(method), + }); + } + extension_method(value.as_bytes())?; + Ok(Self { + method: RequestMethodStore::Extension(value), + }) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{AccessControlRequestMethod, AccessControlRequestMethodOwned, RequestMethodStore, registered_method}; + use crate::headers::cors::test_map::TestMap; + use crate::{DecodeErrorKind, FieldName, FieldValue}; + + #[test] + fn registered_and_extension_constructors_cover_storage_and_accessors() { + for method in ["GET", "PUT", "HEAD", "POST", "PATCH", "TRACE", "DELETE", "CONNECT", "OPTIONS"] { + assert_eq!(registered_method(method.as_bytes()), Some(method), "{method}"); + let owned = AccessControlRequestMethodOwned::from_method(method).expect("registered method"); + assert_eq!(owned.method().expect("method view").as_str(), method); + assert_eq!(owned.to_string(), method); + assert_eq!(owned.into_field_value(), method); + } + assert_eq!(registered_method(b"CUSTOM"), None); + + let extension = AccessControlRequestMethodOwned::try_from(" CUSTOM ").expect("trimmed extension"); + assert_eq!(extension.method().expect("extension").as_str(), "CUSTOM"); + assert_eq!(extension.into_field_value(), " CUSTOM "); + assert_eq!( + AccessControlRequestMethodOwned::try_from(String::from("CUSTOM")) + .expect("owned extension") + .method() + .expect("method") + .as_str(), + "CUSTOM" + ); + assert_eq!( + AccessControlRequestMethodOwned::try_from(FieldValue::from_static("CUSTOM")) + .expect("field extension") + .method() + .expect("method") + .as_str(), + "CUSTOM" + ); + + #[cfg(feature = "http")] + assert_eq!( + AccessControlRequestMethodOwned::try_from(http::Method::PATCH) + .expect("HTTP method") + .method() + .expect("method") + .as_str(), + "PATCH" + ); + } + + #[test] + fn header_decode_insert_absence_duplicates_and_invalid_tokens_are_reported() { + let registered = TestMap::new(&FieldName::AccessControlRequestMethod, vec![FieldValue::from_static("PATCH")]); + let view = AccessControlRequestMethod::view(®istered) + .expect("registered view") + .expect("present"); + assert_eq!(view.method().as_str(), "PATCH"); + assert_eq!(view.as_field_value(), "PATCH"); + let owned = AccessControlRequestMethod::owned(®istered) + .expect("registered owned") + .expect("present"); + assert_eq!(owned.method().expect("method").as_str(), "PATCH"); + + let extension = TestMap::new(&FieldName::AccessControlRequestMethod, vec![FieldValue::from_static(" CUSTOM ")]); + assert_eq!( + AccessControlRequestMethod::view(&extension) + .expect("extension view") + .expect("present") + .method() + .as_str(), + "CUSTOM" + ); + let extension_owned = AccessControlRequestMethod::owned(&extension) + .expect("extension owned") + .expect("present"); + assert_eq!(extension_owned.method().expect("extension").as_str(), "CUSTOM"); + + let mut sink = TestMap::new(&FieldName::Accept, Vec::new()); + AccessControlRequestMethod::insert(&mut sink, extension_owned).expect("insert request method"); + assert_eq!(sink.name, &FieldName::AccessControlRequestMethod); + assert_eq!(sink.values, extension.values); + + let absent = TestMap::new(&FieldName::Accept, Vec::new()); + assert!(AccessControlRequestMethod::view(&absent).expect("absent").is_none()); + assert!(AccessControlRequestMethod::owned(&absent).expect("absent").is_none()); + + let duplicate = TestMap::new( + &FieldName::AccessControlRequestMethod, + vec![FieldValue::from_static("GET"), FieldValue::from_static("POST")], + ); + assert_eq!( + AccessControlRequestMethod::owned(&duplicate).expect_err("singleton header").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!( + AccessControlRequestMethod::view(&duplicate).expect_err("singleton header").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + + for (wire, kind) in [ + ("", DecodeErrorKind::InvalidToken), + (" ", DecodeErrorKind::InvalidToken), + ("bad method", DecodeErrorKind::InvalidToken), + ("bad,method", DecodeErrorKind::InvalidToken), + ] { + let error = AccessControlRequestMethodOwned::try_from(wire).expect_err("invalid request method"); + assert_eq!(error.kind(), kind, "{wire:?}"); + } + let error = AccessControlRequestMethodOwned::try_from(String::from("bad\nvalue")).expect_err("invalid field bytes"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!( + AccessControlRequestMethodOwned::try_from("bad\nvalue") + .expect_err("invalid borrowed field bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn corrupted_extension_storage_is_revalidated_by_accessors() { + let value = AccessControlRequestMethodOwned { + method: RequestMethodStore::Extension(FieldValue::from_static("bad method")), + }; + assert_eq!(value.method().expect_err("corrupt method").kind(), DecodeErrorKind::InvalidToken); + } +} diff --git a/crates/http_headers/src/headers/cors/cors_header_names.rs b/crates/http_headers/src/headers/cors/cors_header_names.rs new file mode 100644 index 000000000..c6a7eff72 --- /dev/null +++ b/crates/http_headers/src/headers/cors/cors_header_names.rs @@ -0,0 +1,54 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::iter::FusedIterator; +use std::{fmt, slice}; + +use super::super::FieldNameView; +use super::cors_tokens::CorsTokens; +use crate::FieldValue; + +/// Borrowed field-name iterator for an owned CORS header-name list. +/// +/// Names retain their original spelling and order. Empty members are skipped. +/// +/// # Examples +/// +/// ``` +/// use http_headers::headers::{AccessControlAllowHeadersOwned, CorsHeaderNames}; +/// +/// let value = AccessControlAllowHeadersOwned::from_header_names(["X-Trace", "content-type"])?; +/// let names: CorsHeaderNames<'_> = (&value).into_iter(); +/// assert_eq!( +/// names.map(|name| name.as_str()).collect::>(), +/// ["X-Trace", "content-type"] +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct CorsHeaderNames<'a> { + tokens: CorsTokens<'a>, +} + +impl<'a> CorsHeaderNames<'a> { + pub(super) const fn new(values: slice::Iter<'a, FieldValue>) -> Self { + Self { + tokens: CorsTokens::new(values), + } + } +} + +impl fmt::Debug for CorsHeaderNames<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("CorsHeaderNames").finish_non_exhaustive() + } +} + +impl<'a> Iterator for CorsHeaderNames<'a> { + type Item = FieldNameView<'a>; + + fn next(&mut self) -> Option { + self.tokens.next().map(FieldNameView::from_validated) + } +} + +impl FusedIterator for CorsHeaderNames<'_> {} diff --git a/crates/http_headers/src/headers/cors/cors_methods.rs b/crates/http_headers/src/headers/cors/cors_methods.rs new file mode 100644 index 000000000..1be5970bd --- /dev/null +++ b/crates/http_headers/src/headers/cors/cors_methods.rs @@ -0,0 +1,54 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::iter::FusedIterator; +use std::{fmt, slice}; + +use super::super::MethodView; +use super::cors_tokens::CorsTokens; +use crate::FieldValue; + +/// Borrowed method iterator for an owned CORS method list. +/// +/// Methods retain their original spelling and order. Empty members are skipped. +/// +/// # Examples +/// +/// ``` +/// use http_headers::headers::{AccessControlAllowMethodsOwned, CorsMethods}; +/// +/// let value = AccessControlAllowMethodsOwned::from_methods(["GET", "X-CUSTOM"])?; +/// let methods: CorsMethods<'_> = (&value).into_iter(); +/// assert_eq!( +/// methods.map(|method| method.as_str()).collect::>(), +/// ["GET", "X-CUSTOM"] +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct CorsMethods<'a> { + tokens: CorsTokens<'a>, +} + +impl<'a> CorsMethods<'a> { + pub(super) const fn new(values: slice::Iter<'a, FieldValue>) -> Self { + Self { + tokens: CorsTokens::new(values), + } + } +} + +impl fmt::Debug for CorsMethods<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("CorsMethods").finish_non_exhaustive() + } +} + +impl<'a> Iterator for CorsMethods<'a> { + type Item = MethodView<'a>; + + fn next(&mut self) -> Option { + self.tokens.next().map(MethodView::from_validated) + } +} + +impl FusedIterator for CorsMethods<'_> {} diff --git a/crates/http_headers/src/headers/cors/cors_tokens.rs b/crates/http_headers/src/headers/cors/cors_tokens.rs new file mode 100644 index 000000000..ab2fbbff8 --- /dev/null +++ b/crates/http_headers/src/headers/cors/cors_tokens.rs @@ -0,0 +1,37 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::slice; + +use crate::{FieldValue, validate}; + +// Concrete state gives IntoIterator a nameable type without boxing a closure-based iterator. +pub(super) struct CorsTokens<'a> { + values: slice::Iter<'a, FieldValue>, + remaining: &'a [u8], +} + +impl<'a> CorsTokens<'a> { + pub(super) const fn new(values: slice::Iter<'a, FieldValue>) -> Self { + Self { values, remaining: &[] } + } +} + +impl<'a> Iterator for CorsTokens<'a> { + type Item = &'a [u8]; + + fn next(&mut self) -> Option { + loop { + if self.remaining.is_empty() { + self.remaining = self.values.next()?.as_bytes(); + } + let end = self.remaining.iter().position(|byte| *byte == b',').unwrap_or(self.remaining.len()); + let (item, rest) = self.remaining.split_at(end); + self.remaining = rest.get(1..).unwrap_or_default(); + let item = validate::trim_ows(item); + if !item.is_empty() { + return Some(item); + } + } + } +} diff --git a/crates/http_headers/src/headers/cors/mod.rs b/crates/http_headers/src/headers/cors/mod.rs new file mode 100644 index 000000000..d78d0f299 --- /dev/null +++ b/crates/http_headers/src/headers/cors/mod.rs @@ -0,0 +1,46 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Cross-Origin Resource Sharing request and response headers. + +mod access_control_allow_credentials; +mod access_control_allow_headers; +mod access_control_allow_methods; +mod access_control_allow_origin; +mod access_control_expose_headers; +mod access_control_max_age; +mod access_control_request_headers; +mod access_control_request_method; +mod cors_header_names; +mod cors_methods; +mod cors_tokens; +mod shared; +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod test_map; + +#[doc(inline)] +pub use access_control_allow_credentials::{ + AccessControlAllowCredentials, AccessControlAllowCredentialsOwned, AccessControlAllowCredentialsView, +}; +#[doc(inline)] +pub use access_control_allow_headers::{AccessControlAllowHeaders, AccessControlAllowHeadersOwned, AccessControlAllowHeadersView}; +#[doc(inline)] +pub use access_control_allow_methods::{AccessControlAllowMethods, AccessControlAllowMethodsOwned, AccessControlAllowMethodsView}; +#[doc(inline)] +pub use access_control_allow_origin::{ + AccessControlAllowOrigin, AccessControlAllowOriginKind, AccessControlAllowOriginOwned, AccessControlAllowOriginView, OriginDomainView, + OriginHost, OriginScheme, SerializedOriginView, +}; +#[doc(inline)] +pub use access_control_expose_headers::{AccessControlExposeHeaders, AccessControlExposeHeadersOwned, AccessControlExposeHeadersView}; +#[doc(inline)] +pub use access_control_max_age::{AccessControlMaxAge, AccessControlMaxAgeOwned}; +#[doc(inline)] +pub use access_control_request_headers::{AccessControlRequestHeaders, AccessControlRequestHeadersOwned, AccessControlRequestHeadersView}; +#[doc(inline)] +pub use access_control_request_method::{AccessControlRequestMethod, AccessControlRequestMethodOwned, AccessControlRequestMethodView}; +#[doc(inline)] +pub use cors_header_names::CorsHeaderNames; +#[doc(inline)] +pub use cors_methods::CorsMethods; diff --git a/crates/http_headers/src/headers/cors/shared.rs b/crates/http_headers/src/headers/cors/shared.rs new file mode 100644 index 000000000..d799cd0ad --- /dev/null +++ b/crates/http_headers/src/headers/cors/shared.rs @@ -0,0 +1,1576 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::hash::{Hash, Hasher}; +use std::ops::Range; +use std::{iter, slice}; + +use http_headers_simd::{EmptyMembers, TokenListScan, scan_token_list}; + +#[cfg(test)] +use super::super::tokens::common_header_name; +pub(super) use super::super::tokens::common_method; +use super::super::{FieldNameView, MethodView}; +use crate::sink::EncodedValues; +use crate::source::FieldLines; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, validate}; + +/// One byte is an RFC 9110 token byte. +const CLASS_TOKEN: u8 = 1 << 0; +/// One byte is optional whitespace. +pub(super) const CLASS_OWS: u8 = 1 << 1; +/// One byte is an ASCII decimal digit. +pub(super) const CLASS_DIGIT: u8 = 1 << 2; +/// One byte is permitted anywhere in a serialized domain label. +pub(super) const CLASS_DOMAIN: u8 = 1 << 3; +/// One byte is permitted at the edges of a serialized domain label. +pub(super) const CLASS_LABEL_EDGE: u8 = 1 << 4; +/// One byte is permitted in a serialized authority outside `:`. +pub(super) const CLASS_AUTHORITY: u8 = 1 << 5; +/// One byte is a hexadecimal digit of an IPv6 address in its serialized form. +pub(super) const CLASS_HEX: u8 = 1 << 6; +/// One byte is permitted in the scheme of a serialized origin. +pub(super) const CLASS_SCHEME: u8 = 1 << 7; + +/// Byte classes shared by every CORS parser in this module. +pub(super) static BYTE_CLASS: [u8; 256] = { + let mut table = [0_u8; 256]; + let mut byte = 0_u8; + loop { + let mut class = 0_u8; + if validate::token_byte(byte) { + class |= CLASS_TOKEN; + } + if byte == b' ' || byte == b'\t' { + class |= CLASS_OWS; + } + if byte.is_ascii_digit() { + class |= CLASS_DIGIT | CLASS_DOMAIN | CLASS_LABEL_EDGE | CLASS_HEX; + } + if byte.is_ascii_lowercase() { + class |= CLASS_DOMAIN | CLASS_LABEL_EDGE; + } + if byte == b'-' { + class |= CLASS_DOMAIN; + } + if byte.is_ascii() && !matches!(byte, b'/' | b'?' | b'#' | b'@' | b':') { + class |= CLASS_AUTHORITY; + } + if matches!(byte, b'a'..=b'f') { + class |= CLASS_HEX; + } + if byte.is_ascii_lowercase() || byte.is_ascii_digit() || matches!(byte, b'+' | b'-' | b'.') { + class |= CLASS_SCHEME; + } + table[byte as usize] = class; + if byte == u8::MAX { + break; + } + byte += 1; + } + table +}; + +/// Returns the byte class at `index`, or no class when `index` is past the end. +#[inline] +#[cfg(test)] +fn byte_class(bytes: &[u8], index: usize) -> u8 { + match bytes.get(index) { + Some(byte) => BYTE_CLASS[usize::from(*byte)], + None => 0, + } +} + +/// Owned list storage keeping the single-field-line case inline. +/// +/// Multiple field lines are rare, so the one-line case avoids both the heap +/// allocation and the length word a growable buffer would need. +#[derive(Clone, Eq)] +pub(super) enum CorsList { + One(FieldValue), + Many(Vec), +} + +impl PartialEq for CorsList { + fn eq(&self, other: &Self) -> bool { + self.field_values().eq(other.field_values()) + } +} + +impl Hash for CorsList { + // LLVM emits an uncallable polymorphized instance for this hasher adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + fn hash(&self, state: &mut H) { + hash_cors_list(self, state); + } +} + +#[inline(never)] +fn hash_cors_list(list: &CorsList, mut state: &mut dyn Hasher) { + for value in list.field_values() { + value.hash(&mut state); + } +} + +pub(super) struct CorsListView<'a> { + pub(super) values: FieldLines<'a>, +} + +impl CorsList { + #[inline] + fn from_values(values: Vec) -> Self { + let mut values = values; + if values.len() == 1 { + Self::One(values.pop().expect("length checked above to contain exactly one value")) + } else { + Self::Many(values) + } + } + + #[inline] + pub(super) fn field_values(&self) -> slice::Iter<'_, FieldValue> { + match self { + Self::One(value) => slice::from_ref(value).iter(), + Self::Many(values) => values.iter(), + } + } + + #[inline] + pub(super) fn value_count(&self) -> usize { + match self { + Self::One(_) => 1, + Self::Many(values) => values.len(), + } + } + + pub(super) fn into_field_values(self) -> Vec { + match self { + Self::One(value) => vec![value], + Self::Many(values) => values, + } + } + + pub(super) fn into_encoded(self) -> EncodedValues { + match self { + Self::One(value) => EncodedValues::single(value), + Self::Many(values) => EncodedValues::from_vec(values), + } + } + + // LLVM emits an uncallable polymorphized instance for this adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + pub(super) fn from_items(name: &'static FieldName, items: I, allow_empty: bool) -> Result + where + I: IntoIterator, + S: AsRef, + { + let items = items.into_iter(); + let capacity = items.size_hint().0.saturating_mul(8); + let mut items = ListItemSourceAdapter(items); + list_from_source(name, &mut items, capacity, allow_empty) + } + + pub(super) fn wildcard() -> Self { + Self::One(FieldValue::from_static("*")) + } + + pub(super) fn empty() -> Self { + Self::One(FieldValue::from_static("")) + } + + pub(super) fn from_field_values(name: &'static FieldName, values: Vec, allow_empty: bool) -> Result { + validate_list(name, values.iter().map(FieldValue::as_field_value_ref), allow_empty)?; + Ok(Self::from_values(values)) + } + + pub(super) fn from_field_value(name: &'static FieldName, value: FieldValue, allow_empty: bool) -> Result { + validate_single_list(name, value.as_field_value_ref(), allow_empty)?; + Ok(Self::One(value)) + } + + /// Validates and adopts field lines in one pass over the map entry. + /// + /// Decoding straight to owned storage keeps a second walk off the owned + /// decoding path. + #[expect(clippy::inline_always, reason = "measured: fusing the decode into the caller saves 32 Ir")] + #[inline(always)] + pub(super) fn decode(name: &'static FieldName, values: &FieldLines<'_>, allow_empty: bool) -> Result { + values.validate_list_item_limit(b',', true)?; + let mut repeated = values.repeated_owned()?; + let (first, first_owned) = repeated.next().expect("FieldLines always contains at least one field line"); + let Some(saw_item) = validate_list_value(first.as_bytes()) else { + return Err(invalid_token(name).at_value(0)); + }; + if let Some((second, second_owned)) = repeated.next() { + return Self::decode_rest(name, first_owned, second, second_owned, repeated, saw_item, allow_empty); + } + if !saw_item && !allow_empty { + return Err(invalid_syntax(name)); + } + Ok(Self::One(first_owned)) + } + + /// Handles the rare field line repetition split out of [`Self::decode`]. + // LLVM emits an uncallable polymorphized instance for this iterator adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + #[inline(never)] + fn decode_rest<'a>( + name: &'static FieldName, + first_owned: FieldValue, + second: FieldValueRef<'a>, + second_owned: FieldValue, + rest: impl Iterator, FieldValue)>, + first_saw_item: bool, + allow_empty: bool, + ) -> Result { + let mut rest = rest; + Self::decode_rest_from_source(name, first_owned, second, second_owned, &mut rest, first_saw_item, allow_empty) + } + + fn decode_rest_from_source<'a>( + name: &'static FieldName, + first_owned: FieldValue, + second: FieldValueRef<'a>, + second_owned: FieldValue, + rest: &mut dyn Iterator, FieldValue)>, + first_saw_item: bool, + allow_empty: bool, + ) -> Result { + let mut saw_item = first_saw_item; + let mut stored = Vec::with_capacity(2_usize.saturating_add(rest.size_hint().0)); + stored.push(first_owned); + for (offset, (value, owned)) in iter::once((second, second_owned)).chain(rest).enumerate() { + let Some(has_item) = validate_list_value(value.as_bytes()) else { + return Err(invalid_token(name).at_value(offset.saturating_add(1))); + }; + saw_item |= has_item; + stored.push(owned); + } + if !saw_item && !allow_empty { + return Err(invalid_syntax(name)); + } + Ok(Self::Many(stored)) + } +} + +struct ListBuilder { + name: &'static FieldName, + bytes: Vec, + allow_empty: bool, +} + +trait ListItemSource { + fn next_with(&mut self, visitor: &mut dyn FnMut(&str) -> Result<(), DecodeError>) -> Result; +} + +struct ListItemSourceAdapter(I); + +impl ListItemSource for ListItemSourceAdapter +where + I: Iterator, + S: AsRef, +{ + // LLVM emits an uncallable polymorphized instance for this type adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + fn next_with(&mut self, visitor: &mut dyn FnMut(&str) -> Result<(), DecodeError>) -> Result { + let Some(item) = self.0.next() else { + return Ok(false); + }; + visitor(item.as_ref())?; + Ok(true) + } +} + +fn list_from_source( + name: &'static FieldName, + items: &mut dyn ListItemSource, + capacity: usize, + allow_empty: bool, +) -> Result { + let mut builder = ListBuilder { + name, + bytes: Vec::with_capacity(capacity), + allow_empty, + }; + while items.next_with(&mut |item| append_list_item(&mut builder, item))? {} + finish_list_items(builder) +} + +fn append_list_item(builder: &mut ListBuilder, item: &str) -> Result<(), DecodeError> { + validate_list_item(builder.name, item.as_bytes())?; + if !builder.bytes.is_empty() { + builder.bytes.extend_from_slice(b", "); + } + builder.bytes.extend_from_slice(item.as_bytes()); + Ok(()) +} + +fn finish_list_items(builder: ListBuilder) -> Result { + if builder.bytes.is_empty() && !builder.allow_empty { + return Err(invalid_syntax(builder.name)); + } + Ok(CorsList::One( + super::super::value_from_bytes(builder.name, builder.bytes) + .expect("validated CORS list items contain only valid field-value bytes"), + )) +} + +pub(super) fn view_list_values<'a>( + name: &'static FieldName, + values: Option>, + allow_empty: bool, +) -> Result>, DecodeError> { + let Some(values) = values else { + return Ok(None); + }; + values.validate_list_item_limit(b',', true)?; + let mut repeated = values.repeated(); + let first = repeated.next().expect("FieldLines always contains at least one field line"); + if repeated.next().is_none() { + validate_single_list(name, first, allow_empty)?; + return Ok(Some(CorsListView { values })); + } + validate_list(name, values.repeated(), allow_empty)?; + Ok(Some(CorsListView { values })) +} + +pub(super) fn owned_list_values( + name: &'static FieldName, + values: Option>, + allow_empty: bool, +) -> Result, DecodeError> { + let Some(values) = values else { + return Ok(None); + }; + CorsList::decode(name, &values, allow_empty).map(Some) +} + +macro_rules! define_header_name_list { + ( + $descriptor:ident, + $owned:ident, + $borrowed:ident, + $header_name:literal, + $name:expr, + $allow_empty:expr, + $specification:literal, + $examples:literal + ) => { + #[doc = concat!("Defines the `", $header_name, "` header.")] + #[doc = ""] + #[doc = "# Specification"] + #[doc = ""] + #[doc = $specification] + #[derive(Debug)] + pub struct $descriptor { + _private: (), + } + + #[doc = concat!("Owned value for the `", $header_name, "` header.")] + /// + /// # Specification + #[doc = $specification] + /// + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let value = AccessControlAllowHeadersOwned::from_header_names(["content-type"])?; + /// let fmt = format!("{value:?}"); + /// assert!(fmt.contains("AccessControlAllowHeadersOwned")); + /// assert!(fmt.contains("header_name_count")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[doc = $examples] + #[derive(Clone, Eq, Hash, PartialEq)] + pub struct $owned(CorsList); + + #[doc = concat!("Borrowed value for the `", $header_name, "` header.")] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowHeaders; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static("content-type"), + /// ); + /// let value = AccessControlAllowHeaders::view(&headers)?.expect("present"); + /// let fmt = format!("{value:?}"); + /// assert!(fmt.contains("AccessControlAllowHeadersView")); + /// assert!(fmt.contains("header_name_count")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub struct $borrowed<'a>(CorsListView<'a>); + + impl fmt::Debug for $owned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct(stringify!($owned)) + .field("value_count", &self.0.value_count()) + .field("header_name_count", &self.len()) + .finish() + } + } + + impl fmt::Display for $owned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + crate::headers::shared::fmt_ascii_values(self.field_values(), f) + } + } + + impl fmt::Debug for $borrowed<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct(stringify!($borrowed)) + .field("value_count", &self.0.values.len()) + .field("header_name_count", &self.len()) + .finish() + } + } + + impl $owned { + /// Constructs one canonical field line from validated field names. + /// + /// Duplicate field names are retained in input order. + /// + /// # Errors + /// + /// Returns an error if an item is not an HTTP field-name token or + /// if this header requires at least one member and input is empty. + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let value = AccessControlAllowHeadersOwned::from_header_names(["content-type", "X-Trace-Id"])?; + /// let names = value.iter().map(|name| name.as_str()).collect::>(); + /// assert_eq!(names, ["content-type", "X-Trace-Id"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + // LLVM emits an uncallable polymorphized instance for this adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + pub fn from_header_names(names: I) -> Result + where + I: IntoIterator, + S: AsRef, + { + CorsList::from_items($name, names, $allow_empty).map(Self) + } + + /// Validates and adopts complete field lines. + /// + /// # Errors + /// + /// Returns an error for no field lines, malformed members, or an + /// empty list when this header requires at least one member. + /// # Examples + /// + /// ``` + /// use http_headers::FieldValue; + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let value = AccessControlAllowHeadersOwned::from_field_values(vec![FieldValue::from_static( + /// "content-type, x-trace-id", + /// )])?; + /// assert_eq!(value.len(), 2); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_field_values(values: Vec) -> Result { + CorsList::from_field_values($name, values, $allow_empty).map(Self) + } + + /// Iterates field names in wire order without allocating. + /// + /// Duplicate names and their original casing are preserved. + /// Names compare and hash case-insensitively through [`FieldNameView`]. + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let value = AccessControlAllowHeadersOwned::from_header_names(["content-type", "x-trace-id"])?; + /// let names = value.iter().map(|name| name.as_str()).collect::>(); + /// assert_eq!(names, ["content-type", "x-trace-id"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn iter(&self) -> super::CorsHeaderNames<'_> { + super::CorsHeaderNames::new(self.0.field_values()) + } + + /// Returns the number of list members, including duplicates. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let value = AccessControlAllowHeadersOwned::from_header_names([ + /// "content-type", + /// "x-trace-id", + /// "content-type", + /// ])?; + /// assert_eq!(value.len(), 3); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn len(&self) -> usize { + self.iter().count() + } + + /// Returns whether the list has no members. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let empty = AccessControlAllowHeadersOwned::empty(); + /// assert!(empty.is_empty()); + /// + /// let value = AccessControlAllowHeadersOwned::from_header_names(["content-type"])?; + /// assert!(!value.is_empty()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn is_empty(&self) -> bool { + self.iter().next().is_none() + } + + /// Iterates original field lines in wire order. + /// # Examples + /// + /// ``` + /// use http_headers::FieldValue; + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let value = AccessControlAllowHeadersOwned::try_from(FieldValue::from_static( + /// "content-type, x-trace-id", + /// ))?; + /// let fields = value + /// .field_values() + /// .map(|field| field.as_bytes()) + /// .collect::>(); + /// assert_eq!(fields, [b"content-type, x-trace-id".as_slice()]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn field_values(&self) -> impl Iterator> { + self.0.field_values().map(FieldValue::as_field_value_ref) + } + + /// Returns the original field lines. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::FieldValue; + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let value = AccessControlAllowHeadersOwned::try_from(vec![ + /// FieldValue::from_static("content-type"), + /// FieldValue::from_static("x-trace-id"), + /// ])?; + /// let fields = value.into_field_values(); + /// assert_eq!(fields.len(), 2); + /// assert_eq!(fields[0].as_bytes(), b"content-type"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_values(self) -> Vec { + self.0.into_field_values() + } + } + + impl<'a> IntoIterator for &'a $owned { + type Item = FieldNameView<'a>; + type IntoIter = super::CorsHeaderNames<'a>; + + fn into_iter(self) -> Self::IntoIter { + self.iter() + } + } + + impl<'a> $borrowed<'a> { + /// Iterates field names in wire order without allocating. + /// + /// Duplicate names and their original casing are preserved. + /// Names compare and hash case-insensitively through [`FieldNameView`]. + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowHeaders; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static("content-type, x-trace-id"), + /// ); + /// let value = AccessControlAllowHeaders::view(&headers)?.expect("present"); + /// let names = value.iter().map(|name| name.as_str()).collect::>(); + /// assert_eq!(names, ["content-type", "x-trace-id"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn iter(&self) -> impl Iterator> + '_ { + self.0 + .values + .repeated() + .flat_map(|value| value.as_bytes().split(|byte| *byte == b',')) + .map(validate::trim_ows) + .filter(|item| !item.is_empty()) + .map(super::shared::header_name_ref_validated) + } + + /// Returns the number of list members, including duplicates. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowHeaders; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static("content-type"), + /// ); + /// headers.append( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static("x-trace-id"), + /// ); + /// let value = AccessControlAllowHeaders::view(&headers)?.expect("present"); + /// assert_eq!(value.len(), 2); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn len(&self) -> usize { + self.iter().count() + } + + /// Returns whether the list has no members. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowHeaders; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static(""), + /// ); + /// let value = AccessControlAllowHeaders::view(&headers)?.expect("present"); + /// assert!(value.is_empty()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn is_empty(&self) -> bool { + self.iter().next().is_none() + } + + /// Iterates original field lines in wire order. + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowHeaders; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static("content-type"), + /// ); + /// headers.append( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static("x-trace-id"), + /// ); + /// let value = AccessControlAllowHeaders::view(&headers)?.expect("present"); + /// let fields = value + /// .field_values() + /// .map(|field| field.as_bytes()) + /// .collect::>(); + /// assert_eq!( + /// fields, + /// [b"content-type".as_slice(), b"x-trace-id".as_slice()] + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn field_values(&self) -> impl Iterator> + '_ { + self.0.values.repeated() + } + } + + impl Field for $descriptor { + type View<'a> = $borrowed<'a>; + type Owned = $owned; + + fn name() -> &'static FieldName { + $name + } + + // LLVM emits an uncallable polymorphized instance for this source adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + super::shared::view_list_values($name, source.lines(Self::name()), $allow_empty).map(|view| view.map($borrowed)) + } + + // LLVM emits an uncallable polymorphized instance for this source adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + super::shared::owned_list_values($name, source.lines(Self::name()), $allow_empty).map(|owned| owned.map($owned)) + } + + // LLVM emits an uncallable polymorphized instance for this sink adapter. + #[cfg_attr(coverage_nightly, coverage(off))] + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), value.0.into_encoded()) + } + } + + impl TryFrom for $owned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + CorsList::from_field_value($name, value, $allow_empty).map(Self) + } + } + + impl TryFrom> for $owned { + type Error = DecodeError; + + fn try_from(values: Vec) -> Result { + Self::from_field_values(values) + } + } + }; +} + +pub(super) use define_header_name_list; + +macro_rules! impl_header_name_wildcard { + ($owned:ident, $borrowed:ident) => { + impl $owned { + /// Constructs the wildcard field value. + /// + /// Whether `*` has wildcard semantics depends on the request and + /// is deliberately not inferred here. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let wildcard = AccessControlAllowHeadersOwned::wildcard(); + /// assert!(wildcard.contains_wildcard()); + /// assert!(wildcard.is_wildcard()); + /// ``` + pub fn wildcard() -> Self { + Self(CorsList::wildcard()) + } + + /// Returns whether any member is `*`. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let wildcard = AccessControlAllowHeadersOwned::wildcard(); + /// assert!(wildcard.contains_wildcard()); + /// + /// let explicit = AccessControlAllowHeadersOwned::from_header_names(["content-type"])?; + /// assert!(!explicit.contains_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn contains_wildcard(&self) -> bool { + self.iter().any(|name| name.as_bytes() == b"*") + } + + /// Returns whether `*` is the only list member. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::AccessControlAllowHeadersOwned; + /// + /// let wildcard = AccessControlAllowHeadersOwned::wildcard(); + /// assert!(wildcard.is_wildcard()); + /// + /// let mixed = AccessControlAllowHeadersOwned::from_header_names(["*", "content-type"])?; + /// assert!(mixed.contains_wildcard()); + /// assert!(!mixed.is_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn is_wildcard(&self) -> bool { + let mut names = self.iter(); + names.next().is_some_and(|name| name.as_bytes() == b"*") && names.next().is_none() + } + } + + impl $borrowed<'_> { + /// Returns whether any member is `*`. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowHeaders; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static("*, content-type"), + /// ); + /// let value = AccessControlAllowHeaders::view(&headers)?.expect("present"); + /// assert!(value.contains_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn contains_wildcard(&self) -> bool { + self.iter().any(|name| name.as_bytes() == b"*") + } + + /// Returns whether `*` is the only list member. + #[must_use] + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::AccessControlAllowHeaders; + /// + /// let mut headers = HeaderMap::new(); + /// headers.insert( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static("*"), + /// ); + /// let wildcard = AccessControlAllowHeaders::view(&headers)?.expect("present"); + /// assert!(wildcard.is_wildcard()); + /// + /// let mut mixed_headers = HeaderMap::new(); + /// mixed_headers.insert( + /// "access-control-allow-headers", + /// http::HeaderValue::from_static("*, content-type"), + /// ); + /// let mixed = AccessControlAllowHeaders::view(&mixed_headers)?.expect("present"); + /// assert!(!mixed.is_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn is_wildcard(&self) -> bool { + let mut names = self.iter(); + names.next().is_some_and(|name| name.as_bytes() == b"*") && names.next().is_none() + } + } + }; +} + +pub(super) use impl_header_name_wildcard; + +#[expect(clippy::inline_always, reason = "preserves pre-split inlining in hot borrowed CORS decoders")] +#[inline(always)] +// LLVM emits an uncallable polymorphized instance for this iterator adapter. +#[cfg_attr(coverage_nightly, coverage(off))] +pub(super) fn validate_list<'a>( + name: &'static FieldName, + values: impl IntoIterator>, + allow_empty: bool, +) -> Result<(), DecodeError> { + let mut values = values.into_iter(); + validate_list_from_source(name, &mut values, allow_empty) +} + +#[inline(never)] +fn validate_list_from_source( + name: &'static FieldName, + values: &mut dyn Iterator>, + allow_empty: bool, +) -> Result<(), DecodeError> { + let Some(first) = values.next() else { + return Err(DecodeError::new(name, DecodeErrorKind::MissingValue)); + }; + let Some(mut saw_item) = validate_list_value(first.as_bytes()) else { + return Err(invalid_token(name).at_value(0)); + }; + for (value_index, value) in (1_usize..).zip(values) { + let Some(has_item) = validate_list_value(value.as_bytes()) else { + return Err(invalid_token(name).at_value(value_index)); + }; + saw_item |= has_item; + } + if !saw_item && !allow_empty { + return Err(invalid_syntax(name)); + } + Ok(()) +} + +/// Validates one comma-separated token list, reporting whether it has members. +/// +/// Returns `None` when a member is not a token after optional whitespace was +/// trimmed. Empty members are skipped, matching `#rule` list expansion. +#[expect(clippy::inline_always, reason = "preserves pre-split inlining in hot CORS list decoding")] +#[inline(always)] +fn validate_list_value(bytes: &[u8]) -> Option { + if let Some(members) = common_list_line(bytes) { + return Some(members); + } + match scan_token_list(bytes, EmptyMembers::Skip) { + TokenListScan::Members => Some(true), + TokenListScan::Empty => Some(false), + TokenListScan::Rejected => None, + } +} + +pub(super) fn validate_single_list(name: &'static FieldName, value: FieldValueRef<'_>, allow_empty: bool) -> Result<(), DecodeError> { + let Some(saw_item) = validate_list_value(value.as_bytes()) else { + return Err(invalid_token(name).at_value(0)); + }; + if !saw_item && !allow_empty { + return Err(invalid_syntax(name)); + } + Ok(()) +} + +/// Recognizes a short field line that is known to be well formed. +/// +/// The lines carried by these headers repeat heavily, and matching one whole +/// settles it without a scan; anything else falls through to one. +#[expect(clippy::inline_always, reason = "measured: folds into the caller's dispatch")] +#[inline(always)] +fn common_list_line(bytes: &[u8]) -> Option { + match bytes { + b"*" + | b"GET" + | b"PUT" + | b"POST" + | b"HEAD" + | b"PATCH" + | b"DELETE" + | b"OPTIONS" + | b"GET, POST" + | b"GET, HEAD" + | b"content-type" + | b"authorization" + | b"content-type, x-request-id" + | b"x-request-id, content-type" + | b"etag, x-request-id" + | b"content-type, authorization" + | b"authorization, content-type" + | b"content-type, x-requested-with" => Some(true), + b"" => Some(false), + _ => None, + } +} + +/// Shortest field line the shared list scanner pays for itself on. +/// +/// The shared finder and token validator process long runs a vector block at +/// a time; below one block, the scalar state machine avoids their call cost. +#[cfg(test)] +const LIST_SCAN_MIN_LEN: usize = 16; + +/// Checks one field line a byte at a time. +/// +/// This scan carries field lines shorter than one vector block, so it is +/// deliberately small: whole blocks are the shared scanner's business. +#[expect(clippy::inline_always, reason = "measured: keeps the short-line scan in the caller")] +#[inline(always)] +#[cfg(test)] +fn scan_list_value(bytes: &[u8]) -> Option { + let mut saw_item = false; + let mut index = 0_usize; + while index < bytes.len() { + index = skip_ows(bytes, index); + let start = index; + while byte_class(bytes, index) & CLASS_TOKEN != 0 { + index += 1; + } + saw_item |= index > start; + index = skip_ows(bytes, index); + match bytes.get(index) { + None => break, + Some(&b',') => index += 1, + Some(_) => return None, + } + } + Some(saw_item) +} + +#[cfg(test)] +fn skip_ows(bytes: &[u8], mut index: usize) -> usize { + while byte_class(bytes, index) & CLASS_OWS != 0 { + index += 1; + } + index +} + +fn validate_list_item(name: &'static FieldName, item: &[u8]) -> Result<(), DecodeError> { + if item.is_empty() || !all_token_bytes(item) { + return Err(invalid_token(name)); + } + Ok(()) +} + +#[expect(clippy::inline_always, reason = "preserves pre-split inlining in hot CORS method iteration")] +#[inline(always)] +pub(super) fn method_ref(bytes: &[u8]) -> Option> { + if let Some(method) = common_method(bytes) { + return Some(MethodView::from_validated(method.as_bytes())); + } + if bytes.is_empty() || !all_token_bytes(bytes) { + return None; + } + Some(MethodView::from_validated(bytes)) +} + +#[inline] +pub(super) fn method_ref_validated(bytes: &[u8]) -> MethodView<'_> { + MethodView::from_validated(bytes) +} + +#[inline] +pub(super) fn header_name_ref_validated(bytes: &[u8]) -> FieldNameView<'_> { + FieldNameView::from_validated(bytes) +} + +#[cfg(test)] +fn header_name_ref(bytes: &[u8]) -> Option> { + if let Some(name) = common_header_name(bytes) { + return Some(FieldNameView::from_validated(name.as_bytes())); + } + if bytes.is_empty() || !all_token_bytes(bytes) { + return None; + } + Some(FieldNameView::from_validated(bytes)) +} + +#[expect(clippy::inline_always, reason = "preserves pre-split inlining in hot CORS list iteration")] +#[inline(always)] +fn all_token_bytes(bytes: &[u8]) -> bool { + bytes.iter().all(|byte| BYTE_CLASS[usize::from(*byte)] & CLASS_TOKEN != 0) +} + +pub(super) fn untrimmed_range(bytes: &[u8]) -> Option> { + let first = *bytes.first()?; + let last = bytes[bytes.len() - 1]; + match (first, last) { + (b' ' | b'\t', _) | (_, b' ' | b'\t') => None, + _ => Some(0..bytes.len()), + } +} + +pub(super) fn trimmed_range(bytes: &[u8]) -> Range { + if bytes.first().is_some_and(|byte| !matches!(byte, b' ' | b'\t')) && bytes.last().is_some_and(|byte| !matches!(byte, b' ' | b'\t')) { + return 0..bytes.len(); + } + let start = bytes.iter().position(|byte| !matches!(byte, b' ' | b'\t')).unwrap_or(bytes.len()); + let end = bytes + .iter() + .rposition(|byte| !matches!(byte, b' ' | b'\t')) + .map_or(start, |index| index + 1); + start..end +} + +#[cold] +pub(super) fn invalid_syntax(name: &'static FieldName) -> DecodeError { + DecodeError::new(name, DecodeErrorKind::InvalidSyntax) +} + +#[cold] +pub(super) fn invalid_token(name: &'static FieldName) -> DecodeError { + DecodeError::new(name, DecodeErrorKind::InvalidToken) +} + +#[cold] +pub(super) fn invalid_number(name: &'static FieldName) -> DecodeError { + DecodeError::new(name, DecodeErrorKind::InvalidNumber) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::collections::hash_map::DefaultHasher; + + use super::super::test_map::TestMap; + use super::*; + use crate::headers::{ + AccessControlAllowHeaders, AccessControlAllowHeadersOwned, AccessControlExposeHeaders, AccessControlExposeHeadersOwned, + AccessControlRequestHeaders, AccessControlRequestHeadersOwned, + }; + use crate::sink::FieldSink; + + /// Checks that the short lines settled from the table reach the verdict + /// the scan would have reached. + #[test] + fn common_short_lines_match_the_scan() { + const ALPHABET: &[u8] = b"GETPOS, *-\t\""; + + let known: &[&[u8]] = &[ + b"*", + b"GET", + b"PUT", + b"POST", + b"HEAD", + b"PATCH", + b"DELETE", + b"OPTIONS", + b"GET, POST", + b"GET, HEAD", + b"content-type", + b"authorization", + b"", + ]; + for line in known { + assert_eq!(common_list_line(line), scan_list_value(line), "table disagrees on {line:?}"); + assert!(line.len() < LIST_SCAN_MIN_LEN, "{line:?} is not a short line"); + } + + let mut line = Vec::new(); + for first in ALPHABET { + for second in ALPHABET { + for third in ALPHABET { + line.clear(); + line.extend_from_slice(&[*first, *second, *third]); + if let Some(members) = common_list_line(&line) { + assert_eq!(Some(members), scan_list_value(&line), "table disagrees on {line:?}"); + } + } + } + } + } + + /// Checks that handing long field lines to the shared scanner keeps the + /// verdict the local scan would have reached. + #[test] + fn long_list_lines_scan_the_same_either_way() { + /// Bytes a member is spelled with. + const TOKEN: &[u8] = b"abcXYZ019-_.!"; + /// Bytes that leave the grammar, mutated into otherwise sound lines. + const FOREIGN: &[u8] = b"\"@\x7f\r ,\t"; + + let mut seed = 0x2545_f491_4f6c_dd1d_u64; + let mut next = move || { + seed ^= seed << 13; + seed ^= seed >> 7; + seed ^= seed << 17; + seed + }; + let mut line = Vec::with_capacity(256); + let mut accepted = 0_usize; + let cases = if cfg!(miri) { 256 } else { 20_000 }; + for _ in 0..cases { + let length = LIST_SCAN_MIN_LEN + usize::try_from(next() % 160).unwrap_or_default(); + line.clear(); + while line.len() < length { + if !line.is_empty() { + line.push(b','); + } + for _ in 0..next() % 3 { + line.push(if next() % 2 == 0 { b' ' } else { b'\t' }); + } + for _ in 0..next() % 9 { + let index = usize::try_from(next() % TOKEN.len() as u64).unwrap_or_default(); + line.push(TOKEN[index]); + } + } + if next() % 3 == 0 { + let at = usize::try_from(next() % line.len() as u64).unwrap_or_default(); + let index = usize::try_from(next() % FOREIGN.len() as u64).unwrap_or_default(); + line[at] = FOREIGN[index]; + } + accepted += usize::from(validate_list_value(&line).is_some()); + assert_eq!(validate_list_value(&line), scan_list_value(&line)); + } + assert!( + accepted > cases / 20, + "the generated lines must exercise the accepting path, not only the fallback" + ); + } + + #[test] + fn byte_classes_and_view_conversions_cover_shared_value_helpers() { + let table = BYTE_CLASS; + assert_ne!(table[usize::from(b'A')] & CLASS_TOKEN, 0); + assert_ne!(table[usize::from(b' ')] & CLASS_OWS, 0); + assert_ne!(table[usize::from(b'9')] & CLASS_DIGIT, 0); + assert_ne!(table[usize::from(b'a')] & CLASS_DOMAIN, 0); + assert_ne!(table[usize::from(b'f')] & CLASS_HEX, 0); + assert_ne!(table[usize::from(b'+')] & CLASS_SCHEME, 0); + + let method = MethodView::new("CUSTOM").unwrap(); + assert_eq!(method.as_str(), "CUSTOM"); + assert_eq!(method.as_bytes(), b"CUSTOM"); + #[cfg(feature = "http")] + assert_eq!(method.try_to_method().expect("HTTP method").as_str(), "CUSTOM"); + + let name = FieldNameView::new("X-Trace-Id").unwrap(); + assert_eq!(name.as_str(), "X-Trace-Id"); + assert_eq!(name.as_bytes(), b"X-Trace-Id"); + assert!(name.eq_ignore_ascii_case("x-trace-id")); + assert_eq!(name.try_to_field_name().expect("native header name").as_str(), "x-trace-id"); + #[cfg(feature = "http")] + assert_eq!(name.try_to_http_header_name().expect("HTTP header name").as_str(), "x-trace-id"); + } + + #[test] + fn cors_list_storage_covers_one_many_equality_hashing_and_errors() { + let one = CorsList::from_field_values( + &FieldName::AccessControlAllowHeaders, + vec![FieldValue::from_static("content-type")], + true, + ) + .expect("one field line"); + assert_eq!(one.value_count(), 1); + assert_eq!(one.field_values().count(), 1); + assert_eq!(one.clone().into_field_values().len(), 1); + assert_eq!(one.clone().into_encoded().len(), 1); + assert!(matches!( + CorsList::from_field_values( + &FieldName::AccessControlRequestHeaders, + vec![FieldValue::from_static("")], + false, + ), + Err(error) if error.kind() == DecodeErrorKind::InvalidSyntax + )); + + let many = CorsList::from_field_values( + &FieldName::AccessControlAllowHeaders, + vec![FieldValue::from_static("content-type"), FieldValue::from_static("x-trace-id")], + true, + ) + .expect("several field lines"); + assert_eq!(many.value_count(), 2); + assert_eq!(many.clone().into_field_values().len(), 2); + assert!(many == many.clone()); + assert!(one != many); + + let mut first_hash = DefaultHasher::new(); + many.hash(&mut first_hash); + let mut second_hash = DefaultHasher::new(); + many.clone().hash(&mut second_hash); + assert_eq!(first_hash.finish(), second_hash.finish()); + let mut helper_hash = DefaultHasher::new(); + hash_cors_list(&many, &mut helper_hash); + assert_eq!(first_hash.finish(), helper_hash.finish()); + assert_eq!(many.clone().into_encoded().len(), 2); + + assert_eq!( + AccessControlRequestHeadersOwned::from_header_names(Vec::::new()) + .expect_err("request list must not be empty") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + AccessControlAllowHeadersOwned::from_header_names(["valid", "bad name"]) + .expect_err("invalid token") + .kind(), + DecodeErrorKind::InvalidToken + ); + let single = AccessControlAllowHeadersOwned::from_header_names(["content-type"]).expect("valid single-item list"); + assert_eq!(single.iter().next().expect("field name").as_str(), "content-type"); + let owned = AccessControlAllowHeadersOwned::from_header_names(vec![String::from("content-type")]).expect("valid owned item list"); + assert_eq!(owned.iter().next().expect("field name").as_str(), "content-type"); + assert_eq!( + CorsList::from_field_values(&FieldName::AccessControlAllowHeaders, Vec::new(), true,) + .err() + .expect("a present header needs a field line") + .kind(), + DecodeErrorKind::MissingValue + ); + assert_eq!(CorsList::wildcard().field_values().next().expect("line"), "*"); + assert_eq!(CorsList::empty().field_values().next().expect("line"), ""); + } + + #[test] + fn allow_and_expose_header_lists_cover_wildcard_and_repeated_paths() { + let allow = + AccessControlAllowHeadersOwned::from_header_names(["content-type", "X-Trace-Id", "content-type"]).expect("header names"); + assert_eq!(allow.len(), 3); + assert!(!allow.is_empty()); + assert_eq!( + allow.iter().map(FieldNameView::as_str).collect::>(), + ["content-type", "X-Trace-Id", "content-type"] + ); + assert_eq!(allow.field_values().count(), 1); + assert!(format!("{allow:?}").contains("header_name_count")); + + let wildcard = AccessControlAllowHeadersOwned::wildcard(); + assert!(wildcard.contains_wildcard()); + assert!(wildcard.is_wildcard()); + let mixed = AccessControlAllowHeadersOwned::from_header_names(["*", "x-trace-id"]).expect("mixed wildcard list"); + assert!(mixed.contains_wildcard()); + assert!(!mixed.is_wildcard()); + assert!(AccessControlAllowHeadersOwned::empty().is_empty()); + assert!(AccessControlExposeHeadersOwned::empty().is_empty()); + let expose_wildcard = AccessControlExposeHeadersOwned::wildcard(); + assert!(expose_wildcard.contains_wildcard()); + assert!(expose_wildcard.is_wildcard()); + assert!(format!("{expose_wildcard:?}").contains("header_name_count")); + let expose_mixed = AccessControlExposeHeadersOwned::from_header_names(["*", "x-trace-id"]).expect("mixed expose list"); + assert!(expose_mixed.contains_wildcard()); + assert!(!expose_mixed.is_wildcard()); + + let repeated = TestMap::new( + &FieldName::AccessControlExposeHeaders, + vec![FieldValue::from_static("content-type"), FieldValue::from_static("*, x-trace-id")], + ); + let view = AccessControlExposeHeaders::view(&repeated).expect("valid view").expect("present"); + assert_eq!(view.len(), 3); + assert!(!view.is_empty()); + assert!(view.contains_wildcard()); + assert!(!view.is_wildcard()); + assert_eq!(view.field_values().count(), 2); + assert!(format!("{view:?}").contains("header_name_count")); + + let expose_wildcard_map = TestMap::new(&FieldName::AccessControlExposeHeaders, vec![FieldValue::from_static("*")]); + let expose_wildcard_view = AccessControlExposeHeaders::view(&expose_wildcard_map) + .expect("valid wildcard view") + .expect("present wildcard view"); + assert!(expose_wildcard_view.is_wildcard()); + + let allow_wildcard_map = TestMap::new(&FieldName::AccessControlAllowHeaders, vec![FieldValue::from_static("*")]); + let allow_wildcard_view = AccessControlAllowHeaders::view(&allow_wildcard_map) + .expect("valid wildcard view") + .expect("present wildcard view"); + assert!(allow_wildcard_view.is_wildcard()); + + let owned = AccessControlExposeHeaders::owned(&repeated).expect("valid owned").expect("present"); + assert_eq!(owned.field_values().count(), 2); + assert_eq!(owned.clone().into_field_values().len(), 2); + + let mut sink = TestMap::new(&FieldName::Accept, Vec::new()); + AccessControlExposeHeaders::insert(&mut sink, owned).expect("insert list"); + assert_eq!(sink.name, &FieldName::AccessControlExposeHeaders); + assert_eq!(sink.values.len(), 2); + } + + #[test] + fn request_and_allow_header_lists_cover_generated_conversion_paths() { + let required = AccessControlRequestHeadersOwned::try_from(vec![ + FieldValue::from_static("content-type"), + FieldValue::from_static("x-trace-id"), + ]) + .expect("required repeated list"); + assert_eq!(required.len(), 2); + let from_one = AccessControlRequestHeadersOwned::try_from(FieldValue::from_static("content-type")).expect("one required name"); + assert_eq!(from_one.len(), 1); + assert!(!from_one.is_empty()); + assert_eq!(from_one.field_values().count(), 1); + assert!(format!("{from_one:?}").contains("header_name_count")); + assert_eq!(from_one.clone().into_field_values().len(), 1); + assert_eq!( + AccessControlRequestHeaders::owned(&TestMap::new( + &FieldName::AccessControlRequestHeaders, + vec![FieldValue::from_static("content-type")], + )) + .expect("one-line owned list") + .expect("present") + .len(), + 1 + ); + + let request_source = TestMap::new( + &FieldName::AccessControlRequestHeaders, + vec![FieldValue::from_static("content-type, x-trace-id")], + ); + let request_view = AccessControlRequestHeaders::view(&request_source) + .expect("request view") + .expect("present"); + assert_eq!(request_view.len(), 2); + assert!(!request_view.is_empty()); + assert_eq!(request_view.iter().count(), 2); + assert_eq!(request_view.field_values().count(), 1); + assert!(format!("{request_view:?}").contains("header_name_count")); + let mut request_sink = TestMap::new(&FieldName::Accept, Vec::new()); + AccessControlRequestHeaders::insert(&mut request_sink, from_one).expect("insert request headers"); + assert_eq!(request_sink.name, &FieldName::AccessControlRequestHeaders); + + let allow_source = TestMap::new( + &FieldName::AccessControlAllowHeaders, + vec![FieldValue::from_static("*, content-type")], + ); + let allow_view = AccessControlAllowHeaders::view(&allow_source) + .expect("allow view") + .expect("present"); + assert_eq!(allow_view.len(), 2); + assert!(!allow_view.is_empty()); + assert!(allow_view.contains_wildcard()); + assert!(!allow_view.is_wildcard()); + assert_eq!(allow_view.field_values().count(), 1); + assert!(format!("{allow_view:?}").contains("header_name_count")); + + let allow_from_field = AccessControlAllowHeadersOwned::try_from(FieldValue::from_static("content-type")).expect("allow field"); + assert_eq!(allow_from_field.len(), 1); + let allow_from_fields = + AccessControlAllowHeadersOwned::try_from(vec![FieldValue::from_static("content-type"), FieldValue::from_static("x-trace-id")]) + .expect("allow fields"); + assert_eq!(allow_from_fields.clone().into_field_values().len(), 2); + let mut allow_sink = TestMap::new(&FieldName::Accept, Vec::new()); + AccessControlAllowHeaders::insert(&mut allow_sink, allow_from_fields).expect("insert allow headers"); + + let expose_from_field = AccessControlExposeHeadersOwned::try_from(FieldValue::from_static("content-type")).expect("expose field"); + assert_eq!(expose_from_field.len(), 1); + let expose_from_fields = + AccessControlExposeHeadersOwned::try_from(vec![FieldValue::from_static("content-type"), FieldValue::from_static("x-trace-id")]) + .expect("expose fields"); + assert_eq!(expose_from_fields.len(), 2); + + let absent = TestMap::new(&FieldName::Accept, Vec::new()); + assert!(AccessControlAllowHeaders::view(&absent).expect("absent").is_none()); + assert!(AccessControlAllowHeaders::owned(&absent).expect("absent").is_none()); + } + + #[test] + fn header_name_lists_report_empty_and_later_field_errors() { + let error = AccessControlRequestHeadersOwned::from_header_names(Vec::::new()).expect_err("required list"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + + let empty = TestMap::new(&FieldName::AccessControlRequestHeaders, vec![FieldValue::from_static("")]); + assert_eq!( + AccessControlRequestHeaders::view(&empty).expect_err("empty required list").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + AccessControlRequestHeaders::owned(&empty).expect_err("empty required list").kind(), + DecodeErrorKind::InvalidSyntax + ); + + let bad_later = TestMap::new( + &FieldName::AccessControlRequestHeaders, + vec![FieldValue::from_static("content-type"), FieldValue::from_static("bad name")], + ); + let error = AccessControlRequestHeaders::view(&bad_later).expect_err("bad second field line"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidToken); + assert_eq!(error.value_index(), Some(1)); + let error = AccessControlRequestHeaders::owned(&bad_later).expect_err("bad second field line"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidToken); + assert_eq!(error.value_index(), Some(1)); + + let bad_first = TestMap::new(&FieldName::AccessControlRequestHeaders, vec![FieldValue::from_static("bad name")]); + assert_eq!( + AccessControlRequestHeaders::view(&bad_first) + .expect_err("bad first field line") + .value_index(), + Some(0) + ); + assert_eq!( + AccessControlRequestHeaders::owned(&bad_first) + .expect_err("bad first field line") + .value_index(), + Some(0) + ); + + let repeated_empty = TestMap::new( + &FieldName::AccessControlRequestHeaders, + vec![FieldValue::from_static(""), FieldValue::from_static(" , ")], + ); + assert_eq!( + AccessControlRequestHeaders::owned(&repeated_empty) + .expect_err("empty repeated required list") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn private_token_helpers_cover_common_fallback_and_range_paths() { + for method in [ + b"GET".as_slice(), + b"PUT", + b"HEAD", + b"POST", + b"PATCH", + b"TRACE", + b"DELETE", + b"CONNECT", + b"OPTIONS", + ] { + assert_eq!(method_ref(method).expect("registered method").as_bytes(), method); + } + assert_eq!(method_ref(b"CUSTOM").expect("extension").as_str(), "CUSTOM"); + assert!(method_ref(b"").is_none()); + assert!(method_ref(b"bad method").is_none()); + assert!(method_ref(&[0xff]).is_none()); + + for name in [ + b"content-type".as_slice(), + b"authorization", + b"x-request-id", + b"etag", + b"origin", + b"accept", + b"x-requested-with", + b"content-length", + b"cache-control", + ] { + assert_eq!(header_name_ref(name).expect("common name").as_bytes(), name); + } + assert_eq!(header_name_ref(b"x-custom").expect("extension name").as_str(), "x-custom"); + assert!(header_name_ref(b"").is_none()); + assert!(header_name_ref(b"bad name").is_none()); + assert!(header_name_ref(&[0xff]).is_none()); + + assert_eq!(untrimmed_range(b"token"), Some(0..5)); + assert_eq!(untrimmed_range(b" token"), None); + assert_eq!(untrimmed_range(b"token "), None); + assert_eq!(untrimmed_range(b""), None); + assert_eq!(trimmed_range(b"\t token \t"), 2..7); + assert_eq!(trimmed_range(b" \t "), 3..3); + + assert_eq!(invalid_syntax(&FieldName::Accept).kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!(invalid_token(&FieldName::Accept).kind(), DecodeErrorKind::InvalidToken); + assert_eq!(invalid_number(&FieldName::Accept).kind(), DecodeErrorKind::InvalidNumber); + assert_eq!(validate_list_value(b" "), Some(false)); + + let mut sink = TestMap::new(&FieldName::AccessControlAllowHeaders, vec![FieldValue::from_static("content-type")]); + sink.remove_values(&FieldName::Accept); + assert_eq!(sink.values.len(), 1); + sink.remove_values(&FieldName::AccessControlAllowHeaders); + assert!(sink.values.is_empty()); + } +} diff --git a/crates/http_headers/src/headers/cors/test_map.rs b/crates/http_headers/src/headers/cors/test_map.rs new file mode 100644 index 000000000..55d1786db --- /dev/null +++ b/crates/http_headers/src/headers/cors/test_map.rs @@ -0,0 +1,47 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{FieldName, FieldValue}; + +pub(super) struct TestMap { + pub(super) name: &'static FieldName, + pub(super) values: Vec, +} + +impl TestMap { + pub(super) fn new(name: &'static FieldName, values: Vec) -> Self { + Self { name, values } + } +} + +impl FieldSource for TestMap { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::from_slice(name, &self.values)).flatten() + } +} + +impl FieldSink for TestMap { + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.name = name; + self.values = values.into_iter().collect(); + Ok(()) + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + if name == self.name { + self.values.extend(values); + } else { + self.name = name; + self.values = values.into_iter().collect(); + } + Ok(()) + } + + fn remove_values(&mut self, name: &'static FieldName) { + if name == self.name { + self.values.clear(); + } + } +} diff --git a/crates/http_headers/src/headers/etag.rs b/crates/http_headers/src/headers/etag.rs new file mode 100644 index 000000000..dea557237 --- /dev/null +++ b/crates/http_headers/src/headers/etag.rs @@ -0,0 +1,606 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Strong and weak entity-tag parsing and construction. + +use crate::{DecodeError, DecodeMode, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +/// Defines the `ETag` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 8.8.3](https://www.rfc-editor.org/rfc/rfc9110#section-8.8.3). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{ETag, ETagOwned}; +/// +/// let mut map = HeaderMap::new(); +/// ETag::insert(&mut map, ETagOwned::weak("revision-42")?)?; +/// assert!(ETag::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct ETag { + _private: (), +} + +/// Owned value for the `ETag` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 8.8.3]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::ETagOwned::weak("revision-42")?; +/// assert_eq!(value.opaque_tag()?, b"revision-42"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `ETag: "revision-42"` is a strong validator. +/// `ETag: W/"revision-42"` is a weak validator; `ETag: ""` is also valid. +/// +/// [RFC 9110 section 8.8.3]: https://www.rfc-editor.org/rfc/rfc9110#section-8.8.3 +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub struct ETagOwned { + value: FieldValue, +} + +/// Borrowed value for the `ETag` header. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{ETag, ETagView}; +/// use http_headers::{FieldValue, SingleValueField}; +/// +/// let wire = FieldValue::from_static(r#""revision""#); +/// let view: ETagView<'_> = ::decode_view(wire.as_field_value_ref())?; +/// assert_eq!(view.opaque_tag(), b"revision"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct ETagView<'a> { + value: FieldValueRef<'a>, +} + +impl ETagOwned { + /// Constructs a strong entity tag from unquoted opaque text. + /// + /// # Errors + /// + /// Returns an error if the opaque tag contains a byte forbidden by the + /// entity-tag grammar. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ETagOwned::strong("revision")?; + /// assert_eq!(value.opaque_tag()?, b"revision"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn strong(opaque: impl AsRef) -> Result { + construct(opaque.as_ref().as_bytes(), false) + } + + /// Constructs a weak entity tag from unquoted opaque text. + /// + /// # Errors + /// + /// Returns an error if the opaque tag contains a byte forbidden by the + /// entity-tag grammar. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ETagOwned; + /// + /// let value = ETagOwned::weak("revision")?; + /// assert!(value.is_weak()); + /// assert_eq!(value.opaque_tag()?, b"revision"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn weak(opaque: impl AsRef) -> Result { + construct(opaque.as_ref().as_bytes(), true) + } + + /// Parses a complete wire-format entity tag. + /// + /// # Errors + /// + /// Returns an error when `wire` is not a valid entity tag. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ETagOwned; + /// + /// let value = ETagOwned::try_from_wire(r#"W/"revision""#)?; + /// assert!(value.is_weak()); + /// assert_eq!(value.opaque_tag()?, b"revision"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn try_from_wire(wire: impl AsRef) -> Result { + Self::try_from(wire.as_ref()) + } + + /// Returns whether this entity tag is weak. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ETagOwned; + /// + /// let strong = ETagOwned::strong("revision")?; + /// let weak = ETagOwned::weak("revision")?; + /// assert!(!strong.is_weak()); + /// assert!(weak.is_weak()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn is_weak(&self) -> bool { + matches!(self.value.as_bytes().first(), Some(b'W' | b'w')) + } + + /// Returns the unquoted opaque tag bytes. + /// + /// # Errors + /// + /// Returns an error if the stored range and wire value disagree. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ETagOwned::strong("revision")?; + /// assert_eq!(value.opaque_tag()?, b"revision"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn opaque_tag(&self) -> Result<&[u8], DecodeError> { + let start = if self.is_weak() { 3 } else { 1 }; + let end = self + .value + .as_bytes() + .len() + .checked_sub(1) + .ok_or_else(|| super::invalid_syntax(&FieldName::Etag))?; + self.value + .as_bytes() + .get(start..end) + .ok_or_else(|| super::invalid_syntax(&FieldName::Etag)) + } + + /// Performs the strong entity-tag comparison. + /// + /// # Errors + /// + /// Returns an error if either stored range disagrees with its wire value. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ETagOwned; + /// + /// let a = ETagOwned::try_from(r#""xyzzy""#)?; + /// let b = ETagOwned::try_from(r#""xyzzy""#)?; + /// let weak = ETagOwned::try_from(r#"W/"xyzzy""#)?; + /// assert!(a.strong_eq(&b)?); + /// assert!(!a.strong_eq(&weak)?); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn strong_eq(&self, other: &Self) -> Result { + Ok(!self.is_weak() && !other.is_weak() && self.opaque_tag()? == other.opaque_tag()?) + } + + /// Performs the weak entity-tag comparison. + /// + /// # Errors + /// + /// Returns an error if either stored range disagrees with its wire value. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ETagOwned; + /// + /// let strong = ETagOwned::strong("xyzzy")?; + /// let weak = ETagOwned::weak("xyzzy")?; + /// assert!(!strong.strong_eq(&weak)?); + /// assert!(strong.weak_eq(&weak)?); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn weak_eq(&self, other: &Self) -> Result { + Ok(self.opaque_tag()? == other.opaque_tag()?) + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::FieldValue; + /// use http_headers::headers::ETagOwned; + /// + /// let value = ETagOwned::strong("revision")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value, FieldValue::from_static(r#""revision""#)); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::shared::impl_field_value_conversion!(ETagOwned, |value| value.value); + +impl<'a> ETagView<'a> { + pub(crate) const fn field_value(self) -> FieldValueRef<'a> { + self.value + } + + /// Returns whether this entity tag is weak. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ETag; + /// use http_headers::{FieldValue, SingleValueField}; + /// + /// let wire = FieldValue::from_static(r#"W/"revision""#); + /// let view = ::decode_view(wire.as_field_value_ref())?; + /// assert!(view.is_weak()); + /// assert_eq!(view.opaque_tag(), b"revision"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn is_weak(self) -> bool { + matches!(self.value.as_bytes().first(), Some(b'W' | b'w')) + } + + /// Returns the unquoted opaque tag bytes. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ETagOwned::strong("revision")?; + /// assert_eq!(value.opaque_tag()?, b"revision"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn opaque_tag(self) -> &'a [u8] { + let bytes = self.value.as_bytes(); + let start = if self.is_weak() { 3 } else { 1 }; + let (_, opaque_and_quote) = bytes.split_at(start); + let (opaque, _) = opaque_and_quote.split_at(opaque_and_quote.len() - 1); + opaque + } + + /// Performs the strong entity-tag comparison. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ETag; + /// use http_headers::{FieldValue, SingleValueField}; + /// + /// let current = FieldValue::from_static(r#""revision""#); + /// let same = FieldValue::from_static(r#""revision""#); + /// let weak = FieldValue::from_static(r#"W/"revision""#); + /// let current = ::decode_view(current.as_field_value_ref())?; + /// let same = ::decode_view(same.as_field_value_ref())?; + /// let weak = ::decode_view(weak.as_field_value_ref())?; + /// assert!(current.strong_eq(same)); + /// assert!(!current.strong_eq(weak)); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn strong_eq(self, other: Self) -> bool { + !self.is_weak() && !other.is_weak() && self.opaque_tag() == other.opaque_tag() + } + + /// Performs the weak entity-tag comparison. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ETag; + /// use http_headers::{FieldValue, SingleValueField}; + /// + /// let strong = FieldValue::from_static(r#""xyzzy""#); + /// let weak = FieldValue::from_static(r#"W/"xyzzy""#); + /// let other = FieldValue::from_static(r#""plugh""#); + /// let strong = ::decode_view(strong.as_field_value_ref())?; + /// let weak = ::decode_view(weak.as_field_value_ref())?; + /// let other = ::decode_view(other.as_field_value_ref())?; + /// assert!(strong.weak_eq(weak)); + /// assert!(!strong.weak_eq(other)); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn weak_eq(self, other: Self) -> bool { + self.opaque_tag() == other.opaque_tag() + } +} + +impl SingleValueField for ETag { + type View<'a> = ETagView<'a>; + type Owned = ETagOwned; + + fn name() -> &'static FieldName { + &FieldName::Etag + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + parse(value.as_bytes())?; + Ok(ETagView { value }) + } + + #[inline] + fn decode_owned(value: FieldValue) -> Result { + ETagOwned::try_from(value) + } + + fn decode_view_with(value: FieldValueRef<'_>, mode: DecodeMode) -> Result, DecodeError> { + parse_with(value.as_bytes(), mode)?; + Ok(ETagView { value }) + } + + fn decode_owned_with(value: FieldValue, mode: DecodeMode) -> Result { + parse_with(value.as_bytes(), mode)?; + Ok(ETagOwned { value }) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +super::shared::impl_string_conversions!(ETagOwned, &FieldName::Etag, super::invalid_syntax, value); + +impl TryFrom for ETagOwned { + type Error = DecodeError; + + #[inline] + fn try_from(value: FieldValue) -> Result { + parse(value.as_bytes())?; + Ok(Self { value }) + } +} + +fn construct(opaque: &[u8], weak: bool) -> Result { + if !opaque.iter().copied().all(valid_opaque_byte) { + return Err(super::invalid_syntax(&FieldName::Etag)); + } + let prefix: &[u8] = if weak { b"W/\"" } else { b"\"" }; + let total_len = etag_wire_len(prefix.len(), opaque.len())?; + let mut wire = Vec::with_capacity(total_len); + wire.extend_from_slice(prefix); + wire.extend_from_slice(opaque); + wire.push(b'"'); + Ok(ETagOwned { + value: super::value_from_bytes(&FieldName::Etag, wire)?, + }) +} + +#[inline] +fn etag_wire_len(prefix_len: usize, opaque_len: usize) -> Result { + prefix_len + .checked_add(opaque_len) + .and_then(|length| length.checked_add(1)) + .ok_or_else(|| super::invalid_syntax(&FieldName::Etag)) +} + +#[inline] +fn parse(bytes: &[u8]) -> Result<(), DecodeError> { + parse_with(bytes, DecodeMode::Strict) +} + +#[inline] +fn parse_with(bytes: &[u8], mode: DecodeMode) -> Result<(), DecodeError> { + let opaque = match bytes { + [b'"', opaque @ .., b'"'] | [b'W', b'/', b'"', opaque @ .., b'"'] => opaque, + [b'w', b'/', b'"', opaque @ .., b'"'] if mode == DecodeMode::Relaxed => opaque, + _ => return Err(super::invalid_syntax(&FieldName::Etag)), + }; + if opaque_is_valid(opaque) { + Ok(()) + } else { + Err(super::invalid_syntax(&FieldName::Etag)) + } +} + +/// Broadcasts `1` into every byte lane of a word. +const LANE_ONES: u64 = 0x0101_0101_0101_0101; + +/// Selects the high bit of every byte lane of a word. +const LANE_HIGH: u64 = 0x8080_8080_8080_8080; + +/// Returns whether the eight packed bytes of `word` are all legal `etagc`. +/// +/// A `FieldValue` byte other than DEL is legal when setting bit one lifts it +/// to `0x23` or above: HTAB becomes `0x0B`, SP and DQUOTE both become `0x22`, +/// `!` becomes `0x23`, and every other byte already exceeds `0x22`. +/// +/// `ored | LANE_HIGH` lifts every lane above the threshold so the subtraction +/// never borrows across lane boundaries, and masking with the untouched high +/// bits keeps `obs-text` lanes out of the comparison. XOR maps DEL to zero; +/// the standard zero-byte test rejects it in any lane. +const fn word_is_valid(word: u64) -> bool { + let ored = word | (0x02 * LANE_ONES); + let lifted = ored | LANE_HIGH; + let below_minimum = !(lifted.wrapping_sub(0x23 * LANE_ONES) | ored) & LANE_HIGH; + let deltas = word ^ (0x7f * LANE_ONES); + let contains_del = deltas.wrapping_sub(LANE_ONES) & !deltas & LANE_HIGH; + below_minimum == 0 && contains_del == 0 +} + +/// Returns whether every byte of `opaque` is a legal `etagc`. +/// +/// Words of eight bytes are classified at once; the final word overlaps the +/// previous one when the length is not a multiple of eight, which re-checks a +/// few bytes rather than paying per-byte for the tail. +#[inline] +pub(super) fn opaque_is_valid(opaque: &[u8]) -> bool { + let length = opaque.len(); + if length < 8 { + return opaque.iter().copied().all(valid_opaque_byte); + } + let mut index = 0; + while index + 8 < length { + if !word_is_valid(read_word(opaque, index)) { + return false; + } + index += 8; + } + word_is_valid(read_word(opaque, length - 8)) +} + +/// Reads the eight bytes of `bytes` starting at `index` as a packed word. +fn read_word(bytes: &[u8], index: usize) -> u64 { + let mut word = [0; 8]; + word.copy_from_slice(&bytes[index..index + 8]); + u64::from_ne_bytes(word) +} + +const fn valid_opaque_byte(byte: u8) -> bool { + matches!(byte, b'!' | b'#'..=b'~' | 0x80..=0xff) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #![expect( + clippy::assertions_on_result_states, + reason = "tests classify parser outcomes without needing successful values" + )] + + use super::{ETag, ETagOwned, etag_wire_len, opaque_is_valid, parse_with, read_word, valid_opaque_byte, word_is_valid}; + use crate::sink::FieldSink; + use crate::{DecodeMode, FieldName, FieldValue, SingleValueField, TestSink}; + + #[test] + fn owned_and_borrowed_tags_cover_accessors_and_comparisons() { + let strong = ETagOwned::strong("revision").expect("valid strong tag"); + let strong_same = ETagOwned::try_from_wire("\"revision\"").expect("valid wire tag"); + let weak = ETagOwned::weak("revision").expect("valid weak tag"); + assert!(!strong.is_weak()); + assert!(weak.is_weak()); + assert_eq!(strong.opaque_tag(), Ok(b"revision".as_slice())); + assert_eq!(strong.strong_eq(&strong_same), Ok(true)); + assert_eq!(strong.strong_eq(&weak), Ok(false)); + assert_eq!(strong.weak_eq(&weak), Ok(true)); + assert_eq!( + String::from("\"owned\"") + .parse::() + .expect("owned string parses") + .opaque_tag(), + Ok(b"owned".as_slice()) + ); + assert_eq!( + "\"parsed\"".parse::().expect("shared FromStr parses").opaque_tag(), + Ok(b"parsed".as_slice()) + ); + + let mut table = TestSink::new(); + ETag::insert(&mut table, strong.clone()).expect("table accepts etag"); + let view = ETag::view(&table).expect("valid tag").expect("present"); + let same = ::decode_view(view.field_value()).expect("valid borrowed tag"); + let weak_view = ::decode_view(weak.value.as_field_value_ref()).expect("valid weak borrowed tag"); + assert!(!view.is_weak()); + assert_eq!(view.opaque_tag(), b"revision"); + assert!(view.strong_eq(same)); + assert!(!view.strong_eq(weak_view)); + assert!(view.weak_eq(weak_view)); + assert_eq!( + ETag::owned(&table).expect("valid owned tag").expect("present").opaque_tag(), + Ok(b"revision".as_slice()) + ); + table.remove_values(&FieldName::Etag); + assert!(ETag::view(&table).expect("absence is valid").is_none()); + assert_eq!(strong.into_field_value(), FieldValue::from_static("\"revision\"")); + } + + #[test] + fn parsing_covers_relaxed_and_invalid_opaque_bytes() { + assert!(parse_with(b"w/\"tag\"", DecodeMode::Strict).is_err()); + assert!(parse_with(b"w/\"tag\"", DecodeMode::Relaxed).is_ok()); + for wire in [ + b"tag".as_slice(), + b"W/tag", + b"\"space inside\"", + b"\"tab\tinside\"", + b"\"del\x7finside\"", + b"\"unterminated", + ] { + assert!(parse_with(wire, DecodeMode::Relaxed).is_err(), "{wire:?}"); + } + assert!(ETagOwned::strong("bad\"tag").is_err()); + assert!(ETagOwned::weak("bad tag").is_err()); + + let relaxed = FieldValue::from_static("w/\"tag\""); + assert!(::decode_view_with(relaxed.as_field_value_ref(), DecodeMode::Relaxed).is_ok()); + assert!(::decode_owned_with(relaxed, DecodeMode::Relaxed).is_ok()); + + let strict = FieldValue::from_static("\"tag\""); + let strict_view = ::decode_view(strict.as_field_value_ref()).expect("valid strict borrowed tag"); + assert_eq!(strict_view.opaque_tag(), b"tag"); + let strict_owned = ::decode_owned(strict.clone()).expect("valid strict owned tag"); + assert_eq!(::as_field_value(&strict_owned), &strict); + assert_eq!( + ETagOwned::try_from(String::from("\"owned\"")) + .expect("valid owned string") + .opaque_tag(), + Ok(b"owned".as_slice()) + ); + assert!(ETagOwned::try_from("\n").is_err()); + assert!(ETagOwned::try_from(String::from("\n")).is_err()); + assert!(ETagOwned::try_from(FieldValue::from_static("invalid")).is_err()); + + let malformed = ETagOwned { + value: FieldValue::from_static(""), + }; + assert!(malformed.opaque_tag().is_err()); + assert!(malformed.strong_eq(&strict_owned).is_err()); + assert!(malformed.weak_eq(&strict_owned).is_err()); + let malformed_range = ETagOwned { + value: FieldValue::from_static("\""), + }; + assert!(malformed_range.opaque_tag().is_err()); + assert_eq!(etag_wire_len(1, 2), Ok(4)); + assert!(etag_wire_len(usize::MAX, 0).is_err()); + assert!(etag_wire_len(usize::MAX - 1, 1).is_err()); + } + + #[test] + fn opaque_validator_covers_scalar_words_tails_and_rejections() { + assert!(opaque_is_valid(b"")); + assert!(opaque_is_valid(b"short")); + assert!(opaque_is_valid(b"12345678")); + assert!(opaque_is_valid(b"123456789")); + assert!(opaque_is_valid(b"12345678901234567")); + assert!(!opaque_is_valid(b"1234567 901234567")); + assert!(!opaque_is_valid(b"1234567890123456 ")); + assert!(!opaque_is_valid(b"short\x7f")); + assert!(!opaque_is_valid(b"1234\x7f678")); + assert!(word_is_valid(read_word(b"12345678", 0))); + assert!(!word_is_valid(read_word(b"1234 678", 0))); + assert!(!word_is_valid(read_word(b"1234\x7f678", 0))); + assert!(valid_opaque_byte(b'!')); + assert!(valid_opaque_byte(0x80)); + assert!(!valid_opaque_byte(b' ')); + + for byte in u8::MIN..=u8::MAX { + assert_eq!(opaque_is_valid(&[byte]), valid_opaque_byte(byte), "scalar byte {byte:#04x}"); + for lane in 0..8 { + let mut bytes = [b'a'; 8]; + bytes[lane] = byte; + assert_eq!( + word_is_valid(u64::from_ne_bytes(bytes)), + valid_opaque_byte(byte), + "packed byte {byte:#04x} in lane {lane}" + ); + } + } + } +} diff --git a/crates/http_headers/src/headers/extension_value.rs b/crates/http_headers/src/headers/extension_value.rs new file mode 100644 index 000000000..5e6dea68f --- /dev/null +++ b/crates/http_headers/src/headers/extension_value.rs @@ -0,0 +1,11 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +/// Distinguishes a flag-style extension from one carrying a value. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub enum ExtensionValue<'a> { + /// Emits only the extension name. + Flag, + /// Emits the extension name followed by `=` and this value. + Value(&'a str), +} diff --git a/crates/http_headers/src/headers/field_name_view.rs b/crates/http_headers/src/headers/field_name_view.rs new file mode 100644 index 000000000..06e76ca95 --- /dev/null +++ b/crates/http_headers/src/headers/field_name_view.rs @@ -0,0 +1,140 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +use std::hash::{Hash, Hasher}; + +use crate::headers::tokens::header_name_text; +use crate::{FieldName, InvalidFieldName, validate}; + +/// A borrowed field-name token with ASCII-case-insensitive equality and hashing. +/// +/// The original spelling is preserved. Unlike the owned [`FieldName`], this +/// view does not impose a length limit or allocate for extension names. +/// +/// Available with either the `headers-cors` or `headers-negotiation` feature. +/// +/// # Examples +/// +/// ``` +/// use http_headers::headers::FieldNameView; +/// +/// let name = FieldNameView::new("X-Trace-Id")?; +/// assert_eq!(name, FieldNameView::new("x-trace-id")?); +/// assert_eq!(name.as_str(), "X-Trace-Id"); +/// assert_eq!(name.try_to_field_name()?.as_str(), "x-trace-id"); +/// # Ok::<(), http_headers::InvalidFieldName>(()) +/// ``` +#[derive(Clone, Copy)] +pub struct FieldNameView<'a>(&'a [u8]); + +impl<'a> FieldNameView<'a> { + /// Validates and borrows a field-name token without allocating. + /// + /// # Errors + /// + /// Returns [`InvalidFieldName`] for empty input or non-token bytes. + #[inline] + pub fn new(name: &'a str) -> Result { + if validate::token(name.as_bytes()) { + Ok(Self(name.as_bytes())) + } else { + Err(InvalidFieldName) + } + } + + #[inline] + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + Self(bytes) + } + + /// Returns the original spelling. + #[must_use] + #[inline] + pub fn as_str(self) -> &'a str { + header_name_text(self.0) + } + + /// Returns the original bytes. + #[must_use] + #[inline] + pub const fn as_bytes(self) -> &'a [u8] { + self.0 + } + + /// Compares to another spelling without allocating or normalizing storage. + #[must_use] + #[inline] + pub fn eq_ignore_ascii_case(self, other: &str) -> bool { + self.0.eq_ignore_ascii_case(other.as_bytes()) + } + + /// Materializes an owned field name, normalizing its spelling. + /// + /// Known names do not allocate; extension names may allocate. + /// + /// # Errors + /// + /// Returns [`InvalidFieldName`] if the token exceeds the owned type's + /// 65,535-byte length limit. + #[inline] + pub fn try_to_field_name(self) -> Result { + FieldName::try_from_bytes(self.as_bytes()) + } + + /// Materializes an external field name, normalizing its spelling. + /// + /// # Errors + /// + /// Returns the external constructor's error if its constraints reject + /// the token. This conversion may allocate. + #[cfg(feature = "http")] + #[inline] + pub fn try_to_http_header_name(self) -> Result { + http::HeaderName::from_bytes(self.as_bytes()) + } +} + +impl<'a> TryFrom<&'a str> for FieldNameView<'a> { + type Error = InvalidFieldName; + + fn try_from(value: &'a str) -> Result { + Self::new(value) + } +} + +impl PartialEq for FieldNameView<'_> { + #[inline] + fn eq(&self, other: &Self) -> bool { + self.0.eq_ignore_ascii_case(other.0) + } +} + +impl Eq for FieldNameView<'_> {} + +impl Hash for FieldNameView<'_> { + fn hash(&self, state: &mut H) { + self.0.len().hash(state); + for byte in self.0 { + state.write_u8(byte.to_ascii_lowercase()); + } + } +} + +impl AsRef for FieldNameView<'_> { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl fmt::Display for FieldNameView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } +} + +impl fmt::Debug for FieldNameView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_tuple("FieldNameView").field(&self.as_str()).finish() + } +} diff --git a/crates/http_headers/src/headers/invalid_method.rs b/crates/http_headers/src/headers/invalid_method.rs new file mode 100644 index 000000000..8c8907648 --- /dev/null +++ b/crates/http_headers/src/headers/invalid_method.rs @@ -0,0 +1,17 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::error::Error; +use std::fmt; + +/// A method was empty or contained a byte outside the HTTP token grammar. +#[derive(Clone, Copy, Debug, Default, Eq, Hash, PartialEq)] +pub struct InvalidMethod; + +impl fmt::Display for InvalidMethod { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("invalid HTTP method") + } +} + +impl Error for InvalidMethod {} diff --git a/crates/http_headers/src/headers/location.rs b/crates/http_headers/src/headers/location.rs new file mode 100644 index 000000000..013724c25 --- /dev/null +++ b/crates/http_headers/src/headers/location.rs @@ -0,0 +1,622 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Validated URI-reference support for the `Location` header. + +use std::hash::{Hash, Hasher}; +use std::{fmt, str}; + +use fluent_uri::Uri; + +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +mod component; +mod construction; +mod metadata; +mod uri_authority; +mod uri_reference; + +use metadata::ComponentRanges; +pub use uri_authority::UriAuthority; +pub use uri_reference::UriReference; + +/// Defines the `Location` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 10.2.2](https://www.rfc-editor.org/rfc/rfc9110#section-10.2.2) +/// using the URI-reference grammar from +/// [RFC 3986 section 4](https://www.rfc-editor.org/rfc/rfc3986#section-4). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{Location, LocationOwned}; +/// +/// let mut map = HeaderMap::new(); +/// Location::insert(&mut map, LocationOwned::try_from("/next")?)?; +/// assert!(Location::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct Location { + _private: (), +} + +/// Owned value for the `Location` header. +/// +/// # Specification +/// +/// The field is defined by [RFC 9110 section 10.2.2] and its URI-reference +/// grammar by [RFC 3986 section 4]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::LocationOwned::try_from("../people?tab=1#profile")?; +/// assert_eq!(value.as_str()?, "../people?tab=1#profile"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Location: https://example.com/people` is absolute, +/// `Location: /accounts/12345` is relative, and `Location: #profile` is a +/// fragment-only reference. An empty field value is also a valid reference. +/// +/// [RFC 9110 section 10.2.2]: https://www.rfc-editor.org/rfc/rfc9110#section-10.2.2 +/// [RFC 3986 section 4]: https://www.rfc-editor.org/rfc/rfc3986#section-4 +#[derive(Clone)] +pub struct LocationOwned { + value: FieldValue, + component_ranges: ComponentRanges, + normalized: Option, +} + +/// Borrowed value for the `Location` header. +/// +/// Ordinary views borrow without allocation. A relaxed value containing +/// backslashes also retains an owned, slash-normalized semantic spelling. +/// Cloning such a view clones that buffer; raw access always borrows the +/// original source and structured access borrows this view. +#[derive(Clone)] +/// # Examples +/// +/// ``` +/// use http_headers::headers::{Location, LocationView}; +/// use http_headers::{FieldValue, SingleValueField}; +/// +/// let field = FieldValue::from_static("/next"); +/// let view: LocationView<'_> = +/// ::decode_view(field.as_field_value_ref())?; +/// assert_eq!(view.as_str()?, "/next"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct LocationView<'a> { + text: &'a str, + component_ranges: ComponentRanges, + normalized: Option, +} + +impl fmt::Debug for LocationOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("LocationOwned").field("redacted", &true).finish_non_exhaustive() + } +} + +impl fmt::Debug for LocationView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("LocationView").field("redacted", &true).finish_non_exhaustive() + } +} + +impl LocationOwned { + /// Constructs a URI-reference from validated authority and encoded components. + /// + /// Components are not percent-encoded automatically. A missing component is + /// distinct from `Some("")`. The authority must already satisfy its grammar, + /// while the remaining components and their contextual constraints are + /// validated here. The assembled reference is not parsed again. + /// + /// # Errors + /// + /// Returns [`crate::DecodeErrorKind::InvalidSyntax`] for invalid components + /// or combinations: an authority requires an empty or slash-prefixed path; + /// without an authority the path cannot start with `//`; without a scheme, + /// a relative path's first segment cannot contain a colon. + /// + /// # Examples + /// + /// ``` + /// use http_headers::headers::{LocationOwned, UriAuthority}; + /// + /// let authority = UriAuthority::new(None, "example.com", Some("443"))?; + /// let location = + /// LocationOwned::from_components(Some("https"), Some(authority), "/a%2Fb", Some(""), None)?; + /// assert_eq!(location.as_str()?, "https://example.com:443/a%2Fb?"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_components( + scheme: Option<&str>, + authority: Option>, + path: &str, + query: Option<&str>, + fragment: Option<&str>, + ) -> Result { + construction::from_components(scheme, authority, path, query, fragment) + } + + /// Returns retained URI components without repeating URI grammar validation. + /// + /// In relaxed mode these describe the backslash-normalized spelling, while + /// [`Self::as_str`] and forwarding retain the original wire spelling. + #[inline] + #[must_use] + pub fn uri_reference(&self) -> UriReference<'_> { + UriReference::from_component_ranges(self.semantic_text(), &self.component_ranges) + } + + fn semantic_text(&self) -> &str { + self.normalized.as_deref().unwrap_or_else(|| { + self.as_str() + .expect("all Location constructors validate UTF-8 and storage is immutable") + }) + } + + /// Whether relaxed decoding replaced backslashes in the semantic spelling. + #[inline] + #[must_use] + pub const fn was_normalized(&self) -> bool { + self.normalized.is_some() + } + + /// Returns the URI-reference bytes. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::LocationOwned; + /// + /// let value = LocationOwned::try_from("/next")?; + /// assert_eq!(value.as_bytes(), b"/next"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_bytes(&self) -> &[u8] { + self.value.as_bytes() + } + + /// Returns the URI-reference as UTF-8. + /// # Errors + /// + /// Returns an error if the stored wire value is unexpectedly non-UTF-8. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::LocationOwned::try_from("/next")?; + /// assert_eq!(value.as_str()?, "/next"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_str(&self) -> Result<&str, DecodeError> { + str::from_utf8(self.value.as_bytes()).map_err(|_invalid| invalid()) + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::LocationOwned; + /// + /// let value = LocationOwned::try_from("https://example.com/people")?; + /// let field = value.into_field_value(); + /// assert_eq!(field.as_bytes(), b"https://example.com/people"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::shared::impl_field_value_conversion!(LocationOwned, |value| value.value); + +impl<'a> LocationView<'a> { + /// Returns retained components, borrowing any normalized backing from this view. + /// + /// Ordinary decoding borrows the source without allocation. Relaxed + /// backslash normalization retains one owned semantic buffer; reading these + /// components neither allocates nor normalizes again. + #[inline] + #[must_use] + pub fn uri_reference(&self) -> UriReference<'_> { + UriReference::from_component_ranges(self.normalized.as_deref().unwrap_or(self.text), &self.component_ranges) + } + + /// Whether relaxed decoding replaced backslashes in the semantic spelling. + #[inline] + #[must_use] + pub const fn was_normalized(&self) -> bool { + self.normalized.is_some() + } + + /// Returns the URI-reference bytes. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::Location; + /// use http_headers::{FieldValue, SingleValueField}; + /// + /// let field = FieldValue::from_static("../people?tab=1#profile"); + /// let view = ::decode_view(field.as_field_value_ref())?; + /// assert_eq!(view.as_bytes(), b"../people?tab=1#profile"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_bytes(&self) -> &'a [u8] { + self.text.as_bytes() + } + + /// Returns the URI-reference as UTF-8. + /// + /// # Errors + /// + /// This validated view always returns `Ok`; the fallible signature is + /// retained for API compatibility with the owned form. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::LocationOwned::try_from("/next")?; + /// assert_eq!(value.as_str()?, "/next"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::unnecessary_wraps, + reason = "the fallible signature intentionally matches the owned representation" + )] + pub const fn as_str(&self) -> Result<&'a str, DecodeError> { + Ok(self.text) + } +} + +impl SingleValueField for Location { + type View<'a> = LocationView<'a>; + type Owned = LocationOwned; + + fn name() -> &'static FieldName { + &FieldName::Location + } + + #[inline] + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + validate(value.as_bytes()) + } + + #[inline] + fn decode_owned(value: FieldValue) -> Result { + LocationOwned::try_from(value) + } + + fn decode_view_with(value: FieldValueRef<'_>, mode: crate::DecodeMode) -> Result, DecodeError> { + validate_with(value.as_bytes(), mode) + } + + fn decode_owned_with(mut value: FieldValue, mode: crate::DecodeMode) -> Result { + let view = validate_with(value.as_bytes(), mode)?; + let component_ranges = view.component_ranges; + let normalized = view.normalized; + value.set_sensitive(true); + Ok(LocationOwned { + value, + component_ranges, + normalized, + }) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +impl TryFrom<&str> for LocationOwned { + type Error = DecodeError; + + fn try_from(value: &str) -> Result { + let mut value = FieldValue::from_str(value).map_err(|_invalid| super::invalid_syntax(&FieldName::Location))?; + value.set_sensitive(true); + Self::try_from(value) + } +} + +impl TryFrom for LocationOwned { + type Error = DecodeError; + + fn try_from(value: String) -> Result { + let mut value = FieldValue::try_from(value).map_err(|_invalid| super::invalid_syntax(&FieldName::Location))?; + value.set_sensitive(true); + Self::try_from(value) + } +} + +impl TryFrom for LocationOwned { + type Error = DecodeError; + + fn try_from(mut value: FieldValue) -> Result { + let component_ranges = validate(value.as_bytes())?.component_ranges; + value.set_sensitive(true); + Ok(Self { + value, + component_ranges, + normalized: None, + }) + } +} + +#[inline] +fn validate(bytes: &[u8]) -> Result, DecodeError> { + if let Some(text) = is_simple_reference(bytes) { + return Ok(LocationView { + text, + component_ranges: ComponentRanges::from_simple(text), + normalized: None, + }); + } + let text = str::from_utf8(bytes).map_err(|_invalid| invalid())?; + let component_ranges = validate_general_reference(text)?; + Ok(LocationView { + text, + component_ranges, + normalized: None, + }) +} + +fn validate_with(bytes: &[u8], mode: crate::DecodeMode) -> Result, DecodeError> { + if mode == crate::DecodeMode::Strict { + return validate(bytes); + } + if let Some(text) = is_simple_reference(bytes) { + return Ok(LocationView { + text, + component_ranges: ComponentRanges::from_simple(text), + normalized: None, + }); + } + let text = str::from_utf8(bytes).map_err(|_invalid| invalid())?; + if let Ok(component_ranges) = validate_general_reference(text) { + return Ok(LocationView { + text, + component_ranges, + normalized: None, + }); + } + if !text.contains('\\') { + return Err(invalid()); + } + let normalized = text.replace('\\', "/"); + let component_ranges = validate_general_reference(&normalized)?; + Ok(LocationView { + text, + component_ranges, + normalized: Some(normalized), + }) +} + +/// Validates everything the origin-relative scanner declines to recognize. +/// +/// The recognized subset already covers the shapes a redirect normally +/// carries, so keeping the RFC 3986 parse cold and behind a call leaves +/// [`validate`] small enough to fold into its callers without changing +/// validation semantics. +#[cold] +#[inline(never)] +fn validate_general_reference(value: &str) -> Result { + validate_general_reference_len(value.len())?; + Uri::parse(value) + .map(|parsed| ComponentRanges::from_parsed(&parsed)) + .map_err(|_invalid| invalid()) +} + +#[inline] +fn validate_general_reference_len(length: usize) -> Result<(), DecodeError> { + if i32::try_from(length).is_ok() { Ok(()) } else { Err(invalid()) } +} + +/// Accepts the references that need no RFC 3986 parse. +/// +/// The scanner recognizes a `path-absolute` with an optional query and one +/// optional fragment written only with `unreserved`, `sub-delims`, `:`, `@`, +/// `/`, and `?`, and the same shape behind a `scheme "://" host [":" port]` +/// prefix. Anything else, including every percent escape, userinfo, and +/// IP-literal, returns `false` and is handed to the full parser unchanged. +fn is_simple_reference(bytes: &[u8]) -> Option<&str> { + http_headers_simd::as_simple_uri_reference(bytes) +} + +fn invalid() -> DecodeError { + super::invalid_syntax(&FieldName::Location) +} + +impl PartialEq for LocationOwned { + fn eq(&self, other: &Self) -> bool { + self.value == other.value + } +} + +impl Eq for LocationOwned {} + +impl Hash for LocationOwned { + fn hash(&self, state: &mut H) { + self.value.hash(state); + } +} + +impl PartialEq for LocationView<'_> { + fn eq(&self, other: &Self) -> bool { + self.text == other.text + } +} + +impl Eq for LocationView<'_> {} + +impl Hash for LocationView<'_> { + fn hash(&self, state: &mut H) { + self.text.hash(state); + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #![expect( + clippy::assertions_on_result_states, + reason = "tests classify parser outcomes without needing successful values" + )] + + use super::{Location, LocationOwned, validate, validate_general_reference, validate_general_reference_len, validate_with}; + use crate::sink::FieldSink; + use crate::{DecodeErrorKind, DecodeMode, FieldName, FieldValue, SingleValueField, TestSink}; + + #[test] + #[should_panic(expected = "all Location constructors validate UTF-8 and storage is immutable")] + fn owned_semantic_projection_checks_its_utf8_invariant() { + let malformed = LocationOwned { + value: FieldValue::try_from(vec![0xff]).unwrap(), + component_ranges: validate(b"").unwrap().component_ranges, + normalized: None, + }; + assert_eq!(malformed.as_str().unwrap_err().kind(), DecodeErrorKind::InvalidSyntax); + let _uri = malformed.uri_reference(); + } + + #[test] + fn constructors_accessors_debug_and_round_trip() { + for wire in ["", "/next?tab=1#profile", "../next", "https://example.com/a"] { + let owned = LocationOwned::try_from(wire).expect("valid URI reference"); + assert_eq!(owned.as_bytes(), wire.as_bytes()); + assert_eq!(owned.as_str(), Ok(wire)); + assert!(format!("{owned:?}").contains("redacted")); + assert!(owned.clone().into_field_value().is_sensitive()); + + let mut table = TestSink::new(); + Location::insert(&mut table, owned).expect("table accepts location"); + let view = Location::view(&table).expect("valid location").expect("present"); + assert_eq!(view.as_bytes(), wire.as_bytes()); + assert_eq!(view.as_str(), Ok(wire)); + assert!(format!("{view:?}").contains("redacted")); + assert!( + Location::owned(&table) + .expect("valid owned location") + .expect("present") + .into_field_value() + .is_sensitive() + ); + table.remove_values(&FieldName::Location); + assert!(Location::view(&table).expect("absence is valid").is_none()); + } + + assert_eq!( + LocationOwned::try_from(String::from("/owned")).expect("valid owned URI").as_str(), + Ok("/owned") + ); + assert_eq!( + "/parsed".parse::().expect("shared FromStr implementation").as_str(), + Ok("/parsed") + ); + } + + /// The fast path replaces the RFC 3986 parse outright, so anything it + /// accepts must be something the general parser would also accept. + #[test] + fn the_recognized_subset_agrees_with_the_general_parser() { + let alphabet = b"aZ0:@/?#%[].+-_~!$&'()*,;=\\ \t"; + let radix = alphabet.len(); + let case_count = radix.pow(if cfg!(miri) { 2 } else { 4 }); + for encoded in 0..case_count { + let mut value = encoded; + let mut indices = [0; 4]; + for index in &mut indices { + *index = value % radix; + value /= radix; + } + if cfg!(miri) { + // Pairwise seeds still place every alphabet byte in every position. + indices[2] = (indices[0] + indices[1]) % radix; + indices[3] = (indices[0] + 2 * indices[1]) % radix; + } + let bytes = indices.map(|index| alphabet[index]); + for prefix in [b"".as_slice(), b"https://h.example".as_slice()] { + let mut candidate = prefix.to_vec(); + candidate.extend_from_slice(&bytes); + let Some(text) = super::is_simple_reference(&candidate) else { + continue; + }; + assert_eq!( + super::ComponentRanges::from_simple(text), + validate_general_reference(text).unwrap(), + "recognized {text:?} with different components" + ); + } + } + } + + #[test] + fn strict_and_relaxed_validation_cover_fallbacks() { + assert!(validate(b"/simple/path?query#fragment").is_ok()); + assert!(validate(b"https://example.com/docs?x=1#top").is_ok()); + assert!(validate(b"https://example.com:8443/docs").is_ok()); + assert!(validate(b"https://user@example.com/docs").is_ok()); + assert!(validate(b"https://example.com/caf%C3%A9").is_ok()); + assert!(validate(b"https://example.com/a b").is_err()); + assert!(validate_general_reference("mailto:user@example.com").is_ok()); + assert!(validate(b"%zz").is_err()); + assert!(validate(b"has space").is_err()); + assert!(validate(&[0xff]).is_err()); + + assert!(validate_with(b"/a\\b", DecodeMode::Strict).is_err()); + assert_eq!(validate_with(b"/simple", DecodeMode::Relaxed).unwrap().as_str(), Ok("/simple")); + assert_eq!(validate_with(b"/a\\b", DecodeMode::Relaxed).unwrap().as_str(), Ok("/a\\b")); + assert!(validate_with(b"bad\\%zz", DecodeMode::Relaxed).is_err()); + assert!(validate_with(b"%zz", DecodeMode::Relaxed).is_err()); + assert!(validate_with(&[0xff], DecodeMode::Relaxed).is_err()); + + let relaxed = FieldValue::from_static("/a\\b"); + assert!(::decode_view(relaxed.as_field_value_ref()).is_err()); + assert!(::decode_view_with(relaxed.as_field_value_ref(), DecodeMode::Relaxed).is_ok()); + let owned = ::decode_owned_with(relaxed, DecodeMode::Relaxed).expect("relaxed owned location"); + assert!(owned.into_field_value().is_sensitive()); + + let strict = FieldValue::from_static("/strict"); + let strict_view = ::decode_view(strict.as_field_value_ref()).expect("valid strict location"); + assert_eq!(strict_view.as_str(), Ok("/strict")); + let strict_owned = ::decode_owned(strict.clone()).expect("valid strict owned location"); + assert_eq!(::as_field_value(&strict_owned), &strict); + + let invalid = FieldValue::try_from(vec![0xff]).expect("obs-text field value"); + assert_eq!( + LocationOwned::try_from(invalid.clone()) + .expect_err("location requires UTF-8") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert!(LocationOwned::try_from("\n").is_err()); + assert!(LocationOwned::try_from(String::from("\n")).is_err()); + assert!(validate_general_reference_len(i32::MAX as usize).is_ok()); + assert!(validate_general_reference_len(i32::MAX as usize + 1).is_err()); + let malformed = LocationOwned { + value: invalid, + component_ranges: validate(b"").unwrap().component_ranges, + normalized: None, + }; + assert_eq!( + malformed.as_str().expect_err("malformed private storage is not UTF-8").kind(), + DecodeErrorKind::InvalidSyntax + ); + } +} diff --git a/crates/http_headers/src/headers/location/component.rs b/crates/http_headers/src/headers/location/component.rs new file mode 100644 index 000000000..269a27f08 --- /dev/null +++ b/crates/http_headers/src/headers/location/component.rs @@ -0,0 +1,57 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Encoded component grammar used by validated construction. + +use super::invalid; +use crate::DecodeError; + +// fluent-uri has no stable component validators. Checking these RFC 3986 +// productions directly avoids assembling a URI just to parse it again. +#[derive(Clone, Copy, Debug)] +pub(super) enum Component { + Scheme, + Userinfo, + RegisteredName, + Path, + QueryFragment, + IpvFuture, +} + +impl Component { + pub(super) fn validate(self, text: &str) -> Result<(), DecodeError> { + let mut bytes = text.bytes(); + while let Some(byte) = bytes.next() { + if byte == b'%' && !matches!(self, Self::Scheme | Self::IpvFuture) { + if !bytes.next().is_some_and(|byte| byte.is_ascii_hexdigit()) || !bytes.next().is_some_and(|byte| byte.is_ascii_hexdigit()) + { + return Err(invalid()); + } + } else if !self.allows(byte) { + return Err(invalid()); + } + } + Ok(()) + } + + fn allows(self, byte: u8) -> bool { + if byte.is_ascii_alphanumeric() { + return true; + } + if matches!(self, Self::Scheme) { + return matches!(byte, b'+' | b'-' | b'.'); + } + if matches!( + byte, + b'-' | b'.' | b'_' | b'~' | b'!' | b'$' | b'&' | b'\'' | b'(' | b')' | b'*' | b'+' | b',' | b';' | b'=' + ) { + return true; + } + match self { + Self::Userinfo | Self::IpvFuture => byte == b':', + Self::Path => matches!(byte, b':' | b'@' | b'/'), + Self::QueryFragment => matches!(byte, b':' | b'@' | b'/' | b'?'), + Self::Scheme | Self::RegisteredName => false, + } + } +} diff --git a/crates/http_headers/src/headers/location/construction.rs b/crates/http_headers/src/headers/location/construction.rs new file mode 100644 index 000000000..97b301afd --- /dev/null +++ b/crates/http_headers/src/headers/location/construction.rs @@ -0,0 +1,158 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Component validation and assembly without parsing the assembled reference again. + +use std::num::NonZeroUsize; + +use super::component::Component; +use super::{ComponentRanges, LocationOwned, UriAuthority, invalid, is_simple_reference, validate_general_reference_len}; +use crate::{DecodeError, FieldValue}; + +pub(super) fn from_components( + scheme: Option<&str>, + authority: Option>, + path: &str, + query: Option<&str>, + fragment: Option<&str>, +) -> Result { + validate_components(scheme, authority, path, query, fragment)?; + + let mut length = path.len(); + for (component, separator_length) in [ + (scheme, 1), + (authority.map(UriAuthority::host), 2), + (authority.and_then(UriAuthority::userinfo), 1), + (authority.and_then(UriAuthority::port), 1), + (query, 1), + (fragment, 1), + ] { + if let Some(component) = component { + length = length + .checked_add(component.len()) + .and_then(|length| length.checked_add(separator_length)) + .ok_or_else(invalid)?; + } + } + let mut text = String::with_capacity(length); + let scheme_end = scheme.and_then(|scheme| { + text.push_str(scheme); + let end = NonZeroUsize::new(text.len()); + text.push(':'); + end + }); + let authority_host = authority.map(|authority| { + text.push_str("//"); + if let Some(userinfo) = authority.userinfo() { + text.push_str(userinfo); + text.push('@'); + } + let host_start = text.len(); + text.push_str(authority.host()); + let host_end = text.len(); + if let Some(port) = authority.port() { + text.push(':'); + text.push_str(port); + } + host_start..host_end + }); + let path_start = text.len(); + text.push_str(path); + let path_end = text.len(); + let query_end = query.and_then(|query| { + text.push('?'); + text.push_str(query); + NonZeroUsize::new(text.len()) + }); + let fragment_start = fragment.and_then(|fragment| { + text.push('#'); + let start = NonZeroUsize::new(text.len()); + text.push_str(fragment); + start + }); + + validate_constructed_reference_len(text.len(), text.as_bytes())?; + let mut value = FieldValue::try_from(text).map_err(|_invalid| invalid())?; + value.set_sensitive(true); + Ok(LocationOwned { + value, + component_ranges: ComponentRanges { + scheme_end, + authority_host, + path: path_start..path_end, + query_end, + fragment_start, + }, + normalized: None, + }) +} + +fn validate_constructed_reference_len(length: usize, bytes: &[u8]) -> Result<(), DecodeError> { + // Only general references have the dependency's existing i32 length bound. + if i32::try_from(length).is_err() && is_simple_reference(bytes).is_none() { + validate_general_reference_len(length)?; + } + Ok(()) +} + +fn validate_components( + scheme: Option<&str>, + authority: Option>, + path: &str, + query: Option<&str>, + fragment: Option<&str>, +) -> Result<(), DecodeError> { + if let Some(scheme) = scheme { + if !scheme.as_bytes().first().is_some_and(u8::is_ascii_alphabetic) { + return Err(invalid()); + } + Component::Scheme.validate(scheme)?; + } + Component::Path.validate(path)?; + if authority.is_some() { + if !path.is_empty() && !path.starts_with('/') { + return Err(invalid()); + } + } else if path.starts_with("//") + || (scheme.is_none() && !path.starts_with('/') && path.split('/').next().is_some_and(|segment| segment.contains(':'))) + { + return Err(invalid()); + } + for component in [query, fragment].into_iter().flatten() { + Component::QueryFragment.validate(component)?; + } + Ok(()) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::validate_constructed_reference_len; + use crate::{DecodeErrorKind, FieldName}; + + #[test] + fn constructed_references_accept_the_general_length_boundary() { + let maximum = usize::try_from(i32::MAX).unwrap(); + for length in [0, 1, maximum - 1, maximum] { + for bytes in [b"/next?x#y".as_slice(), b"../next", b"/next%20path"] { + assert_eq!(validate_constructed_reference_len(length, bytes), Ok(())); + } + } + } + + #[test] + fn oversized_constructed_references_only_accept_simple_shapes() { + let maximum = usize::try_from(i32::MAX).unwrap(); + for length in [maximum + 1, usize::MAX] { + for bytes in [b"/next?x#y".as_slice(), b"https://example.com/path"] { + assert_eq!(validate_constructed_reference_len(length, bytes), Ok(())); + } + for bytes in [b"../next".as_slice(), b"/next%20path"] { + let error = validate_constructed_reference_len(length, bytes).unwrap_err(); + assert_eq!(error.header(), &FieldName::Location); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!(error.value_index(), None); + } + } + } +} diff --git a/crates/http_headers/src/headers/location/metadata.rs b/crates/http_headers/src/headers/location/metadata.rs new file mode 100644 index 000000000..280cb6821 --- /dev/null +++ b/crates/http_headers/src/headers/location/metadata.rs @@ -0,0 +1,83 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Retained component ranges in the validated semantic spelling. + +use std::num::NonZeroUsize; +use std::ops::Range; + +use fluent_uri::Uri; +use http_headers_simd::find_either; + +#[derive(Clone, Debug, Eq, PartialEq)] +pub(super) struct ComponentRanges { + pub(super) scheme_end: Option, + pub(super) authority_host: Option>, + pub(super) path: Range, + pub(super) query_end: Option, + pub(super) fragment_start: Option, +} + +impl ComponentRanges { + pub(super) fn from_parsed(parsed: &Uri<&str>) -> Self { + let scheme_end = parsed.scheme().and_then(|scheme| NonZeroUsize::new(scheme.as_str().len())); + let prefix_end = scheme_end.map_or(0, |end| end.get() + 1); + let authority_host = parsed.authority().map(|authority| { + let start = prefix_end + 2; + let host_start = start + authority.userinfo().map_or(0, |userinfo| userinfo.as_str().len() + 1); + host_start..host_start + authority.host().as_str().len() + }); + let path_start = prefix_end + parsed.authority().map_or(0, |authority| 2 + authority.as_str().len()); + let path_end = path_start + parsed.path().as_str().len(); + let query_end = parsed + .query() + .and_then(|query| NonZeroUsize::new(path_end + 1 + query.as_str().len())); + let fragment_start = parsed + .fragment() + .and_then(|_fragment| NonZeroUsize::new(query_end.map_or(path_end, NonZeroUsize::get) + 1)); + Self { + scheme_end, + authority_host, + path: path_start..path_end, + query_end, + fragment_start, + } + } + + /// Projects the subset already accepted by the SIMD scanner, not arbitrary input. + pub(super) fn from_simple(text: &str) -> Self { + let bytes = text.as_bytes(); + let (scheme_end, authority_host, path_start, path_end) = if bytes.first() == Some(&b'/') { + let path_end = find_either(bytes, b'?', b'#').unwrap_or(bytes.len()); + (None, None, 0, path_end) + } else { + let scheme_end = find_either(bytes, b':', b':').expect("the simple absolute subset includes a scheme colon"); + let start = scheme_end + 3; + let path_end = find_either(&bytes[start..], b'?', b'#').map_or(bytes.len(), |offset| start + offset); + // Colons and slashes in a query or fragment are not authority boundaries. + let host_end = find_either(&bytes[start..path_end], b':', b'/').map_or(path_end, |offset| start + offset); + let path_start = if bytes.get(host_end) == Some(&b':') { + find_either(&bytes[host_end + 1..path_end], b'/', b'/').map_or(path_end, |offset| host_end + 1 + offset) + } else { + host_end + }; + (NonZeroUsize::new(scheme_end), Some(start..host_end), path_start, path_end) + }; + let (query_end, fragment_start) = if bytes.get(path_end) == Some(&b'?') { + let query_end = find_either(&bytes[path_end + 1..], b'#', b'#').map_or(bytes.len(), |offset| path_end + 1 + offset); + ( + NonZeroUsize::new(query_end), + (query_end < bytes.len()).then(|| NonZeroUsize::new(query_end + 1)).flatten(), + ) + } else { + (None, (path_end < bytes.len()).then(|| NonZeroUsize::new(path_end + 1)).flatten()) + }; + Self { + scheme_end, + authority_host, + path: path_start..path_end, + query_end, + fragment_start, + } + } +} diff --git a/crates/http_headers/src/headers/location/uri_authority.rs b/crates/http_headers/src/headers/location/uri_authority.rs new file mode 100644 index 000000000..18063f4fe --- /dev/null +++ b/crates/http_headers/src/headers/location/uri_authority.rs @@ -0,0 +1,107 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Validated RFC 3986 authority components. + +use std::fmt; +use std::net::Ipv6Addr; + +use super::component::Component; +use super::invalid; +use crate::DecodeError; + +/// Validated, encoded authority components of a URI-reference. +/// +/// A host may be empty, a registered name, an IPv4 address, or a bracketed +/// IPv6/`IPvFuture` literal. Ports preserve their decimal spelling, including an +/// empty string, leading zeroes, or a value larger than `u16::MAX`. No DNS, +/// scheme-specific, or safe-redirect policy is implied. +/// +/// Equality and hashing compare the encoded components exactly, including +/// the distinction between absent and empty userinfo or port. They do not +/// fold case, decode escapes, or apply scheme-specific equivalence rules. +/// +/// # Examples +/// +/// ``` +/// use http_headers::headers::UriAuthority; +/// +/// let authority = UriAuthority::new(Some("user%40name"), "[::1]", Some("00443"))?; +/// assert_eq!(authority.userinfo(), Some("user%40name")); +/// assert_eq!(authority.host(), "[::1]"); +/// assert_eq!(authority.port(), Some("00443")); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +#[derive(Clone, Copy, Eq, Hash, PartialEq)] +pub struct UriAuthority<'a> { + pub(super) userinfo: Option<&'a str>, + pub(super) host: &'a str, + pub(super) port: Option<&'a str>, +} + +impl<'a> UriAuthority<'a> { + /// Validates encoded authority components without allocating. + /// + /// Pass brackets as part of an IPv6 or `IPvFuture` host. Components are + /// already encoded: a literal `@` in userinfo must be supplied as `%40`. + /// + /// # Errors + /// + /// Returns [`crate::DecodeErrorKind::InvalidSyntax`] for invalid userinfo, + /// host, or port grammar, including invalid percent escapes. + pub fn new(userinfo: Option<&'a str>, host: &'a str, port: Option<&'a str>) -> Result { + if let Some(userinfo) = userinfo { + Component::Userinfo.validate(userinfo)?; + } + validate_host(host)?; + if port.is_some_and(|port| !port.bytes().all(|byte| byte.is_ascii_digit())) { + return Err(invalid()); + } + Ok(Self { userinfo, host, port }) + } + + /// Returns encoded userinfo, distinguishing absent from present-but-empty. + #[inline] + #[must_use] + pub const fn userinfo(self) -> Option<&'a str> { + self.userinfo + } + + /// Returns the encoded host, retaining brackets around IP literals. + #[inline] + #[must_use] + pub const fn host(self) -> &'a str { + self.host + } + + /// Returns the decimal port spelling, distinguishing absent from empty. + /// + /// The generic URI grammar does not restrict ports to `u16` values. + #[inline] + #[must_use] + pub const fn port(self) -> Option<&'a str> { + self.port + } +} + +impl fmt::Debug for UriAuthority<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("UriAuthority").field("redacted", &true).finish_non_exhaustive() + } +} + +fn validate_host(host: &str) -> Result<(), DecodeError> { + if let Some(literal) = host.strip_prefix('[').and_then(|host| host.strip_suffix(']')) { + if literal.starts_with(['v', 'V']) { + let (version, address) = literal[1..].split_once('.').ok_or_else(invalid)?; + if version.is_empty() || !version.bytes().all(|byte| byte.is_ascii_hexdigit()) || address.is_empty() { + return Err(invalid()); + } + Component::IpvFuture.validate(address) + } else { + literal.parse::().map(|_address| ()).map_err(|_invalid| invalid()) + } + } else { + Component::RegisteredName.validate(host) + } +} diff --git a/crates/http_headers/src/headers/location/uri_reference.rs b/crates/http_headers/src/headers/location/uri_reference.rs new file mode 100644 index 000000000..f0c201b1f --- /dev/null +++ b/crates/http_headers/src/headers/location/uri_reference.rs @@ -0,0 +1,134 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Borrowed projections of retained URI-reference metadata. + +use std::fmt; +use std::hash::{Hash, Hasher}; + +use super::{ComponentRanges, UriAuthority}; + +/// A validated URI-reference with constant-time component projections. +/// +/// Components retain percent escapes: `%2F` stays part of its component, +/// rather than becoming a path separator. Query and fragment accessors +/// distinguish absence from a present-but-empty component. An empty path is +/// valid. No resolving, dot-segment removal, case folding, or percent decoding +/// is performed. +/// +/// In relaxed mode this view describes the spelling with backslashes replaced +/// by slashes. The containing [`super::LocationOwned`] or [`super::LocationView`] +/// still exposes and forwards the original wire value. +/// +/// Equality and hashing compare this semantic spelling exactly. This is not +/// URI equivalence under normalization or scheme-specific rules, and is not a +/// safe-redirect authorization policy. +/// +/// # Examples +/// +/// ``` +/// use http_headers::headers::LocationOwned; +/// +/// let location = LocationOwned::try_from("https://example.com/a%2Fb?#")?; +/// let uri = location.uri_reference(); +/// assert_eq!(uri.scheme(), Some("https")); +/// assert_eq!(uri.authority().unwrap().host(), "example.com"); +/// assert_eq!(uri.path(), "/a%2Fb"); +/// assert_eq!(uri.query(), Some("")); +/// assert_eq!(uri.fragment(), Some("")); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +#[derive(Clone, Copy)] +pub struct UriReference<'a> { + text: &'a str, + scheme: Option<&'a str>, + authority: Option>, + path: &'a str, + query: Option<&'a str>, + fragment: Option<&'a str>, +} + +impl<'a> UriReference<'a> { + pub(super) fn from_component_ranges(text: &'a str, ranges: &ComponentRanges) -> Self { + let authority = ranges.authority_host.as_ref().map(|host| { + let start = ranges.scheme_end.map_or(2, |end| end.get() + 3); + UriAuthority { + userinfo: (host.start != start).then(|| &text[start..host.start - 1]), + host: &text[host.clone()], + port: (host.end != ranges.path.start).then(|| &text[host.end + 1..ranges.path.start]), + } + }); + Self { + text, + scheme: ranges.scheme_end.map(|end| &text[..end.get()]), + authority, + path: &text[ranges.path.clone()], + query: ranges.query_end.map(|end| &text[ranges.path.end + 1..end.get()]), + fragment: ranges.fragment_start.map(|start| &text[start.get()..]), + } + } + + /// Returns the complete semantic spelling, normalized in relaxed mode. + #[inline] + #[must_use] + pub const fn as_str(self) -> &'a str { + self.text + } + + /// Returns the scheme without its colon, preserving case. + #[inline] + #[must_use] + pub const fn scheme(self) -> Option<&'a str> { + self.scheme + } + + /// Returns structured authority components, if `//` introduces an authority. + /// + /// `Some` may contain an empty host; it is distinct from no authority. + #[inline] + #[must_use] + pub const fn authority(self) -> Option> { + self.authority + } + + /// Returns the encoded path, which may be empty. + #[inline] + #[must_use] + pub const fn path(self) -> &'a str { + self.path + } + + /// Returns the encoded query without `?`, distinguishing absent from empty. + #[inline] + #[must_use] + pub const fn query(self) -> Option<&'a str> { + self.query + } + + /// Returns the encoded fragment without `#`, distinguishing absent from empty. + #[inline] + #[must_use] + pub const fn fragment(self) -> Option<&'a str> { + self.fragment + } +} + +impl PartialEq for UriReference<'_> { + fn eq(&self, other: &Self) -> bool { + self.text == other.text + } +} + +impl Eq for UriReference<'_> {} + +impl Hash for UriReference<'_> { + fn hash(&self, state: &mut H) { + self.text.hash(state); + } +} + +impl fmt::Debug for UriReference<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("UriReference").field("redacted", &true).finish_non_exhaustive() + } +} diff --git a/crates/http_headers/src/headers/method_view.rs b/crates/http_headers/src/headers/method_view.rs new file mode 100644 index 000000000..48b2d5858 --- /dev/null +++ b/crates/http_headers/src/headers/method_view.rs @@ -0,0 +1,121 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::InvalidMethod; +use crate::headers::tokens::method_text; +use crate::validate; + +/// A borrowed HTTP method token with case-sensitive equality and hashing. +/// +/// Extension methods retain their spelling. `*` is a method token, not a +/// wildcard, and `get` differs from [`Self::GET`]. +/// +/// Available with either the `headers-cors` or `headers-negotiation` feature. +/// +/// # Examples +/// +/// ``` +/// use http_headers::headers::MethodView; +/// +/// assert_ne!(MethodView::GET, MethodView::new("get")?); +/// assert_eq!(MethodView::new("X-CUSTOM")?.as_str(), "X-CUSTOM"); +/// # Ok::<(), http_headers::headers::InvalidMethod>(()) +/// ``` +#[derive(Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct MethodView<'a>(&'a [u8]); + +impl<'a> MethodView<'a> { + /// The `GET` method. + pub const GET: Self = Self(b"GET"); + /// The `HEAD` method. + pub const HEAD: Self = Self(b"HEAD"); + /// The `POST` method. + pub const POST: Self = Self(b"POST"); + /// The `PUT` method. + pub const PUT: Self = Self(b"PUT"); + /// The `DELETE` method. + pub const DELETE: Self = Self(b"DELETE"); + /// The `CONNECT` method. + pub const CONNECT: Self = Self(b"CONNECT"); + /// The `OPTIONS` method. + pub const OPTIONS: Self = Self(b"OPTIONS"); + /// The `TRACE` method. + pub const TRACE: Self = Self(b"TRACE"); + /// The `PATCH` method. + pub const PATCH: Self = Self(b"PATCH"); + + /// Validates and borrows a method without allocating. + /// + /// # Errors + /// + /// Returns [`InvalidMethod`] for empty input or non-token bytes. + #[inline] + pub fn new(method: &'a str) -> Result { + if validate::token(method.as_bytes()) { + Ok(Self(method.as_bytes())) + } else { + Err(InvalidMethod) + } + } + + #[inline] + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + Self(bytes) + } + + /// Returns the original spelling. + #[must_use] + #[inline] + pub fn as_str(self) -> &'a str { + method_text(self.0) + } + + /// Returns the original bytes. + #[must_use] + #[inline] + pub const fn as_bytes(self) -> &'a [u8] { + self.0 + } + + /// Converts to an owned [`http::Method`]. + /// + /// Standard methods do not allocate. Extension methods may allocate and + /// are checked against the external type's own constraints. + /// + /// # Errors + /// + /// Returns the external constructor's error if it rejects the token. + #[cfg(feature = "http")] + #[inline] + pub fn try_to_method(self) -> Result { + http::Method::from_bytes(self.as_bytes()) + } +} + +impl<'a> TryFrom<&'a str> for MethodView<'a> { + type Error = InvalidMethod; + + fn try_from(value: &'a str) -> Result { + Self::new(value) + } +} + +impl AsRef for MethodView<'_> { + fn as_ref(&self) -> &str { + self.as_str() + } +} + +impl fmt::Display for MethodView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } +} + +impl fmt::Debug for MethodView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_tuple("MethodView").field(&self.as_str()).finish() + } +} diff --git a/crates/http_headers/src/headers/mod.rs b/crates/http_headers/src/headers/mod.rs new file mode 100644 index 000000000..bd0b96993 --- /dev/null +++ b/crates/http_headers/src/headers/mod.rs @@ -0,0 +1,183 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Well-known HTTP header types. +//! +//! Header families expose descriptors plus borrowed `*View` and owned +//! `*Owned` values, re-exported here for convenience. +//! +//! # Examples +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-content-type"))] +//! # fn main() -> Result<(), Box> { +//! use http_headers::headers::ContentType; +//! +//! let mut headers = http::HeaderMap::new(); +//! ContentType::insert(&mut headers, ContentType::json())?; +//! assert!(ContentType::view(&headers)?.is_some()); +//! ContentType::remove(&mut headers); +//! assert!(ContentType::owned(&headers)?.is_none()); +//! # Ok(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-content-type")))] +//! # fn main() {} +//! ``` + +#[cfg(any(test, feature = "headers-authorization"))] +mod authorization; +#[cfg(any(test, feature = "headers-cache-control"))] +mod cache_control; +#[cfg(any(test, feature = "headers-conditional"))] +mod conditional; +#[cfg(any(test, feature = "headers-content-length"))] +mod content_length; +#[cfg(any(test, feature = "headers-content-type"))] +mod content_type; +#[cfg(any(test, feature = "headers-cors"))] +mod cors; +#[cfg(any(test, feature = "headers-etag"))] +mod etag; +mod extension_value; +#[cfg(any(test, feature = "headers-cors", feature = "headers-negotiation"))] +mod field_name_view; +#[cfg(any(test, feature = "headers-cors", feature = "headers-negotiation"))] +mod invalid_method; +#[cfg(any(test, feature = "headers-location"))] +mod location; +#[cfg(any(test, feature = "headers-cors", feature = "headers-negotiation"))] +mod method_view; +#[cfg(any(test, feature = "headers-negotiation"))] +mod negotiation; +#[cfg(any(test, feature = "headers-range"))] +mod range; +#[cfg(any(test, feature = "headers-security"))] +mod security; +#[cfg(any(test, feature = "headers-set-cookie"))] +mod set_cookie; +mod shared; +#[cfg(any(test, feature = "headers-cors", feature = "headers-negotiation"))] +mod tokens; +#[cfg(any(test, feature = "headers-user-agent"))] +mod user_agent; +#[cfg(any(test, feature = "headers-websocket"))] +mod websocket; + +#[cfg(any(test, feature = "headers-authorization"))] +#[doc(inline)] +pub use authorization::{Authorization, AuthorizationOwned, AuthorizationView, Basic, BasicCredentials, Bearer}; +#[cfg(any(test, feature = "headers-cache-control"))] +#[doc(inline)] +pub use cache_control::{CacheControl, CacheControlBuilder, CacheControlOwned, CacheControlView, CacheDirectiveView}; +#[cfg(any(test, feature = "headers-conditional"))] +#[doc(inline)] +pub use conditional::{ + ConditionalTagView, IfMatch, IfMatchOwned, IfMatchView, IfModifiedSince, IfModifiedSinceOwned, IfModifiedSinceView, IfNoneMatch, + IfNoneMatchOwned, IfNoneMatchView, IfRange, IfRangeOwned, IfRangeValueView, IfRangeView, IfUnmodifiedSince, IfUnmodifiedSinceOwned, + IfUnmodifiedSinceView, LastModified, LastModifiedOwned, LastModifiedView, +}; +#[cfg(any(test, feature = "headers-content-length"))] +#[doc(inline)] +pub use content_length::{ContentLength, ContentLengthOwned}; +#[cfg(any(test, feature = "headers-content-type"))] +#[doc(inline)] +pub use content_type::{ContentType, ContentTypeOwned, ContentTypeView, MediaTypeParameterView, MediaTypeParameters}; +#[cfg(any(test, feature = "headers-cors"))] +#[doc(inline)] +pub use cors::{ + AccessControlAllowCredentials, AccessControlAllowCredentialsOwned, AccessControlAllowCredentialsView, AccessControlAllowHeaders, + AccessControlAllowHeadersOwned, AccessControlAllowHeadersView, AccessControlAllowMethods, AccessControlAllowMethodsOwned, + AccessControlAllowMethodsView, AccessControlAllowOrigin, AccessControlAllowOriginKind, AccessControlAllowOriginOwned, + AccessControlAllowOriginView, AccessControlExposeHeaders, AccessControlExposeHeadersOwned, AccessControlExposeHeadersView, + AccessControlMaxAge, AccessControlMaxAgeOwned, AccessControlRequestHeaders, AccessControlRequestHeadersOwned, + AccessControlRequestHeadersView, AccessControlRequestMethod, AccessControlRequestMethodOwned, AccessControlRequestMethodView, + CorsHeaderNames, CorsMethods, OriginDomainView, OriginHost, OriginScheme, SerializedOriginView, +}; +#[cfg(any(test, feature = "headers-etag"))] +#[doc(inline)] +pub use etag::{ETag, ETagOwned, ETagView}; +#[doc(inline)] +pub use extension_value::ExtensionValue; +#[cfg(any(test, feature = "headers-cors", feature = "headers-negotiation"))] +#[doc(inline)] +pub use field_name_view::FieldNameView; +#[cfg(any(test, feature = "headers-cors", feature = "headers-negotiation"))] +#[doc(inline)] +pub use invalid_method::InvalidMethod; +#[cfg(any(test, feature = "headers-location"))] +#[doc(inline)] +pub use location::{Location, LocationOwned, LocationView, UriAuthority, UriReference}; +#[cfg(any(test, feature = "headers-cors", feature = "headers-negotiation"))] +#[doc(inline)] +pub use method_view::MethodView; +#[cfg(any(test, feature = "headers-negotiation"))] +#[doc(inline)] +pub use negotiation::{ + Accept, AcceptEncoding, AcceptEncodingEntry, AcceptEncodingOwned, AcceptEncodingView, AcceptEntry, AcceptLanguage, AcceptLanguageEntry, + AcceptLanguageOwned, AcceptLanguageView, AcceptOwned, AcceptView, Allow, AllowOwned, AllowView, ContentCoding, ContentCodingKind, Host, + HostKind, HostOwned, HostPortView, HostView, InexactQuality, InvalidQuality, IpvFutureView, LanguageRange, MediaRange, MediaRangeKind, + NegotiationParameter, NegotiationParameterValue, NegotiationParameters, NegotiationToken, PortConversionError, PortConversionErrorKind, + Quality, QualityView, RegisteredNameView, Server, ServerOwned, ServerView, Vary, VaryEntryView, VaryOwned, VaryView, +}; +#[cfg(any(test, feature = "headers-range"))] +#[doc(inline)] +pub use range::{ + AcceptRanges, AcceptRangesOwned, AcceptRangesView, ByteContentRange, ByteRangeSpec, CompleteLength, ContentRange, ContentRangeOwned, + ContentRangeView, Range, RangeOwned, RangeView, +}; +#[cfg(any(test, feature = "headers-security"))] +#[doc(inline)] +pub use security::{ + ContentSecurityPolicy, ContentSecurityPolicyOwned, ContentSecurityPolicyView, HstsDirectiveView, ReferrerPolicy, ReferrerPolicyOwned, + ReferrerPolicyTokenView, ReferrerPolicyValue, ReferrerPolicyView, StrictTransportSecurity, StrictTransportSecurityBuilder, + StrictTransportSecurityOwned, StrictTransportSecurityView, XContentTypeOptions, XContentTypeOptionsOwned, XContentTypeOptionsView, +}; +#[cfg(any(test, feature = "headers-set-cookie"))] +#[doc(inline)] +pub use set_cookie::{SetCookie, SetCookieOwned, SetCookieView}; +#[cfg(any(test, feature = "headers-negotiation", feature = "headers-user-agent"))] +use shared::has_non_ows; +#[cfg(any( + test, + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-type", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +use shared::invalid_syntax; +#[cfg(any(test, feature = "headers-cache-control", feature = "headers-range"))] +use shared::normalized_comma_value; +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-range", + feature = "headers-security", + feature = "headers-websocket", +))] +use shared::trim_ows; +#[cfg(any( + test, + feature = "headers-authorization", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-websocket", +))] +use shared::value_from_bytes; +#[cfg(any(test, feature = "headers-user-agent"))] +#[doc(inline)] +pub use user_agent::{UserAgent, UserAgentOwned, UserAgentView}; +#[cfg(any(test, feature = "headers-websocket"))] +#[doc(inline)] +pub use websocket::{ + SecWebSocketAccept, SecWebSocketAcceptOwned, SecWebSocketAcceptView, SecWebSocketExtensions, SecWebSocketExtensionsBuilder, + SecWebSocketExtensionsOwned, SecWebSocketExtensionsView, SecWebSocketKey, SecWebSocketKeyOwned, SecWebSocketKeyView, + SecWebSocketProtocol, SecWebSocketProtocolOwned, SecWebSocketProtocolView, SecWebSocketVersion, SecWebSocketVersionOwned, + WebSocketExtensionParameterView, WebSocketExtensionParameters, WebSocketExtensionView, +}; diff --git a/crates/http_headers/src/headers/negotiation/accept.rs b/crates/http_headers/src/headers/negotiation/accept.rs new file mode 100644 index 000000000..a1bdc7dc8 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/accept.rs @@ -0,0 +1,497 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::accept_entry::AcceptEntry; +use super::negotiation_members::{collect_members, validated_member}; +use super::shared::{ + ListValues, QuotedItems, check_quoted_value, check_quoted_values, invalid, invalid_syntax, parse_parameter, try_plain_items, + validate_quality, +}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Owned value for the `Accept` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 12.5.1]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::AcceptOwned::try_from("text/html, application/json")?; +/// assert_eq!(value.items().count(), 2); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Accept: text/html, application/xhtml+xml, */*;q=0.9` combines exact and +/// wildcard media ranges with a quality weight. +/// +/// # Relaxed decoding +/// +/// [`DecodeMode::Relaxed`](crate::DecodeMode::Relaxed) permits spaces or tabs around the quality +/// parameter's `=`, leading-dot quality values, and more than three +/// fractional digits. All media-range, wildcard, parameter-ordering, quoting, +/// and list rules remain strict, and the original bytes are preserved. +/// +/// [RFC 9110 section 12.5.1]: https://www.rfc-editor.org/rfc/rfc9110#section-12.5.1 +pub struct AcceptOwned { + values: ListValues, +} + +/// Borrowed value for the `Accept` header. +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{Accept, AcceptView}; +/// use http_headers::source::{FieldLines, FieldSource}; +/// use http_headers::{Field, FieldName}; +/// +/// struct Source; +/// +/// impl FieldSource for Source { +/// fn lines(&self, name: &'static FieldName) -> Option> { +/// (name == &FieldName::Accept) +/// .then(|| FieldLines::single(name, b"text/html;q=0.8, application/json")) +/// } +/// } +/// +/// let value: AcceptView<'_> = Accept::view(&Source)?.expect("header is present"); +/// let items = value.items().collect::>(); +/// assert_eq!(items, [&b"text/html;q=0.8"[..], &b"application/json"[..]]); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct AcceptView<'a> { + values: FieldLines<'a>, +} + +super::shared::list_header!( + Accept, + AcceptOwned, + AcceptView, + "Accept", + "Defined by [RFC 9110 section 12.5.1](https://www.rfc-editor.org/rfc/rfc9110#section-12.5.1).", + &FieldName::Accept, + validate_accept, + validate_accept_relaxed, + check_quoted_values, + check_quoted_value, + quoted +); + +impl AcceptOwned { + /// Iterates typed preferences in field-line and member order. + /// + /// Empty list members are ignored. This allocation-free streaming + /// projection retains metadata in each yielded entry; a fresh iterator + /// traverses the wire again. Scalar entry getters do not parse again. + /// + /// # Examples + /// + /// ``` + /// # #[cfg(feature = "headers-negotiation")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http_headers::headers::{AcceptOwned, MediaRangeKind}; + /// + /// let accept = AcceptOwned::try_from("text/html, text/*;q=0.5")?; + /// let mut preferences = accept.entries(); + /// assert_eq!( + /// preferences.next().unwrap().range().subtype().as_str(), + /// "html" + /// ); + /// let wildcard = preferences.next().unwrap(); + /// assert_eq!(wildcard.range().kind(), MediaRangeKind::TypeWildcard); + /// assert_eq!(wildcard.quality().to_quality().unwrap().thousandths(), 500); + /// # Ok(()) + /// # } + /// # #[cfg(not(feature = "headers-negotiation"))] + /// # fn main() {} + /// ``` + pub fn entries(&self) -> impl Iterator> { + self.values + .iter() + .flat_map(|value| QuotedItems::comma(value.as_bytes(), &FieldName::Accept)) + .map(validated_member) + .map(AcceptEntry::from_validated) + } + + /// Constructs one field line from typed preferences without sorting them. + /// + /// Quality spelling and separators are written in canonical form; token spelling and + /// parameter bytes are retained. An empty iterator creates a present empty + /// list. Use raw field values for byte-exact forwarding. + /// + /// # Errors + /// + /// Returns an error if the aggregate wire size or member count exceeds the + /// custom-source budgets. + pub fn from_entries<'a>(entries: impl IntoIterator>) -> Result { + collect_members(entries, &FieldName::Accept, AcceptEntry::encoded_len, AcceptEntry::append_to).map(|values| Self { values }) + } +} + +impl<'a> AcceptView<'a> { + /// Iterates typed preferences without allocating, preserving duplicates. + /// + /// Each new iterator traverses the wire again. Each yielded entry retains + /// its scalar components and media-parameter/extension boundaries. + pub fn entries(&self) -> impl Iterator> + '_ { + self.values + .validated_comma_items() + .map(validated_member) + .map(AcceptEntry::from_validated) + } +} + +fn validate_accept(bytes: &[u8]) -> Result<(), DecodeError> { + validate_accept_with(bytes, false) +} + +fn validate_accept_relaxed(bytes: &[u8]) -> Result<(), DecodeError> { + validate_accept_with(bytes, true) +} + +fn validate_accept_with(bytes: &[u8], relaxed: bool) -> Result<(), DecodeError> { + if validate_accept_plain(bytes, relaxed)? { + return Ok(()); + } + validate_accept_quoted(bytes, relaxed) +} + +fn validate_accept_plain(bytes: &[u8], relaxed: bool) -> Result { + let mut first = true; + let mut quality_seen = false; + try_plain_items(bytes, b';', false, |segment| { + if first { + first = false; + return validate_media_range(segment); + } + validate_accept_parameter(segment, &mut quality_seen, relaxed) + }) +} + +fn validate_accept_quoted(bytes: &[u8], relaxed: bool) -> Result<(), DecodeError> { + let mut parameters = QuotedItems::semicolon(bytes, &FieldName::Accept); + let media_range = parameters.next().expect("semicolon iteration always yields a first item")?; + validate_media_range(media_range)?; + let mut quality_seen = false; + for parameter in parameters { + let (name, value, compact) = parse_parameter(parameter?, quality_seen, &FieldName::Accept)?; + if validate::eq_ignore_ascii_case(name, b"q") { + if quality_seen { + return Err(invalid_syntax(&FieldName::Accept)); + } + if !relaxed && !compact { + return Err(invalid_syntax(&FieldName::Accept)); + } + let quality = value.expect("quality parameters require a value"); + validate_quality(quality, &FieldName::Accept, relaxed)?; + quality_seen = true; + } else if !compact { + return Err(invalid_syntax(&FieldName::Accept)); + } + } + Ok(()) +} + +fn validate_accept_parameter(bytes: &[u8], quality_seen: &mut bool, relaxed: bool) -> Result<(), DecodeError> { + if !*quality_seen && bytes.len() >= 2 && bytes[0].eq_ignore_ascii_case(&b'q') && bytes[1] == b'=' { + validate_quality(&bytes[2..], &FieldName::Accept, relaxed)?; + *quality_seen = true; + return Ok(()); + } + let (name, value, compact) = parse_parameter(bytes, *quality_seen, &FieldName::Accept)?; + if validate::eq_ignore_ascii_case(name, b"q") { + if *quality_seen || (!relaxed && !compact) { + return Err(invalid_syntax(&FieldName::Accept)); + } + let quality = value.expect("quality parameters require a value"); + validate_quality(quality, &FieldName::Accept, relaxed)?; + *quality_seen = true; + } else if !compact { + return Err(invalid_syntax(&FieldName::Accept)); + } + Ok(()) +} + +/// Validates one media range. +/// +/// The range is split at its first slash and each half is checked as a token. +/// A slash is not a token byte, so a second one already fails the subtype +/// check, and the scan that tells a repeated slash apart from an ordinary bad +/// byte runs only when the range is being rejected anyway. +pub(super) fn validate_media_range(bytes: &[u8]) -> Result<(), DecodeError> { + let mut halves = bytes.splitn(2, |byte| *byte == b'/'); + let type_ = halves.next().expect("splitting always yields a first half"); + let Some(subtype) = halves.next() else { + return Err(invalid_syntax(&FieldName::Accept)); + }; + + if !validate::token(type_) || !validate::token(subtype) { + if subtype.contains(&b'/') { + return Err(invalid_syntax(&FieldName::Accept)); + } + return Err(invalid(&FieldName::Accept, DecodeErrorKind::InvalidToken)); + } + + if type_ == b"*" && subtype != b"*" { + return Err(invalid_syntax(&FieldName::Accept)); + } + + Ok(()) +} + +#[cfg(test)] +#[expect( + clippy::assertions_on_result_states, + reason = "the tests classify many parser outcomes without needing their success values" +)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::super::accept_scan::scan_accept_line; + use super::super::recognition_test_support::{self, exhaustive}; + use super::super::shared::{WELL_KNOWN_ACCEPT, invalid, invalid_syntax}; + use super::{ + validate_accept, validate_accept_parameter, validate_accept_quoted, validate_accept_relaxed, validate_accept_with, + validate_media_range, + }; + use crate::{DecodeError, DecodeErrorKind, FieldName, validate}; + + #[test] + fn quoted_parameters_enforce_media_range_and_quality_rules() { + validate_accept(b"text/plain").expect("strict plain media range"); + validate_accept_relaxed(b"text/plain; q = .5").expect("relaxed quality"); + validate_accept_with(b"text/plain;level=\"one\"", false).expect("quoted extension parameter"); + assert!(validate_accept_quoted(b"text/plain;level=\"one\";q=0.5", false).is_ok()); + assert_eq!( + validate_accept_quoted(b"\"unterminated", false) + .expect_err("unterminated first item") + .kind(), + DecodeErrorKind::UnterminatedQuote + ); + assert_eq!( + validate_accept_quoted(b"", false).expect_err("missing media range").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + validate_accept_quoted(b"text/plain;level=\"unterminated", false) + .expect_err("unterminated parameter") + .kind(), + DecodeErrorKind::UnterminatedQuote + ); + assert_eq!( + validate_accept_quoted(b"text/plain;q=0.5;Q=0.4", false) + .expect_err("duplicate quality") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + validate_accept_quoted(b"text/plain;q = 0.5", false) + .expect_err("strict whitespace") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert!(validate_accept_quoted(b"text/plain;charset = \"utf-8\"", false).is_err()); + assert!(validate_accept_quoted(b"text/plain;charset = \"utf-8\"", true).is_err()); + } + + #[test] + fn plain_parameters_and_media_ranges_reject_structural_errors() { + let mut quality_seen = false; + assert!(validate_accept_parameter(b"level=one", &mut quality_seen, false).is_ok()); + assert!(!quality_seen); + assert!(validate_accept_parameter(b"q=0.8", &mut quality_seen, false).is_ok()); + assert!(quality_seen); + assert_eq!( + validate_accept_parameter(b"q=0.7", &mut quality_seen, false) + .expect_err("duplicate quality") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let mut quality_seen = false; + assert_eq!( + validate_accept_parameter(b"q = 0.8", &mut quality_seen, false) + .expect_err("strict whitespace") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert!(validate_accept_parameter(b"charset = utf-8", &mut quality_seen, false).is_err()); + assert!(validate_accept_parameter(b"charset = utf-8", &mut quality_seen, true).is_err()); + assert!(validate_accept_parameter(b"q = .8", &mut quality_seen, true).is_ok()); + + for malformed in [b"text".as_slice(), b"text/plain/more", b"te xt/plain", b"*/plain"] { + assert!(validate_media_range(malformed).is_err(), "{malformed:?}"); + } + assert!(validate_media_range(b"*/*").is_ok()); + assert!(validate_media_range(b"text/*").is_ok()); + } + + /// Builds every concatenation of up to `count` grammar fragments. + fn fragment_lines(count: usize) -> Vec> { + const FRAGMENTS: &[&str] = &[ + "a/b", + "*/*", + "a/*", + "*/a", + ";q=0.9", + ";q=1", + ";q=0", + ";q=", + ";q=1.5", + ";q=0.1234", + ";q=1.000", + ";q=1.001", + ";Q=0.123", + ";x=y", + ";x", + "; q=0.9", + " ;q=0.9", + ";\tq=0.9", + ";q=0.9 ", + "; ", + " , ", + ",", + " ", + "\t", + "/", + ";", + "=", + "\"", + "\\", + "a", + "0.", + ";q=0.", + ]; + + recognition_test_support::fragment_lines(FRAGMENTS, count) + } + + fn assert_recognition_is_sound(lines: &[Vec]) -> usize { + recognition_test_support::assert_recognition_is_sound(lines, scan_accept_line, validate_accept, validate_accept_relaxed) + } + + #[test] + fn recognized_short_lines_satisfy_the_member_grammar() { + let recognized = assert_recognition_is_sound(&exhaustive(b"a*/;,=qQ01. ", 4)); + assert!( + recognized > 100, + "the corpus must actually exercise the recognizer, saw {recognized}" + ); + } + + #[test] + fn recognized_fragment_lines_satisfy_the_member_grammar() { + let recognized = assert_recognition_is_sound(&fragment_lines(3)); + assert!( + recognized > 100, + "the corpus must actually exercise the recognizer, saw {recognized}" + ); + } + + #[test] + fn realistic_lines_are_recognized() { + const LINES: &[&str] = &[ + "*/*", + "text/html", + "application/json", + "text/html, application/json", + "text/html;q=0.9", + "text/html, */*;q=0.8", + "application/json, text/plain;q=0.9, */*;q=0.1", + "text/html,application/xhtml+xml,application/xml;q=0.9,image/avif,image/webp,image/apng,*/*;q=0.8,application/signed-exchange;v=b3;q=0.7", + ]; + + for line in LINES { + assert!( + scan_accept_line(line.as_bytes()), + "the recognizer must settle the common line {line:?}" + ); + } + } + + /// Validates a media range with an independent rule-at-a-time reference parser. + fn media_range_reference(bytes: &[u8]) -> Result<(), DecodeError> { + let Some(slash) = bytes.iter().position(|byte| *byte == b'/') else { + return Err(invalid_syntax(&FieldName::Accept)); + }; + if bytes[slash + 1..].contains(&b'/') { + return Err(invalid_syntax(&FieldName::Accept)); + } + let type_ = &bytes[..slash]; + let subtype = &bytes[slash + 1..]; + if !validate::token(type_) || !validate::token(subtype) { + return Err(invalid(&FieldName::Accept, DecodeErrorKind::InvalidToken)); + } + if type_ == b"*" && subtype != b"*" { + return Err(invalid_syntax(&FieldName::Accept)); + } + Ok(()) + } + + fn media_ranges() -> Vec> { + let mut ranges: Vec> = [ + "", + "/", + "//", + "*/*", + "*/plain", + "text/*", + "text", + "text/", + "/plain", + "text/plain", + "text/plain/extra", + "a/b/c", + "*", + "**/*", + "*/**", + "text//plain", + " text/plain", + "text/plain ", + "text /plain", + "text/ plain", + ] + .iter() + .map(|range| range.as_bytes().to_vec()) + .collect(); + + for byte in crate::test_support::byte_cases(b't') { + ranges.push(vec![byte]); + ranges.push(vec![b't', byte, b'/', b'p']); + ranges.push(vec![b't', b'/', byte, b'p']); + ranges.push(vec![byte, b'/', b'*']); + ranges.push(vec![b'*', b'/', byte]); + } + ranges + } + + #[test] + fn single_pass_media_range_matches_the_rule_at_a_time_walk() { + for range in media_ranges() { + #[cfg(miri)] + assert_eq!(validate_media_range(&range), media_range_reference(&range), "media range {range:?}"); + #[cfg(not(miri))] + assert_eq!( + format!("{:?}", validate_media_range(&range)), + format!("{:?}", media_range_reference(&range)), + "media range {:?} must validate identically", + String::from_utf8_lossy(&range) + ); + } + } + + #[test] + fn every_well_known_line_also_satisfies_the_member_grammar() { + let rejected: Vec<_> = WELL_KNOWN_ACCEPT + .iter() + .filter(|line| validate_accept(line).is_err()) + .map(|line| String::from_utf8_lossy(line).into_owned()) + .collect(); + assert!( + rejected.is_empty(), + "recognized lines must satisfy the grammar they skip: {rejected:?}" + ); + } +} diff --git a/crates/http_headers/src/headers/negotiation/accept_encoding.rs b/crates/http_headers/src/headers/negotiation/accept_encoding.rs new file mode 100644 index 000000000..0bf028769 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/accept_encoding.rs @@ -0,0 +1,250 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::accept_encoding_entry::AcceptEncodingEntry; +use super::negotiation_members::{collect_members, validated_member}; +use super::shared::{ListValues, QuotedItems, check_quoted_value, check_quoted_values, validate_weighted_token}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Owned value for the `Accept-Encoding` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 12.5.3]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::AcceptEncodingOwned::try_from("gzip, br")?; +/// assert_eq!(value.items().count(), 2); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Accept-Encoding: gzip, deflate, br` lists codings. +/// `Accept-Encoding: gzip;q=1.0, *;q=0` adds quality and wildcard preferences. +/// +/// # Relaxed decoding +/// +/// [`DecodeMode::Relaxed`](crate::DecodeMode::Relaxed) permits spaces or tabs around the quality +/// parameter's `=`, leading-dot quality values, and more than three +/// fractional digits. Content-coding tokens, wildcard syntax, parameter +/// count, quoting, and list structure remain strict, and the original bytes +/// are preserved. +/// +/// [RFC 9110 section 12.5.3]: https://www.rfc-editor.org/rfc/rfc9110#section-12.5.3 +pub struct AcceptEncodingOwned { + values: ListValues, +} + +/// Borrowed value for the `Accept-Encoding` header. +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{AcceptEncoding, AcceptEncodingView}; +/// use http_headers::source::{FieldLines, FieldSource}; +/// use http_headers::{Field, FieldName}; +/// +/// struct Source; +/// +/// impl FieldSource for Source { +/// fn lines(&self, name: &'static FieldName) -> Option> { +/// (name == &FieldName::AcceptEncoding) +/// .then(|| FieldLines::single(name, b"gzip, br;q=0.8")) +/// } +/// } +/// +/// let value: AcceptEncodingView<'_> = AcceptEncoding::view(&Source)?.expect("header is present"); +/// let items = value.items().collect::>(); +/// assert_eq!(items, [&b"gzip"[..], &b"br;q=0.8"[..]]); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct AcceptEncodingView<'a> { + values: FieldLines<'a>, +} + +super::shared::list_header!( + AcceptEncoding, + AcceptEncodingOwned, + AcceptEncodingView, + "Accept-Encoding", + "Defined by [RFC 9110 section 12.5.3](https://www.rfc-editor.org/rfc/rfc9110#section-12.5.3).", + &FieldName::AcceptEncoding, + validate_accept_encoding, + validate_accept_encoding_relaxed, + check_quoted_values, + check_quoted_value, + quoted +); + +impl AcceptEncodingOwned { + /// Iterates typed coding preferences in wire order, retaining duplicates. + /// + /// A fresh iterator traverses the wire again without allocating. Each + /// yielded member retains its coding and exact quality for repeated reads. + pub fn entries(&self) -> impl Iterator> { + self.values + .iter() + .flat_map(|value| QuotedItems::comma(value.as_bytes(), &FieldName::AcceptEncoding)) + .map(validated_member) + .map(AcceptEncodingEntry::from_validated) + } + + /// Constructs one field line from typed coding preferences. + /// + /// Order and duplicates are preserved. Separators and quality spellings + /// are written in canonical form. Empty input creates an empty list, not a wildcard. + /// + /// # Errors + /// + /// Returns an error if the aggregate wire size or member count exceeds the + /// custom-source budgets. + pub fn from_entries<'a>(entries: impl IntoIterator>) -> Result { + collect_members( + entries, + &FieldName::AcceptEncoding, + AcceptEncodingEntry::encoded_len, + AcceptEncodingEntry::append_to, + ) + .map(|values| Self { values }) + } +} + +impl<'a> AcceptEncodingView<'a> { + /// Iterates typed coding preferences without allocating or adding entries. + /// + /// Each new iterator traverses the wire again; member getters reuse their + /// retained components. Empty list members are ignored. + pub fn entries(&self) -> impl Iterator> + '_ { + self.values + .validated_comma_items() + .map(validated_member) + .map(AcceptEncodingEntry::from_validated) + } +} + +fn validate_accept_encoding(bytes: &[u8]) -> Result<(), DecodeError> { + validate_weighted_token(bytes, &FieldName::AcceptEncoding, validate::token, false) +} + +fn validate_accept_encoding_relaxed(bytes: &[u8]) -> Result<(), DecodeError> { + validate_weighted_token(bytes, &FieldName::AcceptEncoding, validate::token, true) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::super::recognition_test_support::{self, exhaustive, fragment_lines}; + use super::super::shared::WELL_KNOWN_ACCEPT_ENCODING; + use super::super::weighted_token_scan::scan_accept_encoding_line; + use super::{validate_accept_encoding, validate_accept_encoding_relaxed}; + use crate::DecodeErrorKind; + + #[test] + fn strict_and_relaxed_encoding_weights_use_the_shared_grammar() { + validate_accept_encoding(b"gzip;q=0.5").expect("strict encoding weight"); + validate_accept_encoding_relaxed(b"br; q = .1234").expect("relaxed encoding weight"); + assert_eq!( + validate_accept_encoding(b"bad encoding") + .expect_err("encoding must be a token") + .kind(), + DecodeErrorKind::InvalidToken + ); + } + + #[test] + fn every_well_known_line_also_satisfies_the_member_grammar() { + let rejected: Vec<_> = WELL_KNOWN_ACCEPT_ENCODING + .iter() + .filter(|line| validate_accept_encoding(line).is_err()) + .map(|line| String::from_utf8_lossy(line).into_owned()) + .collect(); + assert!( + rejected.is_empty(), + "recognized lines must satisfy the grammar they skip: {rejected:?}" + ); + } + + fn assert_recognition_is_sound(lines: &[Vec]) -> usize { + recognition_test_support::assert_recognition_is_sound( + lines, + scan_accept_encoding_line, + validate_accept_encoding, + validate_accept_encoding_relaxed, + ) + } + + #[test] + fn recognized_short_lines_satisfy_the_member_grammar() { + let recognized = assert_recognition_is_sound(&exhaustive(b"a*;,=qQ01. ", 4)); + assert!( + recognized > 100, + "the corpus must actually exercise the recognizer, saw {recognized}" + ); + } + + #[test] + fn recognized_fragment_lines_satisfy_the_member_grammar() { + const FRAGMENTS: &[&str] = &[ + "gzip", + "br", + "*", + ";q=0.9", + ";q=1", + ";q=0", + ";q=", + ";q=1.5", + ";q=0.1234", + ";q=1.000", + ";q=1.001", + ";Q=0.123", + ";x=y", + ";q=0.5;q=0.5", + "; q=0.9", + " ;q=0.9", + ";\tq=0.9", + ";q=0.9 ", + "; ", + " , ", + ",", + " ", + "\t", + ";", + "=", + "\"", + "\\", + "/", + "a", + "0.", + ]; + + let recognized = assert_recognition_is_sound(&fragment_lines(FRAGMENTS, 3)); + assert!( + recognized > 100, + "the corpus must actually exercise the recognizer, saw {recognized}" + ); + } + + #[test] + fn realistic_lines_are_recognized() { + const LINES: &[&str] = &[ + "*", + "gzip", + "gzip, deflate", + "gzip, deflate, br", + "gzip, deflate, br, zstd", + "gzip, deflate, br, zstd;q=0.9", + "gzip; q=0.9, deflate ;q=0.8", + "gzip;q=1.0, identity;q=0.5, *;q=0", + "br;q=1.0, gzip;q=0.8, *;q=0.1", + ]; + + for line in LINES { + assert!( + scan_accept_encoding_line(line.as_bytes()), + "the recognizer must settle the common line {line:?}" + ); + } + } +} diff --git a/crates/http_headers/src/headers/negotiation/accept_encoding_entry.rs b/crates/http_headers/src/headers/negotiation/accept_encoding_entry.rs new file mode 100644 index 000000000..ed062be07 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/accept_encoding_entry.rs @@ -0,0 +1,69 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::content_coding::ContentCoding; +use super::negotiation_members::{append_weight, checked_len, weight_len, weighted_parts}; +use super::quality::QualityView; +use crate::{DecodeError, FieldName}; + +/// One validated Accept-Encoding preference with retained coding and quality. +/// +/// Wildcard, identity and extension codings remain distinct. No implicit +/// identity entry, sorting or codec selection is performed. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct AcceptEncodingEntry<'a> { + coding: ContentCoding<'a>, + quality: Option>, +} + +impl<'a> AcceptEncodingEntry<'a> { + /// Constructs a coding preference without revalidating its components. + /// + /// # Errors + /// + /// Returns an error if its serialized size exceeds + /// [`MAX_CUSTOM_FIELD_BYTES`](crate::source::MAX_CUSTOM_FIELD_BYTES). + pub fn new(coding: ContentCoding<'a>, quality: Option>) -> Result { + let entry = Self { coding, quality }; + entry.encoded_len()?; + Ok(entry) + } + + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + let (coding, quality) = weighted_parts(bytes); + Self { + coding: ContentCoding::from_validated(coding), + quality, + } + } + + /// Returns the retained coding classification and original token. + #[must_use] + pub const fn coding(self) -> ContentCoding<'a> { + self.coding + } + + /// Returns the exact quality, defaulting to one when `q` was omitted. + #[must_use] + pub const fn quality(self) -> QualityView<'a> { + match self.quality { + Some(quality) => quality, + None => QualityView::ONE, + } + } + + /// Returns only an explicitly supplied quality. + #[must_use] + pub const fn explicit_quality(self) -> Option> { + self.quality + } + + pub(super) fn encoded_len(&self) -> Result { + checked_len([self.coding.as_str().len(), weight_len(self.quality)], &FieldName::AcceptEncoding) + } + + pub(super) fn append_to(&self, bytes: &mut Vec) { + bytes.extend_from_slice(self.coding.as_str().as_bytes()); + append_weight(bytes, self.quality); + } +} diff --git a/crates/http_headers/src/headers/negotiation/accept_entry.rs b/crates/http_headers/src/headers/negotiation/accept_entry.rs new file mode 100644 index 000000000..902995980 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/accept_entry.rs @@ -0,0 +1,294 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::media_range::MediaRange; +use super::negotiation_members::{append_weight, checked_len, weight_len}; +use super::negotiation_parameter::{NegotiationParameter, NegotiationParameters, ParameterSource}; +use super::quality::QualityView; +use super::shared::{QuotedItems, invalid_syntax}; +use crate::{DecodeError, FieldName, validate}; + +/// One validated Accept preference with retained semantic components. +/// +/// Scalar getters do not parse again. Parameter iteration traverses only its +/// retained region; each fresh iteration or lookup can traverse it again. +/// Media parameters precede `q`, and extensions follow it. Missing `q` means +/// effective quality one, but remains distinguishable from explicit `q=1`. +#[derive(Clone, Copy, Debug)] +pub struct AcceptEntry<'a> { + range: MediaRange<'a>, + quality: Option>, + parameters: ParameterSource<'a>, + extensions: ParameterSource<'a>, +} + +impl<'a> AcceptEntry<'a> { + /// Constructs a preference from validated components. + /// + /// # Errors + /// + /// Media parameters require values. Neither parameter list may contain a + /// `q` name, and extensions require an explicit quality to separate them + /// from media parameters. The serialized member must fit + /// [`MAX_CUSTOM_FIELD_BYTES`](crate::source::MAX_CUSTOM_FIELD_BYTES). + pub fn new( + range: MediaRange<'a>, + parameters: &'a [NegotiationParameter<'a>], + quality: Option>, + extensions: &'a [NegotiationParameter<'a>], + ) -> Result { + if parameters + .iter() + .any(|parameter| parameter.name().eq_ignore_ascii_case("q") || parameter.value().is_none()) + || extensions.iter().any(|parameter| parameter.name().eq_ignore_ascii_case("q")) + || (!extensions.is_empty() && quality.is_none()) + { + return Err(invalid_syntax(&FieldName::Accept)); + } + let entry = Self { + range, + quality, + parameters: ParameterSource::Components(parameters), + extensions: ParameterSource::Components(extensions), + }; + entry.encoded_len()?; + Ok(entry) + } + + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + let head_end = bytes.iter().position(|byte| *byte == b';').unwrap_or(bytes.len()); + let range = MediaRange::from_validated(validate::trim_ows(&bytes[..head_end])); + if head_end == bytes.len() { + return Self { + range, + quality: None, + parameters: ParameterSource::Wire(&[]), + extensions: ParameterSource::Wire(&[]), + }; + } + let parameter_start = head_end + 1; + if let [b'q' | b'Q', b'=', tail @ ..] = &bytes[parameter_start..] { + // Header validation excludes quoted qualities, so the first + // semicolon closes the quality rather than a quoted value. + let (quality, extensions) = match tail.iter().position(|byte| *byte == b';') { + Some(end) => (&tail[..end], validate::trim_ows(&tail[end + 1..])), + None => (tail, &[][..]), + }; + return Self { + range, + quality: Some(QualityView::from_validated(validate::trim_ows(quality))), + parameters: ParameterSource::Wire(&[]), + extensions: ParameterSource::Wire(extensions), + }; + } + let mut quality = None; + let mut parameter_end = bytes.len(); + let mut extensions = &bytes[bytes.len()..]; + for segment in QuotedItems::semicolon(&bytes[parameter_start..], &FieldName::Accept) { + let segment = segment.expect("header decoding validated every parameter's quoting"); + let equals = segment + .iter() + .position(|byte| *byte == b'=') + .expect("validated media parameters before quality require values"); + if validate::trim_ows(&segment[..equals]).eq_ignore_ascii_case(b"q") { + let start = segment.as_ptr().addr() - bytes.as_ptr().addr(); + parameter_end = bytes[..start] + .iter() + .rposition(|byte| *byte == b';') + .expect("a quality parameter follows a semicolon"); + let after_quality = start + segment.len(); + if let Some(delimiter) = bytes[after_quality..].iter().position(|byte| *byte == b';') { + extensions = &bytes[after_quality + delimiter + 1..]; + } + quality = Some(QualityView::from_validated(validate::trim_ows(&segment[equals + 1..]))); + break; + } + } + let parameters = if parameter_end < parameter_start { + &bytes[0..0] + } else { + validate::trim_ows(&bytes[parameter_start..parameter_end]) + }; + Self { + range, + quality, + parameters: ParameterSource::Wire(parameters), + extensions: ParameterSource::Wire(validate::trim_ows(extensions)), + } + } + + /// Returns the validated media range. + #[must_use] + pub const fn range(self) -> MediaRange<'a> { + self.range + } + + /// Returns the exact effective quality, defaulting to one. + #[must_use] + pub const fn quality(self) -> QualityView<'a> { + match self.quality { + Some(quality) => quality, + None => QualityView::ONE, + } + } + + /// Returns only an explicitly supplied quality. + #[must_use] + pub const fn explicit_quality(self) -> Option> { + self.quality + } + + /// Iterates media parameters before `q`, preserving order and duplicates. + #[must_use] + pub fn parameters(self) -> NegotiationParameters<'a> { + self.parameters.iter() + } + + /// Iterates extensions after `q`, including extensions without a value. + #[must_use] + pub fn extensions(self) -> NegotiationParameters<'a> { + self.extensions.iter() + } + + pub(super) fn encoded_len(&self) -> Result { + let components = [ + self.range.type_().as_str().len(), + 1, + self.range.subtype().as_str().len(), + weight_len(self.quality), + ]; + checked_len( + components.into_iter().chain( + self.parameters() + .chain(self.extensions()) + .map(|parameter| 1 + parameter.encoded_len()), + ), + &FieldName::Accept, + ) + } + + pub(super) fn append_to(&self, bytes: &mut Vec) { + bytes.extend_from_slice(self.range.type_().as_str().as_bytes()); + bytes.push(b'/'); + bytes.extend_from_slice(self.range.subtype().as_str().as_bytes()); + for parameter in self.parameters() { + parameter.append_to(bytes); + } + append_weight(bytes, self.quality); + for extension in self.extensions() { + extension.append_to(bytes); + } + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod quality_first_tests { + use crate::headers::Accept; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeErrorKind, DecodeMode, Field, FieldName}; + + struct Source<'a>(&'a [u8]); + + impl FieldSource for Source<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::Accept).then(|| FieldLines::single(name, self.0)) + } + } + + fn compare_to_general(bytes: &[u8], mode: DecodeMode, expected_quality: &str, explicit: bool) { + let source = Source(bytes); + let decoded = Accept::view_with(&source, mode).unwrap().unwrap(); + let actual = decoded.entries().next().unwrap(); + let head_end = bytes.iter().position(|byte| *byte == b';').unwrap_or(bytes.len()); + let mut general_wire = bytes[..head_end].to_vec(); + // A leading media parameter forces the existing general projection. + general_wire.extend_from_slice(b";comparison-probe=present"); + general_wire.extend_from_slice(&bytes[head_end..]); + let general_source = Source(&general_wire); + let general_decoded = Accept::view_with(&general_source, mode).unwrap().unwrap(); + let general = general_decoded.entries().next().unwrap(); + assert_eq!(actual.range(), general.range()); + assert_eq!(actual.quality(), general.quality()); + assert_eq!(actual.quality().to_string(), expected_quality); + assert_eq!(actual.explicit_quality(), general.explicit_quality()); + assert_eq!(actual.explicit_quality().is_some(), explicit); + let mut parameters = general.parameters(); + let probe = parameters.next().unwrap(); + assert_eq!(probe.name().as_str(), "comparison-probe"); + assert_eq!(probe.value().unwrap().raw_bytes(), b"present"); + assert_eq!(actual.parameters().collect::>(), parameters.collect::>()); + assert_eq!(actual.extensions().collect::>(), general.extensions().collect::>()); + } + + #[test] + fn quality_first_and_general_projection_agree() { + let cases: &[(&[u8], DecodeMode, &str, bool)] = &[ + (b"*/*", DecodeMode::Strict, "1", false), + (b"text/plain;level=\"q=0.5;x\"", DecodeMode::Strict, "1", false), + (b"*/*;q=0", DecodeMode::Strict, "0", true), + (b"text/plain;q=0.", DecodeMode::Strict, "0", true), + (b"text/*;Q=1.", DecodeMode::Strict, "1", true), + ( + b"text/html;q=0.500;Flag;empty=\"\";tag=\"a,b;c\\\"\\\\\xff\";Flag", + DecodeMode::Strict, + "0.5", + true, + ), + (b"text/plain;charset=utf-8;q=0.8;flag", DecodeMode::Strict, "0.8", true), + (b"application/json; q=0.125 ;preview", DecodeMode::Strict, "0.125", true), + (b"application/json;q=.5000;flag", DecodeMode::Relaxed, "0.5", true), + ( + b"text/plain;q= \t.0001000000 \t;note=\"\xff;\\\"x\"", + DecodeMode::Relaxed, + "0.0001", + true, + ), + ( + b"application/json;Q=0.5000000000000000000001;Flag", + DecodeMode::Relaxed, + "0.5000000000000000000001", + true, + ), + (b"Text/HTML; q = .12500 ;Flag", DecodeMode::Relaxed, "0.125", true), + (b"text/*;\tQ\t=\t1.0000;X=\"q=0;d,e\"", DecodeMode::Relaxed, "1", true), + ]; + for &(bytes, mode, expected_quality, explicit) in cases { + compare_to_general(bytes, mode, expected_quality, explicit); + } + } + + #[test] + fn quality_first_retains_extension_order_and_byte_values() { + let source = Source(b"text/html;q=0.500;Flag;empty=\"\";tag=\"a,b;c\\\"\\\\\xff\";Flag"); + let decoded = Accept::view(&source).unwrap().unwrap(); + let entry = decoded.entries().next().unwrap(); + assert_eq!(entry.parameters().count(), 0); + let extensions = entry.extensions().collect::>(); + assert_eq!( + extensions.iter().map(|parameter| parameter.name().as_str()).collect::>(), + ["Flag", "empty", "tag", "Flag"] + ); + assert_eq!(extensions[0].value(), None); + assert_eq!(extensions[1].value().unwrap().to_decoded_bytes(), b""); + assert_eq!(extensions[2].value().unwrap().to_decoded_bytes(), b"a,b;c\"\\\xff"); + assert_eq!(extensions[3].value(), None); + } + + #[test] + fn header_validation_still_rejects_invalid_quality_first_members() { + for (bytes, kind) in [ + (b"text/plain;q=.5".as_slice(), DecodeErrorKind::InvalidSyntax), + (b"text/plain;q=\"0.5\"", DecodeErrorKind::InvalidSyntax), + (b"text/plain;q=0.5;Q=0.4", DecodeErrorKind::InvalidSyntax), + (b"text/plain;flag", DecodeErrorKind::InvalidSyntax), + (b"*/plain;q=0.5", DecodeErrorKind::InvalidSyntax), + (b"text/plain;q=0.5;flag=\"unterminated", DecodeErrorKind::UnterminatedQuote), + ] { + let source = Source(bytes); + let error = Accept::view(&source).unwrap_err(); + assert_eq!(error.kind(), kind); + assert_eq!(error.value_index(), None); + } + } +} diff --git a/crates/http_headers/src/headers/negotiation/accept_language.rs b/crates/http_headers/src/headers/negotiation/accept_language.rs new file mode 100644 index 000000000..c694c7c17 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/accept_language.rs @@ -0,0 +1,270 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::accept_language_entry::AcceptLanguageEntry; +use super::negotiation_members::{collect_members, validated_member}; +use super::shared::{ListValues, QuotedItems, check_quoted_value, check_quoted_values, validate_weighted_token}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef}; + +/// Owned value for the `Accept-Language` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 12.5.4]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::AcceptLanguageOwned::try_from("en, fr;q=0.7")?; +/// assert_eq!(value.items().count(), 2); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Accept-Language: en-US, en;q=0.9, fr;q=0.7` expresses ordered preferences; +/// `Accept-Language: *` accepts any language. +/// +/// # Relaxed decoding +/// +/// [`DecodeMode::Relaxed`](crate::DecodeMode::Relaxed) permits spaces or tabs around the quality +/// parameter's `=`, leading-dot quality values, and more than three +/// fractional digits. Language-range syntax, wildcard syntax, parameter +/// count, quoting, and list structure remain strict, and the original bytes +/// are preserved. +/// +/// [RFC 9110 section 12.5.4]: https://www.rfc-editor.org/rfc/rfc9110#section-12.5.4 +pub struct AcceptLanguageOwned { + values: ListValues, +} + +/// Borrowed value for the `Accept-Language` header. +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{AcceptLanguage, AcceptLanguageView}; +/// use http_headers::source::{FieldLines, FieldSource}; +/// use http_headers::{Field, FieldName}; +/// +/// struct Source; +/// +/// impl FieldSource for Source { +/// fn lines(&self, name: &'static FieldName) -> Option> { +/// (name == &FieldName::AcceptLanguage) +/// .then(|| FieldLines::single(name, b"en-US, fr;q=0.7")) +/// } +/// } +/// +/// let value: AcceptLanguageView<'_> = AcceptLanguage::view(&Source)?.expect("header is present"); +/// let items = value.items().collect::>(); +/// assert_eq!(items, [&b"en-US"[..], &b"fr;q=0.7"[..]]); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct AcceptLanguageView<'a> { + values: FieldLines<'a>, +} + +super::shared::list_header!( + AcceptLanguage, + AcceptLanguageOwned, + AcceptLanguageView, + "Accept-Language", + "Defined by [RFC 9110 section 12.5.4](https://www.rfc-editor.org/rfc/rfc9110#section-12.5.4).", + &FieldName::AcceptLanguage, + validate_accept_language, + validate_accept_language_relaxed, + check_quoted_values, + check_quoted_value, + quoted +); + +impl AcceptLanguageOwned { + /// Iterates typed language preferences in wire order without allocating. + /// + /// Empty list members are ignored and duplicates retained. A fresh + /// iterator traverses the wire again; yielded members retain scalar + /// metadata for repeated reads without parsing. + pub fn entries(&self) -> impl Iterator> { + self.values + .iter() + .flat_map(|value| QuotedItems::comma(value.as_bytes(), &FieldName::AcceptLanguage)) + .map(validated_member) + .map(AcceptLanguageEntry::from_validated) + } + + /// Constructs one field line from typed basic language preferences. + /// + /// Order, duplicate ranges and original range spellings are retained. + /// Quality spelling and separators are written in canonical form. Empty input creates + /// a present empty list. + /// + /// # Errors + /// + /// Returns an error if the aggregate wire size or member count exceeds the + /// custom-source budgets. + pub fn from_entries<'a>(entries: impl IntoIterator>) -> Result { + collect_members( + entries, + &FieldName::AcceptLanguage, + AcceptLanguageEntry::encoded_len, + AcceptLanguageEntry::append_to, + ) + .map(|values| Self { values }) + } +} + +impl<'a> AcceptLanguageView<'a> { + /// Iterates typed basic language ranges and qualities without allocating. + /// + /// Each new iterator traverses the wire again; scalar entry getters reuse + /// retained metadata rather than parsing. + pub fn entries(&self) -> impl Iterator> + '_ { + self.values + .validated_comma_items() + .map(validated_member) + .map(AcceptLanguageEntry::from_validated) + } +} + +fn validate_accept_language(bytes: &[u8]) -> Result<(), DecodeError> { + validate_weighted_token(bytes, &FieldName::AcceptLanguage, valid_language_range, false) +} + +fn validate_accept_language_relaxed(bytes: &[u8]) -> Result<(), DecodeError> { + validate_weighted_token(bytes, &FieldName::AcceptLanguage, valid_language_range, true) +} + +pub(super) fn valid_language_range(bytes: &[u8]) -> bool { + if bytes == b"*" { + return true; + } + let mut subtags = bytes.split(|byte| *byte == b'-'); + let primary = subtags.next().expect("slice splitting yields a first subtag"); + (1..=8).contains(&primary.len()) + && primary.iter().all(u8::is_ascii_alphabetic) + && subtags.all(|subtag| (1..=8).contains(&subtag.len()) && subtag.iter().all(u8::is_ascii_alphanumeric)) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::super::recognition_test_support::{self, exhaustive, fragment_lines}; + use super::super::shared::WELL_KNOWN_ACCEPT_LANGUAGE; + use super::super::weighted_token_scan::scan_accept_language_line; + use super::{valid_language_range, validate_accept_language, validate_accept_language_relaxed}; + use crate::DecodeErrorKind; + + #[test] + fn language_ranges_enforce_primary_and_subtag_boundaries() { + validate_accept_language(b"en-US;q=0.8").expect("strict language weight"); + validate_accept_language_relaxed(b"fr; q = .1234").expect("relaxed language weight"); + assert_eq!( + validate_accept_language(b"en_US") + .expect_err("underscore is not a language-range separator") + .kind(), + DecodeErrorKind::InvalidToken + ); + for valid in [b"*".as_slice(), b"en", b"en-US", b"abcdefgh-12345678"] { + assert!(valid_language_range(valid), "{valid:?}"); + } + for invalid in [b"".as_slice(), b"-en", b"abcdefghi", b"en-", b"en-123456789", b"en-US!"] { + assert!(!valid_language_range(invalid), "{invalid:?}"); + } + } + + #[test] + fn every_well_known_line_also_satisfies_the_member_grammar() { + let rejected: Vec<_> = WELL_KNOWN_ACCEPT_LANGUAGE + .iter() + .filter(|line| validate_accept_language(line).is_err()) + .map(|line| String::from_utf8_lossy(line).into_owned()) + .collect(); + assert!( + rejected.is_empty(), + "recognized lines must satisfy the grammar they skip: {rejected:?}" + ); + } + + fn assert_recognition_is_sound(lines: &[Vec]) -> usize { + recognition_test_support::assert_recognition_is_sound( + lines, + scan_accept_language_line, + validate_accept_language, + validate_accept_language_relaxed, + ) + } + + #[test] + fn recognized_short_lines_satisfy_the_member_grammar() { + let recognized = assert_recognition_is_sound(&exhaustive(b"a-*;,=qQ01. ", 4)); + assert!( + recognized > 100, + "the corpus must actually exercise the recognizer, saw {recognized}" + ); + } + + #[test] + fn recognized_fragment_lines_satisfy_the_member_grammar() { + const FRAGMENTS: &[&str] = &[ + "en", + "en-US", + "*", + "en-", + "-en", + "en_US", + "abcdefghi", + "en-abcdefghi", + ";q=0.9", + ";q=1", + ";q=0", + ";q=", + ";q=1.5", + ";q=0.1234", + ";q=1.000", + ";Q=0.123", + ";x=y", + "; q=0.9", + " ;q=0.9", + ";\tq=0.9", + ";q=0.9 ", + "; ", + " , ", + ",", + " ", + "\t", + ";", + "=", + "\"", + "\\", + "a", + "0.", + ]; + + let recognized = assert_recognition_is_sound(&fragment_lines(FRAGMENTS, 3)); + assert!( + recognized > 100, + "the corpus must actually exercise the recognizer, saw {recognized}" + ); + } + + #[test] + fn realistic_lines_are_recognized() { + const LINES: &[&str] = &[ + "*", + "en", + "en-US", + "en-US,en;q=0.9", + "en-US,en;q=0.9,fr;q=0.8", + "en-US,en;q=0.9,fr-FR;q=0.8,fr;q=0.7", + "fr-CH, fr;q=0.9, en;q=0.8, de;q=0.7, *;q=0.5", + "zh-Hans-CN,zh-Hans;q=0.9,en-US;q=0.8,en;q=0.7", + ]; + + for line in LINES { + assert!( + scan_accept_language_line(line.as_bytes()), + "the recognizer must settle the common line {line:?}" + ); + } + } +} diff --git a/crates/http_headers/src/headers/negotiation/accept_language_entry.rs b/crates/http_headers/src/headers/negotiation/accept_language_entry.rs new file mode 100644 index 000000000..f9cb035cc --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/accept_language_entry.rs @@ -0,0 +1,69 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::language_range::LanguageRange; +use super::negotiation_members::{append_weight, checked_len, weight_len, weighted_parts}; +use super::quality::QualityView; +use crate::{DecodeError, FieldName}; + +/// One validated Accept-Language preference with retained range and quality. +/// +/// The basic language-range grammar does not perform locale selection, +/// normalization or fallback. Missing quality remains distinct from `q=1`. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct AcceptLanguageEntry<'a> { + range: LanguageRange<'a>, + quality: Option>, +} + +impl<'a> AcceptLanguageEntry<'a> { + /// Constructs a language preference from validated components. + /// + /// # Errors + /// + /// Returns an error if its serialized size exceeds + /// [`MAX_CUSTOM_FIELD_BYTES`](crate::source::MAX_CUSTOM_FIELD_BYTES). + pub fn new(range: LanguageRange<'a>, quality: Option>) -> Result { + let entry = Self { range, quality }; + entry.encoded_len()?; + Ok(entry) + } + + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + let (range, quality) = weighted_parts(bytes); + Self { + range: LanguageRange::from_validated(range), + quality, + } + } + + /// Returns the validated basic language range. + #[must_use] + pub const fn range(self) -> LanguageRange<'a> { + self.range + } + + /// Returns the exact quality, defaulting to one when `q` was omitted. + #[must_use] + pub const fn quality(self) -> QualityView<'a> { + match self.quality { + Some(quality) => quality, + None => QualityView::ONE, + } + } + + /// Returns only an explicitly supplied quality. + #[must_use] + pub const fn explicit_quality(self) -> Option> { + self.quality + } + + pub(super) fn encoded_len(&self) -> Result { + checked_len([self.range.as_str().len(), weight_len(self.quality)], &FieldName::AcceptLanguage) + } + + pub(super) fn append_to(&self, bytes: &mut Vec) { + bytes.extend_from_slice(self.range.as_str().as_bytes()); + append_weight(bytes, self.quality); + } +} diff --git a/crates/http_headers/src/headers/negotiation/accept_scan.rs b/crates/http_headers/src/headers/negotiation/accept_scan.rs new file mode 100644 index 000000000..261e852a1 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/accept_scan.rs @@ -0,0 +1,411 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Whole-line recognizer for the common shape of an `Accept` field line. + +use crate::validate::token_byte; + +/// The byte classes an `Accept` field line is built from. +/// +/// Dense indices keep the table small, and padding the row to sixteen lets the +/// scan form an index with a shift instead of a multiply. +const CLASS_OTHER: u8 = 0; +const CLASS_TCHAR: u8 = 1; +const CLASS_Q: u8 = 2; +const CLASS_ZERO: u8 = 3; +const CLASS_ONE: u8 = 4; +const CLASS_DIGIT: u8 = 5; +const CLASS_DOT: u8 = 6; +const CLASS_STAR: u8 = 7; +const CLASS_SLASH: u8 = 8; +const CLASS_SEMICOLON: u8 = 9; +const CLASS_COMMA: u8 = 10; +const CLASS_EQUALS: u8 = 11; +const CLASS_OWS: u8 = 12; + +const CLASSES: usize = 16; +const STATES: usize = 32; + +/// The line left the recognized subset. The state absorbs every byte, so the +/// scan needs no branch to leave the loop. +const REJECTED: u8 = 0; +/// A member is due: the line just started or a comma just closed one. +const MEMBER_DUE: u8 = 1; +/// The type read so far is exactly `*`. +const TYPE_STAR: u8 = 2; +/// A type is in progress and is not the bare wildcard. +const TYPE: u8 = 3; +/// A slash closed a `*` type, so only a `*` subtype may follow. +const SUBTYPE_DUE_STAR: u8 = 4; +/// A slash closed a named type. +const SUBTYPE_DUE: u8 = 5; +/// The subtype read so far is exactly `*` and the type was `*`. +const SUBTYPE_STAR: u8 = 6; +/// A subtype is in progress. +const SUBTYPE: u8 = 7; +/// Whitespace closed a complete member, so only `;`, `,`, or the end may follow. +const MEMBER_OWS: u8 = 8; +/// A semicolon opened a parameter. +const PARAM_DUE: u8 = 9; +/// The parameter name read so far is exactly `q`. +const NAME_Q: u8 = 10; +/// A parameter name is in progress and is not the bare `q`. +const NAME: u8 = 11; +/// An equals sign closed an ordinary parameter name. +const VALUE_DUE: u8 = 12; +/// An ordinary parameter value is in progress. +const VALUE: u8 = 13; +/// An equals sign closed the `q` name, so a quality value is due. +const QUALITY_DUE: u8 = 14; +/// The quality value read so far is `0`. +const QUALITY_ZERO: u8 = 15; +/// The quality value read so far is `1`. +const QUALITY_ONE: u8 = 16; +/// The quality value read so far is `0.` with no fraction digit yet. +const QUALITY_ZERO_DOT: u8 = 17; +/// `0.` followed by one fraction digit. +const QUALITY_ZERO_ONE_DIGIT: u8 = 18; +/// `0.` followed by two fraction digits. +const QUALITY_ZERO_TWO_DIGITS: u8 = 19; +/// `0.` followed by three fraction digits, the most the grammar allows. +const QUALITY_ZERO_THREE_DIGITS: u8 = 20; +/// The quality value read so far is `1.` with no fraction digit yet. +const QUALITY_ONE_DOT: u8 = 21; +/// `1.` followed by one fraction zero. +const QUALITY_ONE_ONE_DIGIT: u8 = 22; +/// `1.` followed by two fraction zeros. +const QUALITY_ONE_TWO_DIGITS: u8 = 23; +/// `1.` followed by three fraction zeros, the most the grammar allows. +const QUALITY_ONE_THREE_DIGITS: u8 = 24; +/// Whitespace closed a member that already carries a quality parameter. +const MEMBER_OWS_WEIGHED: u8 = 25; +/// A semicolon opened a parameter in a member that already carries a quality. +const PARAM_DUE_WEIGHED: u8 = 26; +/// A bare `q` name in a member that already carries a quality: a repeat. +const NAME_Q_WEIGHED: u8 = 27; +/// A parameter name in a member that already carries a quality. +const NAME_WEIGHED: u8 = 28; +/// An equals sign closed a parameter name after a quality. +const VALUE_DUE_WEIGHED: u8 = 29; +/// A parameter value is in progress after a quality. +const VALUE_WEIGHED: u8 = 30; + +/// The states that end a line inside the recognized subset. +/// +/// A member must have reached a complete subtype, a complete parameter value, +/// or a complete quality, and the line may also end between members. +const ACCEPTING: u32 = (1 << MEMBER_DUE) + | (1 << SUBTYPE_STAR) + | (1 << SUBTYPE) + | (1 << MEMBER_OWS) + | (1 << VALUE) + | (1 << QUALITY_ZERO) + | (1 << QUALITY_ONE) + | (1 << QUALITY_ZERO_DOT) + | (1 << QUALITY_ZERO_ONE_DIGIT) + | (1 << QUALITY_ZERO_TWO_DIGITS) + | (1 << QUALITY_ZERO_THREE_DIGITS) + | (1 << QUALITY_ONE_DOT) + | (1 << QUALITY_ONE_ONE_DIGIT) + | (1 << QUALITY_ONE_TWO_DIGITS) + | (1 << QUALITY_ONE_THREE_DIGITS) + | (1 << MEMBER_OWS_WEIGHED) + | (1 << VALUE_WEIGHED); + +/// Maps each byte to its role inside an `Accept` field line. +static CLASS: [u8; 256] = class_table(); + +/// Maps a row and a byte class to the next row. +/// +/// A row is a state already multiplied by [`CLASSES`], so a step indexes the +/// table with `row | class` and stores the entry back unchanged. Keeping the +/// multiply in the table removes a shift from every byte of the scan. +static TRANSITION: [u16; STATES * CLASSES] = transition_table(); + +/// The row a state occupies in [`TRANSITION`]. +const fn row(state: u8) -> u16 { + (state as u16) << 4 +} + +const fn class_table() -> [u8; 256] { + let mut table = [CLASS_OTHER; 256]; + let mut byte = 0_usize; + while byte < 256 { + #[expect(clippy::cast_possible_truncation, reason = "the loop bound keeps the index inside a byte")] + let value = byte as u8; + table[byte] = match value { + b'q' | b'Q' => CLASS_Q, + b'0' => CLASS_ZERO, + b'1' => CLASS_ONE, + b'2'..=b'9' => CLASS_DIGIT, + b'.' => CLASS_DOT, + b'*' => CLASS_STAR, + b'/' => CLASS_SLASH, + b';' => CLASS_SEMICOLON, + b',' => CLASS_COMMA, + b'=' => CLASS_EQUALS, + b' ' | b'\t' => CLASS_OWS, + _ if token_byte(value) => CLASS_TCHAR, + _ => CLASS_OTHER, + }; + byte += 1; + } + table +} + +/// Returns whether a class is one of the `tchar` classes. +const fn is_token_class(class: u8) -> bool { + matches!( + class, + CLASS_TCHAR | CLASS_Q | CLASS_ZERO | CLASS_ONE | CLASS_DIGIT | CLASS_DOT | CLASS_STAR + ) +} + +#[expect( + clippy::too_many_lines, + reason = "one arm per state keeps the whole transition relation readable in one place" +)] +const fn transition_table() -> [u16; STATES * CLASSES] { + let mut table = [row(REJECTED); STATES * CLASSES]; + let mut state = 0_usize; + while state < STATES { + let mut class = 0_usize; + while class < CLASSES { + #[expect(clippy::cast_possible_truncation, reason = "the loop bounds keep both indices inside a byte")] + let (state_value, class_value) = (state as u8, class as u8); + let token = is_token_class(class_value); + table[(state << 4) | class] = row(match state_value { + MEMBER_DUE => match class_value { + CLASS_OWS | CLASS_COMMA => MEMBER_DUE, + CLASS_STAR => TYPE_STAR, + _ if token => TYPE, + _ => REJECTED, + }, + TYPE_STAR => match class_value { + CLASS_SLASH => SUBTYPE_DUE_STAR, + _ if token => TYPE, + _ => REJECTED, + }, + TYPE => match class_value { + CLASS_SLASH => SUBTYPE_DUE, + _ if token => TYPE, + _ => REJECTED, + }, + SUBTYPE_DUE_STAR => match class_value { + CLASS_STAR => SUBTYPE_STAR, + _ => REJECTED, + }, + SUBTYPE_DUE => { + if token { + SUBTYPE + } else { + REJECTED + } + } + SUBTYPE => match class_value { + CLASS_OWS => MEMBER_OWS, + CLASS_SEMICOLON => PARAM_DUE, + CLASS_COMMA => MEMBER_DUE, + _ if token => SUBTYPE, + _ => REJECTED, + }, + SUBTYPE_STAR | MEMBER_OWS => match class_value { + CLASS_OWS => MEMBER_OWS, + CLASS_SEMICOLON => PARAM_DUE, + CLASS_COMMA => MEMBER_DUE, + _ => REJECTED, + }, + PARAM_DUE => match class_value { + CLASS_OWS => PARAM_DUE, + CLASS_Q => NAME_Q, + _ if token => NAME, + _ => REJECTED, + }, + NAME_Q => match class_value { + CLASS_EQUALS => QUALITY_DUE, + _ if token => NAME, + _ => REJECTED, + }, + NAME => match class_value { + CLASS_EQUALS => VALUE_DUE, + _ if token => NAME, + _ => REJECTED, + }, + VALUE_DUE => { + if token { + VALUE + } else { + REJECTED + } + } + VALUE => match class_value { + CLASS_OWS => MEMBER_OWS, + CLASS_SEMICOLON => PARAM_DUE, + CLASS_COMMA => MEMBER_DUE, + _ if token => VALUE, + _ => REJECTED, + }, + QUALITY_DUE => match class_value { + CLASS_ZERO => QUALITY_ZERO, + CLASS_ONE => QUALITY_ONE, + _ => REJECTED, + }, + QUALITY_ZERO => match class_value { + CLASS_DOT => QUALITY_ZERO_DOT, + _ => weighed_close(class_value), + }, + QUALITY_ONE => match class_value { + CLASS_DOT => QUALITY_ONE_DOT, + _ => weighed_close(class_value), + }, + QUALITY_ZERO_DOT => match class_value { + CLASS_ZERO | CLASS_ONE | CLASS_DIGIT => QUALITY_ZERO_ONE_DIGIT, + _ => weighed_close(class_value), + }, + QUALITY_ZERO_ONE_DIGIT => match class_value { + CLASS_ZERO | CLASS_ONE | CLASS_DIGIT => QUALITY_ZERO_TWO_DIGITS, + _ => weighed_close(class_value), + }, + QUALITY_ZERO_TWO_DIGITS => match class_value { + CLASS_ZERO | CLASS_ONE | CLASS_DIGIT => QUALITY_ZERO_THREE_DIGITS, + _ => weighed_close(class_value), + }, + QUALITY_ONE_DOT => match class_value { + CLASS_ZERO => QUALITY_ONE_ONE_DIGIT, + _ => weighed_close(class_value), + }, + QUALITY_ONE_ONE_DIGIT => match class_value { + CLASS_ZERO => QUALITY_ONE_TWO_DIGITS, + _ => weighed_close(class_value), + }, + QUALITY_ONE_TWO_DIGITS => match class_value { + CLASS_ZERO => QUALITY_ONE_THREE_DIGITS, + _ => weighed_close(class_value), + }, + QUALITY_ZERO_THREE_DIGITS | QUALITY_ONE_THREE_DIGITS | MEMBER_OWS_WEIGHED => weighed_close(class_value), + PARAM_DUE_WEIGHED => match class_value { + CLASS_OWS => PARAM_DUE_WEIGHED, + CLASS_Q => NAME_Q_WEIGHED, + _ if token => NAME_WEIGHED, + _ => REJECTED, + }, + NAME_Q_WEIGHED => { + if token { + NAME_WEIGHED + } else { + REJECTED + } + } + NAME_WEIGHED => match class_value { + CLASS_EQUALS => VALUE_DUE_WEIGHED, + _ if token => NAME_WEIGHED, + _ => REJECTED, + }, + VALUE_DUE_WEIGHED => { + if token { + VALUE_WEIGHED + } else { + REJECTED + } + } + VALUE_WEIGHED => match class_value { + _ if token => VALUE_WEIGHED, + _ => weighed_close(class_value), + }, + _ => REJECTED, + }); + class += 1; + } + state += 1; + } + table +} + +/// Closes a member that already carries a quality parameter. +/// +/// Whitespace and a further parameter stay in the weighed half of the machine +/// so a repeated `q` is rejected, while a comma starts a fresh member. +const fn weighed_close(class: u8) -> u8 { + match class { + CLASS_OWS => MEMBER_OWS_WEIGHED, + CLASS_SEMICOLON => PARAM_DUE_WEIGHED, + CLASS_COMMA => MEMBER_DUE, + _ => REJECTED, + } +} + +/// Advances the table one byte. +/// +/// Every row is already a multiple of `CLASSES` and every class is smaller +/// than `CLASSES`, so the mask changes no index it is given and exists only to +/// prove the bound, which is what lets the step compile to two loads and no +/// branch. +#[expect(clippy::inline_always, reason = "the step is the unrolled loop body and must not become a call")] +#[inline(always)] +fn step(row: u16, byte: u8) -> u16 { + let class = CLASS[usize::from(byte)]; + let index = (usize::from(row) | usize::from(class)) & (STATES * CLASSES - 1); + TRANSITION[index] +} + +/// Reports whether a whole `Accept` field line is valid under strict decoding. +/// +/// The line is read once and every byte drives one table step, so the common +/// shape — comma separated media ranges carrying token parameters and a +/// quality — is settled without splitting the line into members and walking +/// each one. Stepping eight bytes per iteration amortizes the loop counter and +/// branch, which otherwise cost about as much as the two table loads they +/// carry. +/// +/// A `false` answer means "not recognized" rather than "malformed": quoted +/// strings, whitespace around an equals sign, and parameters without values +/// are all legal yet left to the general parser, which also produces the +/// diagnostic when the line really is malformed. +pub(super) fn scan_accept_line(bytes: &[u8]) -> bool { + let mut row = row(MEMBER_DUE); + + let mut blocks = bytes.chunks_exact(8); + for block in &mut blocks { + for byte in block { + row = step(row, *byte); + } + } + for byte in blocks.remainder() { + row = step(row, *byte); + } + + ACCEPTING & (1 << (row >> 4)) != 0 +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::hint::black_box; + + use super::{ACCEPTING, CLASS, CLASSES, REJECTED, STATES, TRANSITION, class_table, row, transition_table}; + + #[test] + fn runtime_tables_match_the_static_tables() { + assert_eq!(black_box(class_table()), CLASS); + + let generated = black_box(transition_table()); + assert_eq!(generated, TRANSITION); + + for class in 0..CLASSES { + assert_eq!( + generated[usize::from(row(REJECTED)) | class], + row(REJECTED), + "rejection must absorb every later byte" + ); + } + for entry in generated { + assert!(usize::from(entry >> 4) < STATES, "every transition must land on a defined state"); + assert_eq!(entry & 0xf, 0, "every entry must be a state already multiplied by the class count"); + } + assert_eq!( + ACCEPTING & (1 << REJECTED), + 0, + "a rejected line must never end inside the recognized subset" + ); + } +} diff --git a/crates/http_headers/src/headers/negotiation/allow.rs b/crates/http_headers/src/headers/negotiation/allow.rs new file mode 100644 index 000000000..54af453af --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/allow.rs @@ -0,0 +1,132 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::MethodView; +use super::shared::{ListValues, check_token_value, check_token_values, invalid}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Owned value for the `Allow` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 10.2.1]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::AllowOwned::try_from("GET, POST")?; +/// assert_eq!(value.items().count(), 2); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Allow: GET, HEAD, OPTIONS` lists supported methods. An empty `Allow:` +/// field is valid and indicates that no methods are currently allowed. +/// +/// [RFC 9110 section 10.2.1]: https://www.rfc-editor.org/rfc/rfc9110#section-10.2.1 +pub struct AllowOwned { + values: ListValues, +} + +/// Borrowed value for the `Allow` header. +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{Allow, AllowView}; +/// use http_headers::source::{FieldLines, FieldSource}; +/// use http_headers::{Field, FieldName}; +/// +/// struct Source; +/// +/// impl FieldSource for Source { +/// fn lines(&self, name: &'static FieldName) -> Option> { +/// (name == &FieldName::Allow).then(|| FieldLines::single(name, b"GET, HEAD, OPTIONS")) +/// } +/// } +/// +/// let value: AllowView<'_> = Allow::view(&Source)?.expect("header is present"); +/// let items = value.items().collect::>(); +/// assert_eq!(items, [&b"GET"[..], &b"HEAD"[..], &b"OPTIONS"[..]]); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct AllowView<'a> { + values: FieldLines<'a>, +} + +super::shared::list_header!( + Allow, + AllowOwned, + AllowView, + "Allow", + "Defined by [RFC 9110 section 10.2.1](https://www.rfc-editor.org/rfc/rfc9110#section-10.2.1).", + &FieldName::Allow, + validate_allow_item, + validate_allow_item, + check_token_values, + check_token_value, + token +); + +impl AllowOwned { + /// Iterates validated, case-sensitive methods in wire order. + /// + /// Empty list members are ignored. Iteration does not allocate or + /// revalidate token syntax. + #[inline] + pub fn methods(&self) -> impl Iterator> { + self.items().map(MethodView::from_validated) + } + + /// Constructs one field line from validated method tokens. + /// + /// Order and duplicates are retained; an empty iterator produces a valid + /// empty Allow field. Formatting does not reparse the method grammar. + #[must_use] + pub fn from_methods<'a>(methods: impl IntoIterator>) -> Self { + let mut wire = String::new(); + for method in methods { + if !wire.is_empty() { + wire.push_str(", "); + } + wire.push_str(method.as_str()); + } + Self { + values: ListValues::One(FieldValue::from_validated_owned_bytes(wire.into_bytes(), false)), + } + } +} + +impl<'a> AllowView<'a> { + /// Iterates validated methods without allocating or revalidating tokens. + /// + /// Methods retain their case-sensitive spelling and wire order. + #[inline] + pub fn methods(&self) -> impl Iterator> + '_ { + self.items().map(MethodView::from_validated) + } +} + +fn validate_allow_item(bytes: &[u8]) -> Result<(), DecodeError> { + if validate::token(bytes) { + Ok(()) + } else { + Err(invalid(&FieldName::Allow, DecodeErrorKind::InvalidToken)) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::validate_allow_item; + use crate::DecodeErrorKind; + + #[test] + fn allow_items_are_bare_method_tokens() { + validate_allow_item(b"PATCH").expect("extension method token"); + assert_eq!( + validate_allow_item(b"BAD METHOD").expect_err("spaces are not token bytes").kind(), + DecodeErrorKind::InvalidToken + ); + } +} diff --git a/crates/http_headers/src/headers/negotiation/content_coding.rs b/crates/http_headers/src/headers/negotiation/content_coding.rs new file mode 100644 index 000000000..1ecd287d9 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/content_coding.rs @@ -0,0 +1,113 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::negotiation_token::NegotiationToken; +use super::shared::invalid; +use crate::{DecodeError, DecodeErrorKind, FieldName, validate}; + +/// The semantic classification of a content coding. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +#[non_exhaustive] +pub enum ContentCodingKind { + /// The standalone wildcard `*`. + Wildcard, + /// The representation without content encoding, `identity`. + Identity, + /// The `gzip` coding. + Gzip, + /// The `compress` coding. + Compress, + /// The `deflate` coding. + Deflate, + /// The Brotli `br` coding. + Br, + /// The Zstandard `zstd` coding. + Zstd, + /// The dictionary-compressed Brotli `dcb` coding. + Dcb, + /// The dictionary-compressed Zstandard `dcz` coding. + Dcz, + /// Another validated coding token, including embedded stars. + Extension, +} + +/// A validated content coding with case-insensitive equality. +/// +/// Wildcard and identity are distinct; no implicit coding entries or codec +/// selection policies are applied. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct ContentCoding<'a> { + token: NegotiationToken<'a>, + kind: ContentCodingKind, +} + +impl<'a> ContentCoding<'a> { + /// Validates a coding token and classifies known spellings. + /// + /// # Errors + /// + /// Returns an error for an empty or malformed HTTP token. + pub fn parse(value: &'a str) -> Result { + if !validate::token(value.as_bytes()) { + return Err(invalid(&FieldName::AcceptEncoding, DecodeErrorKind::InvalidToken)); + } + Ok(Self::new(NegotiationToken::from_validated(value.as_bytes()))) + } + + /// Classifies a validated coding token without changing its spelling. + #[must_use] + pub fn new(token: NegotiationToken<'a>) -> Self { + let kind = if token.as_str() == "*" { + ContentCodingKind::Wildcard + } else if token.eq_ignore_ascii_case("identity") { + ContentCodingKind::Identity + } else if token.eq_ignore_ascii_case("gzip") { + ContentCodingKind::Gzip + } else if token.eq_ignore_ascii_case("compress") { + ContentCodingKind::Compress + } else if token.eq_ignore_ascii_case("deflate") { + ContentCodingKind::Deflate + } else if token.eq_ignore_ascii_case("br") { + ContentCodingKind::Br + } else if token.eq_ignore_ascii_case("zstd") { + ContentCodingKind::Zstd + } else if token.eq_ignore_ascii_case("dcb") { + ContentCodingKind::Dcb + } else if token.eq_ignore_ascii_case("dcz") { + ContentCodingKind::Dcz + } else { + ContentCodingKind::Extension + }; + Self { token, kind } + } + + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + Self::new(NegotiationToken::from_validated(bytes)) + } + + /// Returns the retained coding classification. + #[must_use] + pub const fn kind(self) -> ContentCodingKind { + self.kind + } + + /// Returns the validated token with its original spelling. + #[must_use] + pub const fn token(self) -> NegotiationToken<'a> { + self.token + } + + /// Returns the original coding spelling. + #[must_use] + pub const fn as_str(self) -> &'a str { + self.token.as_str() + } +} + +impl fmt::Display for ContentCoding<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.token.fmt(f) + } +} diff --git a/crates/http_headers/src/headers/negotiation/host.rs b/crates/http_headers/src/headers/negotiation/host.rs new file mode 100644 index 000000000..3544a59c6 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/host.rs @@ -0,0 +1,1175 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::borrow::Cow; +use std::fmt::{self, Write as _}; +use std::net::{Ipv4Addr, Ipv6Addr}; +use std::ops::Range; +use std::str::{self, FromStr as _}; + +use idna::domain_to_ascii; + +use super::shared::{invalid, invalid_syntax}; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, SingleValueField, validate}; + +mod components; +pub use components::{HostKind, HostPortView, IpvFutureView, PortConversionError, PortConversionErrorKind, RegisteredNameView}; + +/// Defines the `Host` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 7.2](https://www.rfc-editor.org/rfc/rfc9110#section-7.2). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{Host, HostOwned}; +/// +/// let mut map = HeaderMap::new(); +/// Host::insert(&mut map, HostOwned::try_from("example.com:8080")?)?; +/// assert!(Host::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct Host { + _private: (), +} + +/// Owned value for the `Host` header. +/// +/// Debug output omits authority components when the field value is sensitive. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 7.2]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::HostOwned::try_from("example.com:8080")?; +/// assert_eq!(value.host()?, "example.com"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Host: example.com` names a host, `Host: example.com:8080` includes a port, +/// and `Host: [2001:db8::1]:443` uses an IPv6 literal. An empty field value +/// represents a target URI without an authority component. +/// +/// [RFC 9110 section 7.2]: https://www.rfc-editor.org/rfc/rfc9110#section-7.2 +#[derive(Clone, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct HostOwned { + value: FieldValue, + parsed: ParsedHost, +} + +/// A view of the `Host` header with retained validated components. +/// +/// Ordinary ASCII values borrow their storage without allocating. Relaxed +/// international names retain an owned IDNA normalization, separately from the +/// borrowed original bytes. [`HostOwned::as_view`] borrows that normalization +/// instead of copying it. +/// +/// Debug output omits authority components when the field value is sensitive. +#[derive(Clone, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{Host, HostView}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let view: HostView<'_> = +/// ::decode_view(FieldValueRef::new(b"example.com:443"))?; +/// assert_eq!(view.host(), "example.com"); +/// assert_eq!(view.port(), Some("443")); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct HostView<'a> { + value: FieldValueRef<'a>, + host: &'a str, + port: Option<&'a str>, + kind: ParsedHostKind, + normalized: Option>, + numeric_port: Option>, +} + +impl fmt::Debug for HostOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut debug = f.debug_struct("HostOwned"); + debug.field("value", &self.value); + if self.value.is_sensitive() { + debug.finish_non_exhaustive() + } else { + debug.field("parsed", &self.parsed).finish() + } + } +} + +impl fmt::Debug for HostView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut debug = f.debug_struct("HostView"); + debug.field("value", &self.value); + if self.value.is_sensitive() { + debug.finish_non_exhaustive() + } else { + debug + .field("host", &self.host) + .field("port", &self.port) + .field("kind", &self.kind) + .field("normalized", &self.normalized) + .field("numeric_port", &self.numeric_port) + .finish() + } + } +} + +impl HostOwned { + /// Constructs an authority from validated components without parsing them again. + /// + /// Registered names preserve their spelling and retained normalization. + /// IP addresses use their standard textual representation; `IPvFuture` uses + /// a lowercase `v` prefix and preserves its version and address spelling. + /// + /// # Errors + /// + /// Returns an error if a port is supplied for an empty registered name, or + /// if adding a port to a relaxed international name would introduce a + /// second port delimiter in its retained IDNA normalization. + #[expect( + clippy::missing_panics_doc, + reason = "validated components and String formatting cannot violate field-value invariants" + )] + pub fn from_parts(host: HostKind<'_>, port: Option>) -> Result { + if let HostKind::RegisteredName(name) = host + && port.is_some() + { + if name.as_str().is_empty() { + return Err(invalid_syntax(&FieldName::Host)); + } + if name.normalized() != name.as_str() { + let normalized = name.normalized(); + if normalized.starts_with('[') { + if !normalized.ends_with(']') { + return Err(invalid(&FieldName::Host, DecodeErrorKind::InvalidNumber)); + } + } else if normalized.contains(':') { + return Err(invalid_syntax(&FieldName::Host)); + } + } + } + let host_capacity = match host { + HostKind::RegisteredName(name) => name.as_str().len(), + HostKind::Ipv4(_) => 15, + HostKind::Ipv6(_) => 41, + HostKind::IpvFuture(address) => address.version().len().saturating_add(address.address().len()).saturating_add(4), + }; + let port_capacity = port.map_or(0, |port| port.as_str().len().saturating_add(1)); + let mut wire = String::with_capacity(host_capacity.saturating_add(port_capacity)); + let mut normalized = None; + let kind = match host { + HostKind::RegisteredName(name) => { + wire.push_str(name.as_str()); + if name.normalized() != name.as_str() { + normalized = Some(name.normalized().to_owned()); + } + ParsedHostKind::RegisteredName + } + HostKind::Ipv4(address) => { + write!(wire, "{address}").expect("writing an IP address into a String cannot fail"); + ParsedHostKind::Ipv4(address) + } + HostKind::Ipv6(address) => { + write!(wire, "[{address}]").expect("writing an IP address into a String cannot fail"); + ParsedHostKind::Ipv6(address) + } + HostKind::IpvFuture(address) => { + write!(wire, "[{address}]").expect("writing a validated IPvFuture address into a String cannot fail"); + ParsedHostKind::IpvFuture { + dot: address.version().len() + 2, + } + } + }; + let host_end = wire.len(); + let port_start = port.map(|port| { + wire.push(':'); + wire.push_str(port.as_str()); + host_end + 1 + }); + Ok(Self { + value: FieldValue::try_from(wire).expect("validated host and decimal port components are valid field value bytes"), + parsed: ParsedHost { + host_end, + port_start, + kind, + normalized, + numeric_port: port.map(HostPortView::to_u16), + }, + }) + } + + /// Constructs an IPv4 authority with an optional checked network port. + #[must_use] + pub fn from_ipv4(address: Ipv4Addr, port: Option) -> Self { + Self::from_address(HostKind::Ipv4(address), port) + } + + /// Constructs a bracketed IPv6 authority with an optional checked network port. + #[must_use] + pub fn from_ipv6(address: Ipv6Addr, port: Option) -> Self { + Self::from_address(HostKind::Ipv6(address), port) + } + + fn from_address(host: HostKind<'_>, port: Option) -> Self { + let mut buffer = itoa::Buffer::new(); + let port = port.map(|port| HostPortView { + text: buffer.format(port), + numeric: Ok(port), + }); + Self::from_parts(host, port).expect("IP address construction cannot introduce an IDNA-normalized port delimiter") + } + + /// Borrows the retained parsing results without repeating validation or IDNA. + #[must_use] + #[inline] + #[expect(clippy::missing_panics_doc, reason = "the offsets are private and validated at construction")] + pub fn as_view(&self) -> HostView<'_> { + HostView { + value: self.value.as_field_value_ref(), + host: self.host().expect("host offsets were validated when constructing this value"), + port: self.port().expect("port offsets were validated when constructing this value"), + kind: self.parsed.kind, + normalized: self.parsed.normalized.as_deref().map(Cow::Borrowed), + numeric_port: self.parsed.numeric_port, + } + } + + /// Returns the validated host kind and retained address or name components. + /// + /// # Examples + /// + /// ``` + /// use std::net::Ipv6Addr; + /// + /// use http_headers::headers::{HostKind, HostOwned}; + /// + /// let value = HostOwned::from_ipv6(Ipv6Addr::LOCALHOST, Some(443)); + /// assert_eq!(value.kind(), HostKind::Ipv6(Ipv6Addr::LOCALHOST)); + /// assert_eq!(value.network_port(), Ok(Some(443))); + /// ``` + #[must_use] + #[inline] + #[expect(clippy::missing_panics_doc, reason = "the offsets are private and validated at construction")] + pub fn kind(&self) -> HostKind<'_> { + self.parsed.kind.project( + self.host().expect("host offsets were validated when constructing this value"), + self.parsed.normalized.as_deref(), + ) + } + + /// Returns the validated textual port, including an explicitly empty port. + #[must_use] + #[inline] + #[expect(clippy::missing_panics_doc, reason = "the offsets are private and validated at construction")] + pub fn port_view(&self) -> Option> { + self.port() + .expect("port offsets were validated when constructing this value") + .zip(self.parsed.numeric_port) + .map(|(text, numeric)| HostPortView { text, numeric }) + } + + /// Returns the checked network port, or `None` only when no port is present. + /// + /// # Errors + /// + /// Returns an error for an explicitly empty port or a value above `u16::MAX`. + /// + /// # Examples + /// + /// ``` + /// use http_headers::headers::{HostOwned, PortConversionErrorKind}; + /// + /// let value = HostOwned::try_from("example.com:65536")?; + /// assert_eq!(value.port()?, Some("65536")); + /// assert_eq!( + /// value.network_port().unwrap_err().kind(), + /// PortConversionErrorKind::Overflow + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[inline] + pub fn network_port(&self) -> Result, PortConversionError> { + self.parsed.numeric_port.transpose() + } + + /// Constructs a host without a port. + /// + /// # Errors + /// + /// Returns an error when `host` is not a valid URI host. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::HostOwned; + /// + /// let value = HostOwned::new("example.com")?; + /// assert_eq!(value.host()?, "example.com"); + /// assert_eq!(value.port()?, None); + /// + /// assert!(HostOwned::new("example.com:443").is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn new(host: impl AsRef) -> Result { + let parsed = Self::try_from(host.as_ref())?; + if parsed.port_start().is_some() { + return Err(invalid_syntax(&FieldName::Host)); + } + Ok(parsed) + } + + /// Constructs a host with a decimal port. + /// + /// IPv6 literals must include their square brackets. + /// + /// # Errors + /// + /// Returns an error when `host` is invalid. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::HostOwned; + /// + /// let value = HostOwned::with_port("[2001:db8::1]", 443)?; + /// assert_eq!(value.host()?, "[2001:db8::1]"); + /// assert_eq!(value.port()?, Some("443")); + /// assert_eq!(value.as_str()?, "[2001:db8::1]:443"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn with_port(host: impl AsRef, port: u16) -> Result { + let host = host.as_ref(); + if host.is_empty() { + return Err(invalid_syntax(&FieldName::Host)); + } + let mut wire = String::with_capacity(host.len().saturating_add(6)); + wire.push_str(host); + wire.push(':'); + write!(&mut wire, "{port}").map_err(|_invalid| invalid_syntax(&FieldName::Host))?; + let mut parsed = match parse_host(host.as_bytes()) { + Ok(parsed) if parsed.port_start.is_none() => parsed, + _ => return Self::try_from(wire), + }; + parsed.port_start = Some(host.len() + 1); + parsed.numeric_port = Some(Ok(port)); + Ok(Self { + value: FieldValue::try_from(wire).map_err(|_invalid| invalid_syntax(&FieldName::Host))?, + parsed, + }) + } + + /// Returns the URI host, including brackets around an IP literal. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the wire value. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::HostOwned::try_from("example.com:443")?; + /// assert_eq!(value.host()?, "example.com"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn host(&self) -> Result<&str, DecodeError> { + component(self.value.as_bytes(), 0..self.parsed.host_end) + } + + /// Returns the optional decimal port. + /// + /// The port is textual because URI syntax does not impose a `u16` range. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the wire value. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::HostOwned; + /// + /// let value = HostOwned::try_from("example.com:8443")?; + /// assert_eq!(value.port()?, Some("8443")); + /// + /// let default_port = HostOwned::try_from("example.com")?; + /// assert_eq!(default_port.port()?, None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn port(&self) -> Result, DecodeError> { + self.port_start() + .map(|start| component(self.value.as_bytes(), start..self.value.as_bytes().len())) + .transpose() + } + + /// Returns the retained port offset. + const fn port_start(&self) -> Option { + self.parsed.port_start + } + + /// Returns the complete authority. + /// + /// # Errors + /// + /// Returns an error if the stored wire value is unexpectedly non-UTF-8. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::HostOwned; + /// + /// let value = HostOwned::with_port("example.com", 8443)?; + /// assert_eq!(value.as_str()?, "example.com:8443"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_str(&self) -> Result<&str, DecodeError> { + str::from_utf8(self.value.as_bytes()).map_err(|_invalid| invalid(&FieldName::Host, DecodeErrorKind::InvalidUtf8)) + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::HostOwned; + /// + /// let value = HostOwned::try_from("example.com:443")?; + /// assert_eq!(value.as_field_value().as_bytes(), b"example.com:443"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::HostOwned; + /// + /// let value = HostOwned::try_from("example.com")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"example.com"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(HostOwned, |value| value.value); + +impl<'a> HostView<'a> { + /// Returns the validated host kind and retained address or name components. + /// + /// A relaxed international name borrows its normalization from this view. + #[must_use] + #[inline] + pub fn kind(&self) -> HostKind<'_> { + self.kind.project(self.host, self.normalized.as_deref()) + } + + /// Returns the validated textual port, including an explicitly empty port. + #[must_use] + #[inline] + pub fn port_view(&self) -> Option> { + self.port + .zip(self.numeric_port) + .map(|(text, numeric)| HostPortView { text, numeric }) + } + + /// Returns the checked network port, or `None` only when no port is present. + /// + /// # Errors + /// + /// Returns an error for an explicitly empty port or a value above `u16::MAX`. + #[inline] + pub fn network_port(&self) -> Result, PortConversionError> { + self.numeric_port.transpose() + } + + /// Returns the URI host, including brackets around an IP literal. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::HostOwned::try_from("example.com:443")?; + /// assert_eq!(value.host()?, "example.com"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn host(&self) -> &'a str { + self.host + } + + /// Returns the optional decimal port. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::Host; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let with_port = ::decode_view(FieldValueRef::new(b"[::1]:8080"))?; + /// assert_eq!(with_port.port(), Some("8080")); + /// + /// let without_port = ::decode_view(FieldValueRef::new(b"example.net"))?; + /// assert_eq!(without_port.port(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn port(&self) -> Option<&'a str> { + self.port + } + + /// Returns the complete authority. + /// + /// # Errors + /// + /// Returns an error if the wire value is unexpectedly non-UTF-8. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::Host; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = ::decode_view(FieldValueRef::new(b"example.com:8080"))?; + /// assert_eq!(view.as_str()?, "example.com:8080"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_str(&self) -> Result<&'a str, DecodeError> { + self.value + .to_str() + .map_err(|_invalid| invalid(&FieldName::Host, DecodeErrorKind::InvalidUtf8)) + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::Host; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = ::decode_view(FieldValueRef::new(b"example.org"))?; + /// assert_eq!(view.as_field_value().as_bytes(), b"example.org"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(&self) -> FieldValueRef<'a> { + self.value + } +} + +impl SingleValueField for Host { + type View<'a> = HostView<'a>; + type Owned = HostOwned; + + fn name() -> &'static FieldName { + &FieldName::Host + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + let bytes = value.as_bytes(); + let parsed = parse_host(bytes)?; + let authority = value.to_str().expect("strict host validation accepts only ASCII"); + let (host, port) = match parsed.port_start { + Some(start) => (&authority[..parsed.host_end], Some(&authority[start..])), + // Without a port the host runs to the end of the authority. + None => (authority, None), + }; + Ok(HostView { + value, + host, + port, + kind: parsed.kind, + normalized: parsed.normalized.map(Cow::Owned), + numeric_port: parsed.numeric_port, + }) + } + + #[expect( + clippy::inline_always, + reason = "measured: Criterion otherwise outlines this cross-crate hot path while Callgrind inlines it" + )] + #[inline(always)] + fn decode_owned(value: FieldValue) -> Result { + HostOwned::try_from(value) + } + + fn decode_view_with(value: FieldValueRef<'_>, mode: crate::DecodeMode) -> Result, DecodeError> { + let authority = value + .to_str() + .map_err(|_invalid| invalid(&FieldName::Host, DecodeErrorKind::InvalidUtf8))?; + let parsed = parse_host_with(value.as_bytes(), mode)?; + let (host, port) = match parsed.port_start { + Some(start) => (&authority[..parsed.host_end], Some(&authority[start..])), + None => (authority, None), + }; + Ok(HostView { + value, + host, + port, + kind: parsed.kind, + normalized: parsed.normalized.map(Cow::Owned), + numeric_port: parsed.numeric_port, + }) + } + + fn decode_owned_with(value: FieldValue, mode: crate::DecodeMode) -> Result { + let parsed = parse_host_with(value.as_bytes(), mode)?; + Ok(HostOwned { value, parsed }) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +super::super::shared::impl_string_conversions!(HostOwned, &FieldName::Host, invalid_syntax, value); + +impl TryFrom for HostOwned { + type Error = DecodeError; + + #[expect( + clippy::inline_always, + reason = "measured: outlining this conversion leaves large Result traffic in Criterion" + )] + #[inline(always)] + fn try_from(value: FieldValue) -> Result { + let parsed = parse_host(value.as_bytes())?; + Ok(Self { value, parsed }) + } +} + +#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +struct ParsedHost { + host_end: usize, + port_start: Option, + kind: ParsedHostKind, + normalized: Option, + numeric_port: Option>, +} + +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +enum ParsedHostKind { + RegisteredName, + Ipv4(Ipv4Addr), + Ipv6(Ipv6Addr), + IpvFuture { dot: usize }, +} + +impl ParsedHostKind { + fn project<'a>(self, host: &'a str, normalized: Option<&'a str>) -> HostKind<'a> { + match self { + Self::RegisteredName => HostKind::RegisteredName(RegisteredNameView { + original: host, + normalized: normalized.unwrap_or(host), + }), + Self::Ipv4(address) => HostKind::Ipv4(address), + Self::Ipv6(address) => HostKind::Ipv6(address), + Self::IpvFuture { dot } => HostKind::IpvFuture(IpvFutureView { + version: &host[2..dot], + address: &host[dot + 1..host.len() - 1], + }), + } + } +} + +#[expect( + clippy::inline_always, + reason = "measured: Callgrind inlines this parser while Criterion otherwise emits a real call" +)] +#[inline(always)] +fn parse_host(bytes: &[u8]) -> Result { + let Some(first) = bytes.first().copied() else { + return Ok(ParsedHost { + host_end: 0, + port_start: None, + kind: ParsedHostKind::RegisteredName, + normalized: None, + numeric_port: None, + }); + }; + if first == b'[' { + return parse_ip_literal_host(bytes); + } + let mut index = 0; + let mut ipv4_candidate = first.is_ascii_digit(); + while index < bytes.len() { + // The table already covers the alphanumerics, `-`, and `.` that make + // up nearly every host, so one load settles a byte where a range test + // ahead of the load would cost several compares. + match HOST_CLASS[usize::from(bytes[index])] { + HOST_REG_NAME => { + ipv4_candidate = ipv4_candidate && (bytes[index].is_ascii_digit() || bytes[index] == b'.'); + index += 1; + } + HOST_PERCENT => { + if !bytes.get(index + 1).is_some_and(u8::is_ascii_hexdigit) || !bytes.get(index + 2).is_some_and(u8::is_ascii_hexdigit) { + return Err(invalid_syntax(&FieldName::Host)); + } + ipv4_candidate = false; + index += 3; + } + HOST_COLON => { + if index == 0 { + return Err(invalid_syntax(&FieldName::Host)); + } + let numeric_port = validate_port(&bytes[index + 1..])?; + return Ok(ParsedHost { + host_end: index, + port_start: Some(index + 1), + kind: registered_host_kind(&bytes[..index], ipv4_candidate), + normalized: None, + numeric_port: Some(numeric_port), + }); + } + _ => return Err(invalid_syntax(&FieldName::Host)), + } + } + Ok(ParsedHost { + host_end: bytes.len(), + port_start: None, + kind: registered_host_kind(bytes, ipv4_candidate), + normalized: None, + numeric_port: None, + }) +} + +fn registered_host_kind(bytes: &[u8], ipv4_candidate: bool) -> ParsedHostKind { + if ipv4_candidate { + let text = str::from_utf8(bytes).expect("the registered-name scan accepts only ASCII"); + if let Ok(address) = text.parse() { + return ParsedHostKind::Ipv4(address); + } + } + ParsedHostKind::RegisteredName +} + +fn parse_host_with(bytes: &[u8], mode: crate::DecodeMode) -> Result { + if mode == crate::DecodeMode::Strict { + return parse_host(bytes); + } + if let Ok(parsed) = parse_host(bytes) { + return Ok(parsed); + } + parse_international_host(bytes) +} + +fn parse_international_host(bytes: &[u8]) -> Result { + let authority = str::from_utf8(bytes).map_err(|_invalid| invalid(&FieldName::Host, DecodeErrorKind::InvalidUtf8))?; + if authority.is_ascii() || authority.contains('@') || authority.starts_with('[') { + return Err(invalid_syntax(&FieldName::Host)); + } + let (host, port, port_start) = match authority.rsplit_once(':') { + Some((host, port)) if !host.contains(':') && port.bytes().all(|byte| byte.is_ascii_digit()) => { + (host, Some(port), Some(host.len() + 1)) + } + Some(_) if authority.contains(':') => return Err(invalid_syntax(&FieldName::Host)), + _ => (authority, None, None), + }; + let mut normalized = domain_to_ascii(host).map_err(|_invalid| invalid_syntax(&FieldName::Host))?; + if normalized.is_empty() { + return Err(invalid_syntax(&FieldName::Host)); + } + let normalized_end = normalized.len(); + if let Some(port) = port { + normalized.push(':'); + normalized.push_str(port); + } + let parsed = parse_host(normalized.as_bytes())?; + normalized.truncate(normalized_end); + Ok(ParsedHost { + host_end: host.len(), + port_start, + kind: ParsedHostKind::RegisteredName, + normalized: Some(normalized), + numeric_port: port.and(parsed.numeric_port), + }) +} + +/// Checks the authority's port, mirroring the whole-value checks it replaces. +/// +/// A byte that a field value may never contain, a byte outside ASCII, and a +/// second colon are syntax errors wherever they appear, so they outrank the +/// numeric error reported for a merely non-decimal port. +fn validate_port(bytes: &[u8]) -> Result, DecodeError> { + let mut decimal = true; + let mut numeric = Some(0_u16); + for byte in bytes { + if byte.is_ascii_digit() { + // A u16 accumulator and a decimal digit produce at most 655_359. + numeric = numeric.and_then(|value| u16::try_from(u32::from(value) * 10 + u32::from(*byte - b'0')).ok()); + continue; + } + if *byte == b':' || *byte == b'@' || *byte == 0x7f || (*byte < b' ' && *byte != b'\t') { + return Err(invalid_syntax(&FieldName::Host)); + } + if !byte.is_ascii() { + return Err(invalid_syntax(&FieldName::Host)); + } + decimal = false; + } + if decimal { + Ok(if bytes.is_empty() { + Err(PortConversionError { + kind: PortConversionErrorKind::Empty, + }) + } else { + numeric.ok_or(PortConversionError { + kind: PortConversionErrorKind::Overflow, + }) + }) + } else { + Err(invalid(&FieldName::Host, DecodeErrorKind::InvalidNumber)) + } +} + +/// Byte classes recognized while scanning an authority's registered name. +/// Consecutive small integers keep the 256-byte lookup table compact. +const HOST_OTHER: u8 = 0; +const HOST_REG_NAME: u8 = 1; +const HOST_PERCENT: u8 = 2; +const HOST_COLON: u8 = 3; + +/// Maps every byte to its role in an authority's registered name. +static HOST_CLASS: [u8; 256] = { + let mut table = [HOST_OTHER; 256]; + let mut index = 0_u8; + loop { + table[index as usize] = if is_unreserved(index) || is_sub_delim(index) { + HOST_REG_NAME + } else if index == b'%' { + HOST_PERCENT + } else if index == b':' { + HOST_COLON + } else { + HOST_OTHER + }; + if index == u8::MAX { + break; + } + index += 1; + } + table +}; + +fn parse_ip_literal_host(bytes: &[u8]) -> Result { + if !validate::field_value(bytes) || !bytes.is_ascii() || bytes.contains(&b'@') { + return Err(invalid_syntax(&FieldName::Host)); + } + let close = bytes + .iter() + .position(|byte| *byte == b']') + .ok_or_else(|| invalid_syntax(&FieldName::Host))?; + let literal = &bytes[1..close]; + let kind = parse_ip_literal(literal).ok_or_else(|| invalid_syntax(&FieldName::Host))?; + let suffix = &bytes[close + 1..]; + let mut numeric_port = None; + let port_start = if suffix.is_empty() { + None + } else if let Some(port) = suffix.strip_prefix(b":") { + let mut numeric = Some(0_u16); + for byte in port { + if !byte.is_ascii_digit() { + return Err(invalid(&FieldName::Host, DecodeErrorKind::InvalidNumber)); + } + // A u16 accumulator and a decimal digit produce at most 655_359. + numeric = numeric.and_then(|value| u16::try_from(u32::from(value) * 10 + u32::from(*byte - b'0')).ok()); + } + numeric_port = Some(if port.is_empty() { + Err(PortConversionError { + kind: PortConversionErrorKind::Empty, + }) + } else { + numeric.ok_or(PortConversionError { + kind: PortConversionErrorKind::Overflow, + }) + }); + Some(close + 2) + } else { + return Err(invalid_syntax(&FieldName::Host)); + }; + Ok(ParsedHost { + host_end: close + 1, + port_start, + kind, + normalized: None, + numeric_port, + }) +} + +fn parse_ip_literal(bytes: &[u8]) -> Option { + let value = str::from_utf8(bytes).ok()?; + let Some(versioned) = bytes.strip_prefix(b"v").or_else(|| bytes.strip_prefix(b"V")) else { + return Ipv6Addr::from_str(value).ok().map(ParsedHostKind::Ipv6); + }; + let dot = versioned.iter().position(|byte| *byte == b'.')?; + let version = &versioned[..dot]; + let address = &versioned[dot + 1..]; + let valid = !version.is_empty() + && version.iter().all(u8::is_ascii_hexdigit) + && !address.is_empty() + && address + .iter() + .copied() + .all(|byte| is_unreserved(byte) || is_sub_delim(byte) || byte == b':'); + valid.then_some(ParsedHostKind::IpvFuture { dot: dot + 2 }) +} + +#[cfg(test)] +fn valid_ip_literal(bytes: &[u8]) -> bool { + parse_ip_literal(bytes).is_some() +} +const fn is_unreserved(byte: u8) -> bool { + byte.is_ascii_alphanumeric() || matches!(byte, b'-' | b'.' | b'_' | b'~') +} + +const fn is_sub_delim(byte: u8) -> bool { + matches!(byte, b'!' | b'$' | b'&' | b'\'' | b'(' | b')' | b'*' | b'+' | b',' | b';' | b'=') +} + +fn component(bytes: &[u8], range: Range) -> Result<&str, DecodeError> { + let bytes = bytes.get(range).ok_or_else(|| invalid_syntax(&FieldName::Host))?; + str::from_utf8(bytes).map_err(|_invalid| invalid(&FieldName::Host, DecodeErrorKind::InvalidUtf8)) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{ + Host, HostOwned, HostView, component, parse_host, parse_host_with, parse_international_host, valid_ip_literal, validate_port, + }; + use crate::{DecodeErrorKind, DecodeMode, FieldName, FieldValue, FieldValueRef, SingleValueField}; + + #[test] + fn constructors_and_accessors_preserve_authority_components() { + let plain = HostOwned::new("example.com").expect("host without port"); + assert_eq!(plain.host(), Ok("example.com")); + assert_eq!(plain.port(), Ok(None)); + assert_eq!(plain.as_str(), Ok("example.com")); + assert_eq!(plain.as_field_value().as_bytes(), b"example.com"); + assert_eq!(plain.into_field_value().as_bytes(), b"example.com"); + + let with_port = HostOwned::with_port("[2001:db8::1]", 443).expect("IPv6 and port"); + assert_eq!(with_port.host(), Ok("[2001:db8::1]")); + assert_eq!(with_port.port(), Ok(Some("443"))); + assert_eq!( + HostOwned::new("example.com:80").expect_err("new excludes ports").kind(), + DecodeErrorKind::InvalidSyntax + ); + + let view = ::decode_view(FieldValueRef::new(b"example.com:8080")).expect("borrowed authority"); + assert_eq!(view.host(), "example.com"); + assert_eq!(view.port(), Some("8080")); + assert_eq!(view.as_str(), Ok("example.com:8080")); + assert_eq!(view.as_field_value().as_bytes(), b"example.com:8080"); + + let no_port = ::decode_view(FieldValueRef::new(b"example.net")).expect("borrowed authority without port"); + assert_eq!(no_port.host(), "example.net"); + assert_eq!(no_port.port(), None); + + let decoded = ::decode_owned(FieldValue::from_static("example.org:8443")).expect("owned authority"); + assert_eq!(::as_field_value(&decoded).as_bytes(), b"example.org:8443"); + assert_eq!( + ::into_field_value(decoded).as_bytes(), + b"example.org:8443" + ); + assert_eq!(::name(), &FieldName::Host); + } + + #[test] + fn strict_parser_covers_registered_names_ports_and_literal_forms() { + for valid in [ + b"".as_slice(), + b"example.com".as_slice(), + b"exa_mple~host", + b"example%20host", + b"example.com:", + b"[2001:db8::1]", + b"[v1.alpha:beta]:443", + ] { + assert!(parse_host(valid).is_ok(), "{valid:?}"); + } + for invalid in [ + b":80".as_slice(), + b"host%2", + b"host%zz", + b"host@name", + b"host:bad", + b"host:8:0", + b"[", + b"[]", + b"[2001:db8::1", + b"[2001:db8::1]tail", + b"[2001:db8::1]:bad", + ] { + assert!(parse_host(invalid).is_err(), "{invalid:?}"); + } + } + + #[test] + fn relaxed_parser_accepts_idna_but_not_ambiguous_authorities() { + let parsed = parse_host_with("münich.example:443".as_bytes(), DecodeMode::Relaxed).expect("international host"); + assert_eq!(parsed.host_end, "münich.example".len()); + assert_eq!(parsed.port_start, Some("münich.example:".len())); + + let strict = parse_host_with(b"example.com", DecodeMode::Strict).expect("strict host"); + assert_eq!(strict.host_end, 11); + assert_eq!(strict.port_start, None); + let relaxed_ascii = parse_host_with(b"example.net", DecodeMode::Relaxed).expect("strict syntax fast path"); + assert_eq!(relaxed_ascii.host_end, 11); + + for invalid in [ + b"bad host".as_slice(), + "münich@example".as_bytes(), + "[münich]".as_bytes(), + "münich:bad".as_bytes(), + b"\xff", + ] { + assert!(parse_host_with(invalid, DecodeMode::Relaxed).is_err(), "{invalid:?}"); + } + assert_eq!( + ::decode_view_with(FieldValueRef::new(b"\xff"), DecodeMode::Relaxed,) + .expect_err("relaxed views still require UTF-8") + .kind(), + DecodeErrorKind::InvalidUtf8 + ); + assert!( + parse_international_host("\u{200d}.example".as_bytes()).is_err(), + "context-invalid IDNA label" + ); + assert!( + parse_international_host("\u{00ad}".as_bytes()).is_err(), + "IDNA mappings may not erase the complete host" + ); + + let value = FieldValue::from_bytes("münich.example:443".as_bytes()).expect("UTF-8 field value"); + let view = + ::decode_view_with(value.as_field_value_ref(), DecodeMode::Relaxed).expect("relaxed borrowed host"); + assert_eq!(view.host(), "münich.example"); + assert_eq!(view.port(), Some("443")); + drop(view); + let owned = ::decode_owned_with(value, DecodeMode::Relaxed).expect("relaxed owned host"); + assert_eq!(owned.as_str(), Ok("münich.example:443")); + assert_eq!(owned.host(), Ok("münich.example")); + assert_eq!(owned.port(), Ok(Some("443"))); + + let value = FieldValue::from_bytes("münich.example".as_bytes()).expect("UTF-8 field value"); + let view = ::decode_view_with(value.as_field_value_ref(), DecodeMode::Relaxed) + .expect("relaxed host without port"); + assert_eq!(view.host(), "münich.example"); + assert_eq!(view.port(), None); + drop(view); + let owned = ::decode_owned_with(value, DecodeMode::Relaxed).expect("relaxed owned host without port"); + assert_eq!(owned.host(), Ok("münich.example")); + assert_eq!(owned.port(), Ok(None)); + } + + #[test] + fn private_component_validators_report_precise_failures() { + for valid in [b"".as_slice(), b"0", b"65536"] { + assert!(validate_port(valid).is_ok(), "{valid:?}"); + } + for syntax in [b"8:0".as_slice(), b"user@", b"\x7f", b"\x01", b"\xff"] { + assert_eq!( + validate_port(syntax).expect_err("syntax error").kind(), + DecodeErrorKind::InvalidSyntax + ); + } + assert_eq!( + validate_port(b"http").expect_err("not decimal").kind(), + DecodeErrorKind::InvalidNumber + ); + + for valid in [b"2001:db8::1".as_slice(), b"vF.alpha:beta"] { + assert!(valid_ip_literal(valid), "{valid:?}"); + } + for invalid in [b"\xff".as_slice(), b"not-ip", b"v.", b"v1", b"v1.", b"v1.bad?"] { + assert!(!valid_ip_literal(invalid), "{invalid:?}"); + } + + assert_eq!(component(b"abc", 0..3), Ok("abc")); + assert_eq!( + component(b"abc", 0..4).expect_err("out of bounds").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + component(b"\xff", 0..1).expect_err("invalid UTF-8").kind(), + DecodeErrorKind::InvalidUtf8 + ); + assert_eq!( + HostOwned::try_from(String::from("line\nbreak")) + .expect_err("invalid field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + HostOwned::try_from("line\nbreak").expect_err("invalid borrowed field value").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + HostOwned::try_from(FieldValue::from_static("bad host")) + .expect_err("invalid authority") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn defensive_owned_and_view_accessors_reject_invalid_private_storage() { + let malformed = HostOwned { + value: FieldValue::from_static("short"), + parsed: super::ParsedHost { + host_end: 6, + port_start: Some(6), + kind: super::ParsedHostKind::RegisteredName, + normalized: None, + numeric_port: None, + }, + }; + assert_eq!( + malformed.host().expect_err("invalid host metadata").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + malformed.port().expect_err("invalid port metadata").kind(), + DecodeErrorKind::InvalidSyntax + ); + + let non_utf8 = FieldValue::from_bytes(b"\xff").expect("obs-text is a valid field value"); + let malformed = HostOwned { + value: non_utf8.clone(), + parsed: super::ParsedHost { + host_end: 1, + port_start: None, + kind: super::ParsedHostKind::RegisteredName, + normalized: None, + numeric_port: None, + }, + }; + assert_eq!(malformed.as_str().expect_err("invalid UTF-8").kind(), DecodeErrorKind::InvalidUtf8); + let view = HostView { + value: non_utf8.as_field_value_ref(), + host: "", + port: None, + kind: super::ParsedHostKind::RegisteredName, + normalized: None, + numeric_port: None, + }; + assert_eq!(view.as_str().expect_err("invalid UTF-8 view").kind(), DecodeErrorKind::InvalidUtf8); + } +} diff --git a/crates/http_headers/src/headers/negotiation/host/components.rs b/crates/http_headers/src/headers/negotiation/host/components.rs new file mode 100644 index 000000000..5307150e4 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/host/components.rs @@ -0,0 +1,165 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::error::Error; +use std::fmt; +use std::net::{Ipv4Addr, Ipv6Addr}; + +/// The validated kind of a URI host. +/// +/// Registered names are not necessarily DNS names. Numeric-looking names that +/// are not strict dotted-decimal IPv4 addresses remain registered names. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub enum HostKind<'a> { + /// A URI registered name, without percent decoding or DNS resolution. + RegisteredName(RegisteredNameView<'a>), + /// A strict dotted-decimal IPv4 address. + Ipv4(Ipv4Addr), + /// A parsed IPv6 address; brackets remain in the header's textual accessor. + Ipv6(Ipv6Addr), + /// A validated `IPvFuture` literal. + IpvFuture(IpvFutureView<'a>), +} + +/// A validated URI registered name and its retained ASCII normalization. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct RegisteredNameView<'a> { + pub(super) original: &'a str, + pub(super) normalized: &'a str, +} + +impl<'a> RegisteredNameView<'a> { + /// Returns the original name, including any percent escapes. + #[must_use] + #[inline] + pub const fn as_str(self) -> &'a str { + self.original + } + + /// Returns the ASCII name produced by relaxed IDNA validation. + /// + /// For names accepted by strict parsing this is the unchanged original + /// name, not a promise of DNS validity, lowercasing, or percent decoding. + /// Relaxed IDNA mappings are retained exactly as validated; any mapped + /// delimiters are not reinterpreted as ports in the original wire value. + #[must_use] + #[inline] + pub const fn normalized(self) -> &'a str { + self.normalized + } +} + +impl fmt::Display for RegisteredNameView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.original) + } +} + +/// Validated version and address components of an `IPvFuture` literal. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct IpvFutureView<'a> { + pub(super) version: &'a str, + pub(super) address: &'a str, +} + +impl<'a> IpvFutureView<'a> { + /// Returns the hexadecimal version, without imposing an integer size limit. + #[must_use] + #[inline] + pub const fn version(self) -> &'a str { + self.version + } + + /// Returns the address, excluding the version and square brackets. + #[must_use] + #[inline] + pub const fn address(self) -> &'a str { + self.address + } +} + +impl fmt::Display for IpvFutureView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "v{}.{}", self.version, self.address) + } +} + +/// A validated textual URI port, which may be empty or exceed `u16`. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct HostPortView<'a> { + pub(super) text: &'a str, + pub(super) numeric: Result, +} + +impl<'a> HostPortView<'a> { + /// Validates a textual port independently of its network-port range. + /// + /// # Errors + /// + /// Returns a `Host` decode error for non-decimal text. + pub fn new(text: &'a str) -> Result { + Ok(Self { + text, + numeric: super::validate_port(text.as_bytes())?, + }) + } + + /// Returns the original decimal spelling, including leading zeros. + #[must_use] + #[inline] + pub const fn as_str(self) -> &'a str { + self.text + } + + /// Returns the retained checked network-port conversion. + /// + /// # Errors + /// + /// Returns [`PortConversionErrorKind::Empty`] for an empty port, or + /// [`PortConversionErrorKind::Overflow`] when its value exceeds `u16`. + #[inline] + pub const fn to_u16(self) -> Result { + self.numeric + } +} + +impl fmt::Display for HostPortView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.text) + } +} + +/// Why a validated textual URI port is not a network port. +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub enum PortConversionErrorKind { + /// The colon is present but its port text is empty. + Empty, + /// The decimal value exceeds `u16::MAX`. + Overflow, +} + +/// A checked conversion failure for a syntactically valid URI port. +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct PortConversionError { + pub(super) kind: PortConversionErrorKind, +} + +impl PortConversionError { + /// Returns the reason conversion failed. + #[must_use] + #[inline] + pub const fn kind(self) -> PortConversionErrorKind { + self.kind + } +} + +impl fmt::Display for PortConversionError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(match self.kind { + PortConversionErrorKind::Empty => "the URI port is empty", + PortConversionErrorKind::Overflow => "the URI port exceeds u16::MAX", + }) + } +} + +impl Error for PortConversionError {} diff --git a/crates/http_headers/src/headers/negotiation/language_range.rs b/crates/http_headers/src/headers/negotiation/language_range.rs new file mode 100644 index 000000000..dcbe37fba --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/language_range.rs @@ -0,0 +1,92 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::accept_language::valid_language_range; +use super::negotiation_token::NegotiationToken; +use super::shared::invalid; +use crate::{DecodeError, DecodeErrorKind, FieldName}; + +/// A validated basic language range, or the standalone wildcard `*`. +/// +/// This models the basic HTTP range grammar, not a `BCP 47` registry or extended +/// language ranges. Equality is ASCII case-insensitive. No case normalization, +/// locale fallback or matching policy is applied. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct LanguageRange<'a> { + token: NegotiationToken<'a>, + primary_len: usize, +} + +impl<'a> LanguageRange<'a> { + /// Validates a wildcard or a primary subtag followed by optional subtags. + /// + /// # Errors + /// + /// The primary must contain one to eight ASCII letters; subsequent + /// subtags must contain one to eight ASCII letters or digits. Internal + /// wildcards and empty subtags are rejected. + pub fn parse(value: &'a str) -> Result { + if !valid_language_range(value.as_bytes()) { + return Err(invalid(&FieldName::AcceptLanguage, DecodeErrorKind::InvalidToken)); + } + Ok(Self::from_validated(value.as_bytes())) + } + + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + let primary_len = if bytes == b"*" { + 0 + } else { + bytes.iter().position(|byte| *byte == b'-').unwrap_or(bytes.len()) + }; + Self { + token: NegotiationToken::from_validated(bytes), + primary_len, + } + } + + /// Returns the complete case-preserved range. + #[must_use] + pub const fn as_str(self) -> &'a str { + self.token.as_str() + } + + /// Whether the range is exactly `*`. + #[must_use] + pub const fn is_wildcard(self) -> bool { + self.primary_len == 0 + } + + /// Returns the primary subtag, or `None` for a wildcard. + #[must_use] + pub fn primary(self) -> Option> { + if self.is_wildcard() { + None + } else { + Some(NegotiationToken::from_validated_str(&self.as_str()[..self.primary_len])) + } + } + + /// Iterates subsequent subtags without allocating, excluding the primary. + /// + /// A wildcard and a primary-only range both yield an empty iterator. + pub fn subtags(self) -> impl Iterator> { + let rest = (!self.is_wildcard() && self.primary_len < self.as_str().len()).then(|| &self.as_str()[self.primary_len + 1..]); + rest.into_iter() + .flat_map(|rest| rest.split('-')) + .map(NegotiationToken::from_validated_str) + } + + /// Compares a complete range without ASCII case distinctions. + #[must_use] + pub fn eq_ignore_ascii_case(self, other: &str) -> bool { + self.token.eq_ignore_ascii_case(other) + } +} + +impl fmt::Display for LanguageRange<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.token.fmt(f) + } +} diff --git a/crates/http_headers/src/headers/negotiation/media_range.rs b/crates/http_headers/src/headers/negotiation/media_range.rs new file mode 100644 index 000000000..65bddf132 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/media_range.rs @@ -0,0 +1,100 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::{fmt, str}; + +use super::accept::validate_media_range; +use super::negotiation_token::NegotiationToken; +use super::shared::invalid_syntax; +use crate::{DecodeError, FieldName}; + +/// The wildcard semantics of a validated media range. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub enum MediaRangeKind { + /// `*/*`, accepting any media type. + Any, + /// A named type and `*` subtype, such as `text/*`. + TypeWildcard, + /// A type and subtype without a standalone wildcard component. + Exact, +} + +/// A validated media range with case-insensitive component equality. +/// +/// Only a complete `*` component is a wildcard. Embedded stars remain ordinary +/// token characters, and the original spelling is retained. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct MediaRange<'a> { + type_: NegotiationToken<'a>, + subtype: NegotiationToken<'a>, + kind: MediaRangeKind, +} + +impl<'a> MediaRange<'a> { + /// Validates a media range without parameters or surrounding whitespace. + /// + /// # Errors + /// + /// Returns an error for invalid components or a wildcard type with a + /// non-wildcard subtype. + pub fn parse(value: &'a str) -> Result { + validate_media_range(value.as_bytes())?; + Ok(Self::from_validated(value.as_bytes())) + } + + /// Constructs a media range from validated components. + /// + /// # Errors + /// + /// A wildcard type requires a wildcard subtype. + pub fn new(type_: NegotiationToken<'a>, subtype: NegotiationToken<'a>) -> Result { + let kind = if type_.as_str() == "*" { + if subtype.as_str() != "*" { + return Err(invalid_syntax(&FieldName::Accept)); + } + MediaRangeKind::Any + } else if subtype.as_str() == "*" { + MediaRangeKind::TypeWildcard + } else { + MediaRangeKind::Exact + }; + Ok(Self { type_, subtype, kind }) + } + + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + let slash = bytes + .iter() + .position(|byte| *byte == b'/') + .expect("validated media ranges contain a slash"); + let text = str::from_utf8(bytes).expect("validated media ranges contain only ASCII bytes"); + Self::new( + NegotiationToken::from_validated_str(&text[..slash]), + NegotiationToken::from_validated_str(&text[slash + 1..]), + ) + .expect("validated media ranges satisfy the cross-component wildcard rule") + } + + /// Returns the range's retained wildcard classification. + #[must_use] + pub const fn kind(self) -> MediaRangeKind { + self.kind + } + + /// Returns the type token, including `*` for an any-type range. + #[must_use] + pub const fn type_(self) -> NegotiationToken<'a> { + self.type_ + } + + /// Returns the subtype token, including `*` for wildcard ranges. + #[must_use] + pub const fn subtype(self) -> NegotiationToken<'a> { + self.subtype + } +} + +impl fmt::Display for MediaRange<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + write!(f, "{}/{}", self.type_, self.subtype) + } +} diff --git a/crates/http_headers/src/headers/negotiation/mod.rs b/crates/http_headers/src/headers/negotiation/mod.rs new file mode 100644 index 000000000..2318e8cb4 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/mod.rs @@ -0,0 +1,66 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +mod accept; +mod accept_encoding; +mod accept_encoding_entry; +mod accept_entry; +mod accept_language; +mod accept_language_entry; +mod accept_scan; +mod allow; +mod content_coding; +mod host; +mod language_range; +mod media_range; +mod negotiation_members; +mod negotiation_parameter; +mod negotiation_token; +mod quality; +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod recognition_test_support; +mod server; +mod shared; +mod vary; +mod vary_entry_view; +mod weighted_token_scan; + +#[doc(inline)] +pub use accept::{Accept, AcceptOwned, AcceptView}; +#[doc(inline)] +pub use accept_encoding::{AcceptEncoding, AcceptEncodingOwned, AcceptEncodingView}; +#[doc(inline)] +pub use accept_encoding_entry::AcceptEncodingEntry; +#[doc(inline)] +pub use accept_entry::AcceptEntry; +#[doc(inline)] +pub use accept_language::{AcceptLanguage, AcceptLanguageOwned, AcceptLanguageView}; +#[doc(inline)] +pub use accept_language_entry::AcceptLanguageEntry; +#[doc(inline)] +pub use allow::{Allow, AllowOwned, AllowView}; +#[doc(inline)] +pub use content_coding::{ContentCoding, ContentCodingKind}; +#[doc(inline)] +pub use host::{ + Host, HostKind, HostOwned, HostPortView, HostView, IpvFutureView, PortConversionError, PortConversionErrorKind, RegisteredNameView, +}; +#[doc(inline)] +pub use language_range::LanguageRange; +#[doc(inline)] +pub use media_range::{MediaRange, MediaRangeKind}; +#[doc(inline)] +pub use negotiation_parameter::{NegotiationParameter, NegotiationParameterValue, NegotiationParameters}; +#[doc(inline)] +pub use negotiation_token::NegotiationToken; +#[doc(inline)] +pub use quality::{InexactQuality, InvalidQuality, Quality, QualityView}; +#[doc(inline)] +pub use server::{Server, ServerOwned, ServerView}; +#[doc(inline)] +pub use vary::{Vary, VaryOwned, VaryView}; +#[doc(inline)] +pub use vary_entry_view::VaryEntryView; + +use super::{FieldNameView, MethodView}; diff --git a/crates/http_headers/src/headers/negotiation/negotiation_members.rs b/crates/http_headers/src/headers/negotiation/negotiation_members.rs new file mode 100644 index 000000000..eca63be29 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/negotiation_members.rs @@ -0,0 +1,89 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::quality::QualityView; +use super::shared::{ListValues, invalid_syntax}; +use crate::source::{MAX_CUSTOM_FIELD_BYTES, MAX_CUSTOM_LIST_ITEMS}; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, validate}; + +pub(super) fn validated_member(item: Result<&[u8], DecodeError>) -> &[u8] { + item.expect("header decoding validated every list member before constructing the immutable header") +} + +pub(super) fn weighted_parts(bytes: &[u8]) -> (&[u8], Option>) { + let Some(semicolon) = bytes.iter().position(|byte| *byte == b';') else { + return (bytes, None); + }; + let parameter = &bytes[semicolon + 1..]; + let equals = parameter + .iter() + .position(|byte| *byte == b'=') + .expect("validated weights contain an equals sign"); + ( + validate::trim_ows(&bytes[..semicolon]), + Some(QualityView::from_validated(validate::trim_ows(¶meter[equals + 1..]))), + ) +} + +pub(super) fn checked_len(lengths: impl IntoIterator, name: &'static FieldName) -> Result { + let mut total = 0_usize; + for length in lengths { + total = total + .checked_add(length) + .filter(|total| *total <= MAX_CUSTOM_FIELD_BYTES) + .ok_or_else(|| DecodeError::new(name, DecodeErrorKind::SourceLimitExceeded))?; + } + Ok(total) +} + +pub(super) fn weight_len(quality: Option>) -> usize { + quality.map_or(0, |value| 3 + value.encoded_len()) +} + +pub(super) fn append_weight(bytes: &mut Vec, quality: Option>) { + if let Some(quality) = quality { + bytes.extend_from_slice(b";q="); + quality.append_to(bytes); + } +} + +pub(super) fn collect_members( + entries: impl IntoIterator, + name: &'static FieldName, + length: impl Fn(&T) -> Result, + append: impl Fn(&T, &mut Vec), +) -> Result { + let mut bytes = Vec::new(); + for (index, entry) in entries.into_iter().enumerate() { + if index >= MAX_CUSTOM_LIST_ITEMS { + return Err(DecodeError::new(name, DecodeErrorKind::SourceLimitExceeded)); + } + let separator = usize::from(index != 0) * 2; + let len = checked_len([bytes.len(), separator, length(&entry)?], name)?; + bytes.reserve(len - bytes.len()); + if separator != 0 { + bytes.extend_from_slice(b", "); + } + append(&entry, &mut bytes); + } + let value = FieldValue::try_from(bytes).map_err(|_invalid| invalid_syntax(name))?; + Ok(ListValues::One(value)) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::checked_len; + use crate::{DecodeError, DecodeErrorKind, FieldName}; + + #[test] + fn member_length_bounds_and_arithmetic_overflow_are_admission_errors() { + assert_eq!(checked_len([65_535, 1], &FieldName::Accept), Ok(65_536)); + for lengths in [[65_536, 1], [1, usize::MAX]] { + assert_eq!( + checked_len(lengths, &FieldName::Accept), + Err(DecodeError::new(&FieldName::Accept, DecodeErrorKind::SourceLimitExceeded)) + ); + } + } +} diff --git a/crates/http_headers/src/headers/negotiation/negotiation_parameter.rs b/crates/http_headers/src/headers/negotiation/negotiation_parameter.rs new file mode 100644 index 000000000..db5e3a6e9 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/negotiation_parameter.rs @@ -0,0 +1,185 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::iter::{self, FusedIterator}; +use std::slice; + +use super::negotiation_token::NegotiationToken; +use super::shared::{QuotedItems, invalid_syntax, valid_parameter_value}; +use crate::{DecodeError, FieldName, validate}; + +/// A validated token or quoted negotiation parameter value. +/// +/// Values are bytes, not necessarily UTF-8. Equality and hashing preserve the +/// raw spelling; use [`decoded_bytes()`](Self::decoded_bytes) to compare +/// unescaped bytes. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct NegotiationParameterValue<'a> { + raw: &'a [u8], + contents: &'a [u8], + quoted: bool, +} + +impl<'a> NegotiationParameterValue<'a> { + /// Validates the raw token or quoted-string spelling. + /// + /// # Errors + /// + /// Returns an error for invalid quoting, escapes, controls or token bytes. + pub fn parse(raw: &'a [u8]) -> Result { + if !valid_parameter_value(raw) { + return Err(invalid_syntax(&FieldName::Accept)); + } + Ok(Self::from_validated(raw)) + } + + fn from_validated(raw: &'a [u8]) -> Self { + let quoted = raw.first() == Some(&b'"'); + let contents = if quoted { &raw[1..raw.len() - 1] } else { raw }; + Self { raw, contents, quoted } + } + + /// Returns the original spelling, including quotes and backslashes. + #[must_use] + pub const fn raw_bytes(self) -> &'a [u8] { + self.raw + } + + /// Whether the wire value was a quoted string rather than a token. + #[must_use] + pub const fn is_quoted(self) -> bool { + self.quoted + } + + /// Iterates decoded bytes without allocating or interpreting them as UTF-8. + /// + /// Quoted-pair backslashes are removed. Each new iterator traverses the + /// value again; the parameter name and value boundaries are retained. + #[must_use] + pub fn decoded_bytes(self) -> impl FusedIterator + 'a { + let mut bytes = self.contents.iter().copied(); + iter::from_fn(move || Self::next_decoded(&mut bytes, self.quoted)).fuse() + } + + fn next_decoded(bytes: &mut impl Iterator, quoted: bool) -> Option { + let byte = bytes.next()?; + Some(if quoted && byte == b'\\' { + bytes.next().expect("validated quoted pairs always contain an escaped byte") + } else { + byte + }) + } + + /// Explicitly allocates the unescaped byte value. + #[must_use] + pub fn to_decoded_bytes(self) -> Vec { + self.decoded_bytes().collect() + } +} + +/// A validated negotiation parameter name and optional byte-valued value. +/// +/// An absent value is permitted for Accept extensions, but not media +/// parameters. [`AcceptEntry::new`](super::accept_entry::AcceptEntry::new) +/// validates that distinction and reserves the `q` name. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct NegotiationParameter<'a> { + name: NegotiationToken<'a>, + value: Option>, +} + +impl<'a> NegotiationParameter<'a> { + /// Combines a validated name and optional validated value. + #[must_use] + pub const fn new(name: NegotiationToken<'a>, value: Option>) -> Self { + Self { name, value } + } + + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + let (name, value) = match bytes.iter().position(|byte| *byte == b'=') { + Some(equals) => ( + validate::trim_ows(&bytes[..equals]), + Some(NegotiationParameterValue::from_validated(validate::trim_ows(&bytes[equals + 1..]))), + ), + None => (bytes, None), + }; + Self { + name: NegotiationToken::from_validated(name), + value, + } + } + + /// Returns the case-insensitive name token with original spelling. + #[must_use] + pub const fn name(self) -> NegotiationToken<'a> { + self.name + } + + /// Returns the value, distinguishing a bare extension from `name=""`. + #[must_use] + pub const fn value(self) -> Option> { + self.value + } + + pub(super) fn encoded_len(self) -> usize { + self.name.as_str().len() + self.value.map_or(0, |value| 1 + value.raw.len()) + } + + pub(super) fn append_to(self, bytes: &mut Vec) { + bytes.push(b';'); + bytes.extend_from_slice(self.name.as_str().as_bytes()); + if let Some(value) = self.value { + bytes.push(b'='); + bytes.extend_from_slice(value.raw); + } + } +} + +#[derive(Clone, Copy, Debug)] +pub(super) enum ParameterSource<'a> { + Wire(&'a [u8]), + Components(&'a [NegotiationParameter<'a>]), +} + +impl<'a> ParameterSource<'a> { + pub(super) fn iter(self) -> NegotiationParameters<'a> { + let repr = match self { + Self::Wire([]) => ParameterIterator::Empty, + Self::Wire(bytes) => ParameterIterator::Wire(QuotedItems::semicolon(bytes, &FieldName::Accept)), + Self::Components(parameters) => ParameterIterator::Components(parameters.iter()), + }; + NegotiationParameters { repr } + } +} + +/// Allocation-free traversal of media parameters or Accept extensions. +/// +/// Order and duplicates are preserved. Each item retains its component +/// boundaries, so its getters do not repeat parameter parsing. +#[derive(Debug)] +pub struct NegotiationParameters<'a> { + repr: ParameterIterator<'a>, +} + +#[derive(Debug)] +enum ParameterIterator<'a> { + Empty, + Wire(QuotedItems<'a>), + Components(slice::Iter<'a, NegotiationParameter<'a>>), +} + +impl<'a> Iterator for NegotiationParameters<'a> { + type Item = NegotiationParameter<'a>; + + fn next(&mut self) -> Option { + match &mut self.repr { + ParameterIterator::Empty => None, + ParameterIterator::Components(parameters) => parameters.next().copied(), + ParameterIterator::Wire(parameters) => parameters + .next() + .map(|parameter| NegotiationParameter::from_validated(parameter.expect("header decoding validated every parameter"))), + } + } +} + +impl FusedIterator for NegotiationParameters<'_> {} diff --git a/crates/http_headers/src/headers/negotiation/negotiation_token.rs b/crates/http_headers/src/headers/negotiation/negotiation_token.rs new file mode 100644 index 000000000..84f86a47a --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/negotiation_token.rs @@ -0,0 +1,88 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::cmp::Ordering; +use std::hash::{Hash, Hasher}; +use std::{fmt, str}; + +use super::shared::invalid; +use crate::{DecodeError, DecodeErrorKind, FieldName, validate}; + +/// A validated HTTP token with case-insensitive semantic equality. +/// +/// Original spelling is preserved. This type is for negotiation components, +/// not case-sensitive HTTP methods or arbitrary parameter values. +#[derive(Clone, Copy, Debug)] +pub struct NegotiationToken<'a>(&'a str); + +impl<'a> NegotiationToken<'a> { + /// Validates an HTTP token. + /// + /// # Errors + /// + /// Returns an error if the string is empty or contains a non-token byte. + pub fn new(value: &'a str) -> Result { + if !validate::token(value.as_bytes()) { + return Err(invalid(&FieldName::Accept, DecodeErrorKind::InvalidToken)); + } + Ok(Self(value)) + } + + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + Self(str::from_utf8(bytes).expect("validated HTTP tokens contain only ASCII bytes")) + } + + pub(super) const fn from_validated_str(value: &'a str) -> Self { + Self(value) + } + + /// Returns the original case-preserved spelling. + #[must_use] + pub const fn as_str(self) -> &'a str { + self.0 + } + + /// Compares a spelling without ASCII case distinctions. + #[must_use] + pub fn eq_ignore_ascii_case(self, other: &str) -> bool { + self.0.eq_ignore_ascii_case(other) + } +} + +impl PartialEq for NegotiationToken<'_> { + fn eq(&self, other: &Self) -> bool { + self.0.eq_ignore_ascii_case(other.0) + } +} + +impl Eq for NegotiationToken<'_> {} + +impl PartialOrd for NegotiationToken<'_> { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl Ord for NegotiationToken<'_> { + fn cmp(&self, other: &Self) -> Ordering { + self.0 + .bytes() + .map(|byte| byte.to_ascii_lowercase()) + .cmp(other.0.bytes().map(|byte| byte.to_ascii_lowercase())) + } +} + +impl Hash for NegotiationToken<'_> { + fn hash(&self, state: &mut H) { + self.0.len().hash(state); + for byte in self.0.bytes() { + byte.to_ascii_lowercase().hash(state); + } + } +} + +impl fmt::Display for NegotiationToken<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.0) + } +} diff --git a/crates/http_headers/src/headers/negotiation/quality.rs b/crates/http_headers/src/headers/negotiation/quality.rs new file mode 100644 index 000000000..9886e897f --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/quality.rs @@ -0,0 +1,296 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::cmp::Ordering; +use std::error::Error; +use std::fmt; +use std::hash::{Hash, Hasher}; + +use super::shared::validate_quality; +use crate::{DecodeMode, FieldName, validate}; + +/// An exact HTTP quality in integer thousandths, between zero and one. +/// +/// Use [`QualityView`] for relaxed fractions that are not exact thousandths. +#[derive(Clone, Copy, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct Quality(u16); + +impl Quality { + /// Zero quality: the member is unacceptable. + pub const ZERO: Self = Self(0); + + /// Full quality, also the effective value when `q` is absent. + pub const ONE: Self = Self(1000); + + /// Constructs an exact strict quality. + /// + /// # Errors + /// + /// Returns an error when `value` exceeds 1000. + pub const fn from_thousandths(value: u16) -> Result { + if value <= 1000 { + Ok(Self(value)) + } else { + Err(InvalidQuality { _private: () }) + } + } + + /// Returns the exact integer number of thousandths. + #[must_use] + pub const fn thousandths(self) -> u16 { + self.0 + } +} + +impl fmt::Display for Quality { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + QualityView::from(*self).fmt(f) + } +} + +/// An invalid quality spelling or out-of-range number of thousandths. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct InvalidQuality { + _private: (), +} + +impl fmt::Display for InvalidQuality { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("quality must be an accepted decimal between zero and one") + } +} + +impl Error for InvalidQuality {} + +/// An exact quality that cannot be represented in integer thousandths. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct InexactQuality { + _private: (), +} + +impl fmt::Display for InexactQuality { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str("quality is not an exact number of thousandths") + } +} + +impl Error for InexactQuality {} + +/// An exact borrowed quality, including arbitrarily long relaxed fractions. +/// +/// Values exactly representable in thousandths use compact storage. Other +/// values borrow validated fractional digits. Equality, ordering and hashing +/// are numeric: `0.5`, `.50` and `0.5000` are equal. No rounding or floating +/// point is used. The original spelling remains in the header's raw members. +#[derive(Clone, Copy, Debug)] +pub struct QualityView<'a> { + repr: Representation<'a>, +} + +#[derive(Clone, Copy, Debug)] +enum Representation<'a> { + Compact(Quality), + Fraction(&'a [u8]), +} + +impl<'a> QualityView<'a> { + /// Zero quality. + pub const ZERO: Self = Self { + repr: Representation::Compact(Quality::ZERO), + }; + + /// Full quality. + pub const ONE: Self = Self { + repr: Representation::Compact(Quality::ONE), + }; + + /// Parses one quality using the header's strict or relaxed grammar. + /// + /// Relaxed mode accepts leading-dot fractions and additional fractional + /// digits, with no additional precision limit. + /// + /// # Errors + /// + /// Returns an error for a malformed or out-of-range quality. + pub fn parse(bytes: &'a [u8], mode: DecodeMode) -> Result { + validate_quality(bytes, &FieldName::Accept, mode == DecodeMode::Relaxed).map_err(|_invalid| InvalidQuality { _private: () })?; + Ok(Self::from_validated(validate::trim_ows(bytes))) + } + + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + if bytes[0] == b'1' { + return Self::ONE; + } + let ([b'0', b'.', fraction @ ..] | [b'.', fraction @ ..]) = bytes else { + return Self::ZERO; + }; + let significant = fraction.iter().rposition(|digit| *digit != b'0').map_or(0, |last| last + 1); + let fraction = &fraction[..significant]; + if significant > 3 { + return Self { + repr: Representation::Fraction(fraction), + }; + } + let mut thousandths = 0; + for index in 0..3 { + thousandths = thousandths * 10 + u16::from(fraction.get(index).copied().unwrap_or(b'0') - b'0'); + } + Self::from(Quality(thousandths)) + } + + /// Converts without rounding. + /// + /// # Errors + /// + /// Returns [`InexactQuality`] if any digit beyond thousandths is nonzero. + pub const fn to_quality(self) -> Result { + match self.repr { + Representation::Compact(value) => Ok(value), + Representation::Fraction(_) => Err(InexactQuality { _private: () }), + } + } + + /// Whether the exact quality is zero. + #[must_use] + pub const fn is_zero(self) -> bool { + matches!(self.repr, Representation::Compact(Quality::ZERO)) + } + + /// Whether the exact quality is one. + #[must_use] + pub const fn is_one(self) -> bool { + matches!(self.repr, Representation::Compact(Quality::ONE)) + } + + fn fraction_len(self) -> usize { + match self.repr { + Representation::Fraction(digits) => digits.len(), + Representation::Compact(Quality(0 | 1000)) => 0, + Representation::Compact(Quality(value)) if value.is_multiple_of(100) => 1, + Representation::Compact(Quality(value)) if value.is_multiple_of(10) => 2, + Representation::Compact(_) => 3, + } + } + + fn digit(self, index: usize) -> u8 { + match self.repr { + Representation::Fraction(digits) => digits.get(index).copied().unwrap_or(b'0'), + Representation::Compact(Quality(value)) => { + let digit = match index { + 0 => value / 100 % 10, + 1 => value / 10 % 10, + 2 => value % 10, + _ => 0, + }; + u8::try_from(digit).expect("a decimal digit is at most nine") + b'0' + } + } + } + + pub(super) fn encoded_len(self) -> usize { + let fraction = self.fraction_len(); + if fraction == 0 { 1 } else { fraction + 2 } + } + + pub(super) fn append_to(self, bytes: &mut Vec) { + bytes.push(if self.is_one() { b'1' } else { b'0' }); + let fraction = self.fraction_len(); + if fraction != 0 { + bytes.push(b'.'); + for index in 0..fraction { + bytes.push(self.digit(index)); + } + } + } +} + +impl From for QualityView<'_> { + fn from(value: Quality) -> Self { + Self { + repr: Representation::Compact(value), + } + } +} + +impl TryFrom> for Quality { + type Error = InexactQuality; + + fn try_from(value: QualityView<'_>) -> Result { + value.to_quality() + } +} + +impl PartialEq for QualityView<'_> { + fn eq(&self, other: &Self) -> bool { + self.cmp(other) == Ordering::Equal + } +} + +impl Eq for QualityView<'_> {} + +impl PartialOrd for QualityView<'_> { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl Ord for QualityView<'_> { + fn cmp(&self, other: &Self) -> Ordering { + match (self.repr, other.repr) { + (Representation::Compact(left), Representation::Compact(right)) => left.cmp(&right), + // Canonical fractions have no trailing zeroes, so a shorter + // prefix sorts before the longer exact decimal value. + (Representation::Fraction(left), Representation::Fraction(right)) => left.cmp(right), + (Representation::Compact(left), Representation::Fraction(right)) => compare_compact_fraction(left, right), + (Representation::Fraction(left), Representation::Compact(right)) => compare_compact_fraction(right, left).reverse(), + } + } +} + +fn compare_compact_fraction(compact: Quality, fraction: &[u8]) -> Ordering { + let thousandths = u16::from(fraction[0] - b'0') * 100 + u16::from(fraction[1] - b'0') * 10 + u16::from(fraction[2] - b'0'); + // Fraction storage has a nonzero digit beyond thousandths; equality of + // the first three digits therefore puts the compact value strictly first. + compact.thousandths().cmp(&thousandths).then(Ordering::Less) +} + +impl Hash for QualityView<'_> { + fn hash(&self, state: &mut H) { + self.is_one().hash(state); + self.fraction_len().hash(state); + for index in 0..self.fraction_len() { + self.digit(index).hash(state); + } + } +} + +impl fmt::Display for QualityView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(if self.is_one() { "1" } else { "0" })?; + let fraction = self.fraction_len(); + if fraction != 0 { + f.write_str(".")?; + for index in 0..fraction { + write!(f, "{}", char::from(self.digit(index)))?; + } + } + Ok(()) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{Quality, QualityView}; + use crate::DecodeMode; + + #[test] + fn fractional_digits_are_zero_padded_beyond_the_retained_precision() { + let compact = QualityView::from(Quality::from_thousandths(125).unwrap()); + let fraction = QualityView::parse(b"0.1251", DecodeMode::Relaxed).unwrap(); + assert_eq!([0, 1, 2, 3, 4].map(|index| compact.digit(index)), *b"12500"); + assert_eq!([0, 1, 2, 3, 4].map(|index| fraction.digit(index)), *b"12510"); + assert_eq!(compact.digit(usize::MAX), b'0'); + assert_eq!(fraction.digit(usize::MAX), b'0'); + } +} diff --git a/crates/http_headers/src/headers/negotiation/recognition_test_support.rs b/crates/http_headers/src/headers/negotiation/recognition_test_support.rs new file mode 100644 index 000000000..5cd572cf4 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/recognition_test_support.rs @@ -0,0 +1,88 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::shared::try_plain_items; +use crate::DecodeError; + +/// Builds every string up to `length` bytes over `alphabet`. +/// Miri keeps every one/two-byte case and samples longer token/range spellings. +pub(super) fn exhaustive(alphabet: &[u8], length: usize) -> Vec> { + const MIRI_CONTINUATIONS: &[u8] = b"a*qQ01/-"; + + let mut all = vec![Vec::new()]; + let mut frontier = vec![Vec::new()]; + for depth in 0..length { + let mut next = Vec::new(); + for prefix in &frontier { + if cfg!(miri) + && depth >= 2 + && (!matches!(prefix.first(), Some(b'a' | b'q' | b'Q' | b'*')) + || !prefix.iter().skip(1).all(|byte| MIRI_CONTINUATIONS.contains(byte))) + { + continue; + } + for byte in alphabet { + if cfg!(miri) && depth >= 2 && !MIRI_CONTINUATIONS.contains(byte) { + continue; + } + let mut candidate = prefix.clone(); + candidate.push(*byte); + next.push(candidate); + } + } + all.extend_from_slice(&next); + frontier = next; + } + all +} + +/// Builds every concatenation of up to `count` of the given fragments. +/// Miri keeps all pairs, then samples quality, parameter, separator and quote continuations. +pub(super) fn fragment_lines(fragments: &[&str], count: usize) -> Vec> { + const MIRI_CONTINUATIONS: &[&str] = &[";q=0.9", ";q=1.001", ";q=0.1234", ";x", ";x=y", ",", " ", "\t", "\"", "\\"]; + + let mut lines = vec![Vec::new()]; + for depth in 0..count { + let mut next = Vec::new(); + for prefix in &lines { + if cfg!(miri) && depth >= 2 && !fragments.iter().take(3).any(|fragment| prefix.starts_with(fragment.as_bytes())) { + continue; + } + for fragment in fragments { + if cfg!(miri) && depth >= 2 && !MIRI_CONTINUATIONS.contains(fragment) { + continue; + } + let mut candidate = prefix.clone(); + candidate.extend_from_slice(fragment.as_bytes()); + next.push(candidate); + } + } + lines.extend_from_slice(&next); + } + lines +} + +pub(super) fn assert_recognition_is_sound( + lines: &[Vec], + recognize: fn(&[u8]) -> bool, + strict: fn(&[u8]) -> Result<(), DecodeError>, + relaxed: fn(&[u8]) -> Result<(), DecodeError>, +) -> usize { + let mut recognized = 0; + for line in lines { + if !recognize(line) { + continue; + } + recognized += 1; + let shown = String::from_utf8_lossy(line); + assert!( + matches!(try_plain_items(line, b',', true, strict), Ok(true)), + "recognized line {shown:?} must satisfy the strict member grammar" + ); + assert!( + matches!(try_plain_items(line, b',', true, relaxed), Ok(true)), + "recognized line {shown:?} must satisfy the relaxed member grammar" + ); + } + recognized +} diff --git a/crates/http_headers/src/headers/negotiation/server.rs b/crates/http_headers/src/headers/negotiation/server.rs new file mode 100644 index 000000000..171bc70d3 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/server.rs @@ -0,0 +1,380 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::shared::{invalid, invalid_syntax}; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +/// Defines the `Server` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 10.2.4](https://www.rfc-editor.org/rfc/rfc9110#section-10.2.4). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{Server, ServerOwned}; +/// +/// let mut map = HeaderMap::new(); +/// Server::insert(&mut map, ServerOwned::try_from("example/1")?)?; +/// assert!(Server::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct Server { + _private: (), +} + +/// Owned value for the `Server` header. +/// +/// The product/comment grammar is deliberately not exposed semantically; +/// callers can inspect the preserved wire bytes. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 10.2.4]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::ServerOwned::try_from("example/1")?; +/// assert_eq!(value.as_bytes(), b"example/1"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Server: nginx/1.25.3` identifies one product. +/// `Server: example-server/2.0 (internal)` includes a comment. +/// +/// [RFC 9110 section 10.2.4]: https://www.rfc-editor.org/rfc/rfc9110#section-10.2.4 +#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct ServerOwned(FieldValue); + +/// Borrowed value for the `Server` header. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{Server, ServerView}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let value: ServerView<'_> = Server::decode_view(FieldValueRef::new(b"nginx/1.25.3"))?; +/// assert_eq!(value.as_bytes(), b"nginx/1.25.3"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct ServerView<'a>(FieldValueRef<'a>); + +impl ServerOwned { + /// Returns the preserved wire bytes. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ServerOwned; + /// + /// let value = ServerOwned::try_from("nginx/1.25.3")?; + /// assert_eq!(value.as_bytes(), b"nginx/1.25.3"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_bytes(&self) -> &[u8] { + self.0.as_bytes() + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ServerOwned; + /// + /// let value = ServerOwned::try_from("Apache/2.4.58 (Unix)")?; + /// assert_eq!(value.as_field_value().as_bytes(), b"Apache/2.4.58 (Unix)"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.0 + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ServerOwned; + /// + /// let value = ServerOwned::try_from("nginx/1.25.3")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"nginx/1.25.3"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(ServerOwned, |value| value.0); + +impl<'a> ServerView<'a> { + /// Returns the preserved wire bytes. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::Server; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let value = Server::decode_view(FieldValueRef::new(b"Apache/2.4.58 (Unix)"))?; + /// assert_eq!(value.as_bytes(), b"Apache/2.4.58 (Unix)"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_bytes(self) -> &'a [u8] { + self.0.as_bytes() + } + + /// Returns the value as UTF-8. + /// + /// # Errors + /// + /// Returns an error when the field contains non-UTF-8 `obs-text`. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::Server; + /// use http_headers::{DecodeErrorKind, FieldValueRef, SingleValueField}; + /// + /// let value = Server::decode_view(FieldValueRef::new(b"nginx/1.25.3"))?; + /// assert_eq!(value.as_str()?, "nginx/1.25.3"); + /// + /// let binary = Server::decode_view(FieldValueRef::new(b"\xff"))?; + /// assert_eq!( + /// binary.as_str().unwrap_err().kind(), + /// DecodeErrorKind::InvalidUtf8 + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_str(self) -> Result<&'a str, DecodeError> { + self.0 + .to_str() + .map_err(|_invalid| invalid(&FieldName::Server, DecodeErrorKind::InvalidUtf8)) + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::Server; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let value = Server::decode_view(FieldValueRef::new(b"Apache/2.4.58 (Unix)"))?; + /// assert_eq!(value.as_field_value().as_bytes(), b"Apache/2.4.58 (Unix)"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.0 + } +} + +impl SingleValueField for Server { + type View<'a> = ServerView<'a>; + type Owned = ServerOwned; + + fn name() -> &'static FieldName { + &FieldName::Server + } + + #[inline] + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + validate_server(value)?; + Ok(ServerView(value)) + } + + #[inline] + fn decode_owned(value: FieldValue) -> Result { + ServerOwned::try_from(value) + } + + #[inline] + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.0 + } + + #[inline] + fn into_field_value(value: Self::Owned) -> FieldValue { + value.0 + } +} + +super::super::shared::impl_string_conversions!(ServerOwned, &FieldName::Server, invalid_syntax, value); + +impl TryFrom for ServerOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + validate_server(value.as_field_value_ref())?; + Ok(Self(value)) + } +} + +#[inline] +fn validate_server(value: FieldValueRef<'_>) -> Result<(), DecodeError> { + if crate::validate::field_value(value.as_bytes()) && super::super::has_non_ows(value.as_bytes()) { + Ok(()) + } else { + Err(invalid_syntax(&FieldName::Server)) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::collections::hash_map::DefaultHasher; + use std::hash::{Hash, Hasher}; + + use super::{Server, ServerOwned, validate_server}; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef, SingleValueField}; + + struct RawSource<'a>(FieldValueRef<'a>); + + impl FieldSource for RawSource<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_borrowed(name, std::slice::from_ref(&self.0)) + } + } + + #[test] + fn direct_and_source_decoders_reject_every_forbidden_field_byte() { + for byte in (0..=0x1f).chain(std::iter::once(0x7f)).filter(|byte| *byte != b'\t') { + let wire = [b'x', byte, b'y']; + let value = FieldValueRef::new(&wire); + let source = RawSource(value); + assert_eq!(Server::decode_view(value).unwrap_err().kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!( + ServerOwned::try_from(String::from_utf8(wire.to_vec()).unwrap()).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!( + Server::decode_view_with(value, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!(Server::view_with(&source, mode).unwrap_err().kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!( + Server::owned_with(&source, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + } + } + } + + #[test] + fn opaque_wire_bytes_and_sensitivity_survive_both_decode_modes() { + for wire in [b" \tserver/1 (opaque)\t ".as_slice(), b" \t\x80\xff\t "] { + for sensitive in [false, true] { + let value = FieldValueRef::new(wire).with_sensitive(sensitive); + let source = RawSource(value); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let direct = Server::decode_view_with(value, mode).unwrap(); + let view = Server::view_with(&source, mode).unwrap().unwrap(); + let owned = Server::decode_owned_with(value.try_to_field_value().unwrap(), mode).unwrap(); + let sourced = Server::owned_with(&source, mode).unwrap().unwrap(); + assert_eq!(direct.as_bytes(), wire); + assert_eq!(view.as_bytes(), wire); + assert_eq!(direct.as_field_value().is_sensitive(), sensitive); + assert_eq!(view.as_field_value().is_sensitive(), sensitive); + assert_eq!(owned, sourced); + assert_eq!(owned.as_field_value().as_bytes(), wire); + assert_eq!(owned.as_field_value().is_sensitive(), sensitive); + if wire.contains(&0xff) { + assert_eq!(direct.as_str().unwrap_err().kind(), DecodeErrorKind::InvalidUtf8); + } + #[cfg(feature = "http")] + { + let mut map = http::HeaderMap::new(); + map.insert(http::header::SERVER, http::HeaderValue::try_from(value).unwrap()); + let view = Server::view_with(&map, mode).unwrap().unwrap(); + let mapped = Server::owned_with(&map, mode).unwrap().unwrap(); + assert_eq!(view.as_bytes(), wire); + assert_eq!(view.as_field_value().is_sensitive(), sensitive); + assert_eq!(mapped, owned); + assert_eq!(mapped.as_field_value().is_sensitive(), sensitive); + } + } + } + } + } + + #[test] + fn owned_and_borrowed_server_values_preserve_opaque_wire_data() { + assert_eq!( + ServerOwned::try_from("borrowed/1").expect("valid borrowed server").as_bytes(), + b"borrowed/1" + ); + let owned = ServerOwned::try_from(String::from("example/1 (test)")).expect("nonempty server is valid"); + assert_eq!(owned.as_bytes(), b"example/1 (test)"); + assert_eq!(owned.as_field_value().as_bytes(), b"example/1 (test)"); + assert!(format!("{owned:?}").contains("example/1 (test)")); + let mut hasher = DefaultHasher::new(); + owned.hash(&mut hasher); + assert_ne!(hasher.finish(), 0); + assert_eq!(owned.into_field_value().as_bytes(), b"example/1 (test)"); + + let view = ::decode_view(FieldValueRef::new(b"server/2")).expect("borrowed server is valid"); + assert_eq!(view.as_bytes(), b"server/2"); + assert_eq!(view.as_str(), Ok("server/2")); + assert_eq!(view.as_field_value().as_bytes(), b"server/2"); + + let decoded = ::decode_owned(FieldValue::from_static("server/3")).expect("owned server"); + assert_eq!(::as_field_value(&decoded).as_bytes(), b"server/3"); + assert_eq!(::into_field_value(decoded).as_bytes(), b"server/3"); + } + + #[test] + fn server_rejects_empty_values_and_reports_non_utf8_views() { + for empty in [b"".as_slice(), b" ", b"\t"] { + assert_eq!( + validate_server(FieldValueRef::new(empty)) + .expect_err("OWS-only values are invalid") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + assert_eq!( + ServerOwned::try_from(String::from("line\nbreak")) + .expect_err("invalid field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ServerOwned::try_from("line\nbreak") + .expect_err("invalid borrowed field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ::decode_owned(FieldValue::from_static(" ")) + .expect_err("OWS-only owned value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ServerOwned::try_from(FieldValue::from_bytes(b"\xff").expect("obs-text is field-safe")) + .expect("opaque obs-text is valid") + .as_bytes(), + b"\xff" + ); + let view = ::decode_view(FieldValueRef::new(b"\xff")).expect("opaque obs-text is valid"); + assert_eq!( + view.as_str().expect_err("obs-text is not UTF-8").kind(), + DecodeErrorKind::InvalidUtf8 + ); + assert_eq!(::name(), &FieldName::Server); + } +} diff --git a/crates/http_headers/src/headers/negotiation/shared.rs b/crates/http_headers/src/headers/negotiation/shared.rs new file mode 100644 index 000000000..c20a2cf67 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/shared.rs @@ -0,0 +1,2193 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::hash::{Hash, Hasher}; +use std::slice; + +use http_headers_simd::{EmptyMembers, TokenListScan}; + +use crate::source::{FieldLines, MAX_CUSTOM_LIST_ITEMS}; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, validate}; + +/// The preserved field lines of a list header. +/// +/// A repeated list header is rare, so the sole line is held on its own rather +/// than behind the length and capacity a growable buffer would carry. +#[derive(Clone, Debug)] +pub(super) enum ListValues { + One(FieldValue), + Many(Vec), +} + +impl ListValues { + /// Borrows the field lines as one contiguous run. + #[inline] + fn as_slice(&self) -> &[FieldValue] { + match self { + Self::One(line) => slice::from_ref(line), + Self::Many(lines) => lines, + } + } + + #[inline] + pub(super) fn len(&self) -> usize { + self.as_slice().len() + } + + #[inline] + pub(super) fn iter(&self) -> slice::Iter<'_, FieldValue> { + self.as_slice().iter() + } +} + +impl Eq for ListValues {} + +impl PartialEq for ListValues { + fn eq(&self, other: &Self) -> bool { + self.as_slice() == other.as_slice() + } +} + +impl Hash for ListValues { + fn hash(&self, state: &mut H) { + self.as_slice().hash(state); + } +} + +/// Copies the preserved field lines, without allocating for a single line. +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the validator into the field line walk" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +fn clone_checked_values( + values: &FieldLines<'_>, + mut validate_value: impl FnMut(usize, FieldValueRef<'_>) -> Result<(), DecodeError>, +) -> Result { + let mut lines = values.repeated_owned()?; + let (first, first_owned) = lines.next().expect("untyped values always contain a field line"); + validate_value(0, first)?; + let Some((second, second_owned)) = lines.next() else { + return Ok(ListValues::One(first_owned)); + }; + validate_value(1, second)?; + clone_remaining_values(lines, first_owned, second_owned, validate_value) +} + +/// Copies the field lines after the second, which only repeated headers have. +/// +/// Keeping the growable buffer out of line leaves the single-line path, which +/// is the overwhelmingly common one, free of its register pressure. +#[inline(never)] +fn clone_remaining_values<'a>( + lines: impl Iterator, FieldValue)>, + first: FieldValue, + second: FieldValue, + mut validate_value: impl FnMut(usize, FieldValueRef<'_>) -> Result<(), DecodeError>, +) -> Result { + let mut copied = vec![first, second]; + for (index, (line, owned)) in (2_usize..).zip(lines) { + validate_value(index, line)?; + copied.push(owned); + } + Ok(ListValues::Many(copied)) +} + +macro_rules! decode_owned_list { + (quoted, $values:expr, $name:expr, $mode:expr, $strict:expr, $relaxed:expr) => { + super::shared::checked_owned_values($values, $name, $mode, $strict, $relaxed, super::shared::decode_quoted_owned_list) + }; + (token, $values:expr, $name:expr, $mode:expr, $strict:expr, $relaxed:expr) => { + super::shared::checked_owned_values($values, $name, $mode, $strict, $relaxed, super::shared::decode_token_owned_list) + }; +} + +/// Validates one already-delimited list member. +pub(super) type ItemValidator = fn(&[u8]) -> Result<(), DecodeError>; + +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the header name and item validator into the field line walk" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn decode_quoted_owned_list( + name: &'static FieldName, + values: &FieldLines<'_>, + validator: ItemValidator, +) -> Result { + clone_checked_values(values, |index, value| { + check_quoted_value(name, value, validator).map_err(|error| { + if error.kind() == DecodeErrorKind::UnterminatedQuote { + error.at_value(index) + } else { + error + } + }) + }) +} + +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the header name and item validator into the field line walk" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn decode_token_owned_list( + name: &'static FieldName, + values: &FieldLines<'_>, + validator: ItemValidator, +) -> Result { + clone_checked_values(values, |index, value| { + if scan_token_list(value.as_bytes()) { + Ok(()) + } else { + Err(token_list_error(name, value.as_bytes(), validator, Some(index))) + } + }) +} + +pub(super) type ValuesValidator = + for<'a> fn(&'static FieldName, &FieldLines<'a>, ItemValidator) -> Result>, DecodeError>; + +pub(super) type OwnedValuesDecoder = fn(&'static FieldName, &FieldLines<'_>, ItemValidator) -> Result; + +fn validate_custom_values(values: &FieldLines<'_>, name: &'static FieldName, validator: ItemValidator) -> Result<(), DecodeError> { + values.validate_custom_source_bounds()?; + + let mut item_count = 0_usize; + for value in values.repeated() { + for item in QuotedItems::comma(value.as_bytes(), name) { + item_count = item_count + .checked_add(1) + .ok_or_else(|| invalid(name, DecodeErrorKind::SourceLimitExceeded))?; + if item_count > MAX_CUSTOM_LIST_ITEMS { + return Err(invalid(name, DecodeErrorKind::SourceLimitExceeded)); + } + validator(item?)?; + } + } + Ok(()) +} + +#[expect( + clippy::inline_always, + reason = "specializing per call site turns the decoder and validator pointers into direct calls" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn checked_view_values<'a>( + values: Option>, + name: &'static FieldName, + mode: crate::DecodeMode, + strict: ItemValidator, + relaxed: ItemValidator, + validate_values: ValuesValidator, +) -> Result>, DecodeError> { + let Some(values) = values else { + return Ok(None); + }; + if values.has_custom_source_limits() { + validate_custom_values( + &values, + name, + match mode { + crate::DecodeMode::Strict => strict, + crate::DecodeMode::Relaxed => relaxed, + }, + )?; + return Ok(Some(values)); + } + match mode { + crate::DecodeMode::Strict => validate_values(name, &values, strict), + crate::DecodeMode::Relaxed => validate_values(name, &values, relaxed), + } + .map(|_single| Some(values)) +} + +#[expect( + clippy::inline_always, + reason = "specializing per call site turns the decoder and validator pointers into direct calls" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn checked_owned_values( + values: Option>, + name: &'static FieldName, + mode: crate::DecodeMode, + strict: ItemValidator, + relaxed: ItemValidator, + decode_values: OwnedValuesDecoder, +) -> Result, DecodeError> { + let Some(values) = values else { + return Ok(None); + }; + if values.has_custom_source_limits() { + validate_custom_values( + &values, + name, + match mode { + crate::DecodeMode::Strict => strict, + crate::DecodeMode::Relaxed => relaxed, + }, + )?; + return clone_checked_values(&values, |_index, _value| Ok(())).map(Some); + } + match mode { + crate::DecodeMode::Strict => decode_values(name, &values, strict), + crate::DecodeMode::Relaxed => decode_values(name, &values, relaxed), + } + .map(Some) +} + +pub(super) fn encoded_list(values: ListValues) -> crate::sink::EncodedValues { + match values { + ListValues::One(line) => crate::sink::EncodedValues::single(line), + ListValues::Many(lines) => crate::sink::EncodedValues::from_vec(lines), + } +} + +macro_rules! owned_list_items { + (quoted, $self:expr, $name:expr) => { + $self + .values + .iter() + .flat_map(|value| super::shared::QuotedItems::comma(value.as_bytes(), $name)) + .filter_map(Result::ok) + }; + (token, $self:expr, $name:expr) => { + $self + .values + .iter() + .flat_map(|value| value.as_bytes().split(|byte| *byte == b',')) + .map(validate::trim_ows) + .filter(|item| !item.is_empty()) + }; +} + +macro_rules! borrowed_list_items { + (quoted, $self:expr) => { + $self.values.comma_items().filter_map(Result::ok) + }; + (token, $self:expr) => { + $self + .values + .repeated() + .flat_map(|value| value.as_bytes().split(|byte| *byte == b',')) + .map(validate::trim_ows) + .filter(|item| !item.is_empty()) + }; +} + +macro_rules! list_header { + ( + $descriptor:ident, + $owned:ident, + $view:ident, + $header_name:literal, + $specification:literal, + $name:expr, + $validator:ident, + $relaxed_validator:ident, + $check_values:ident, + $check_value:ident, + $item_style:ident + ) => { + #[doc = concat!("Defines the `", $header_name, "` header.")] + #[doc = ""] + #[doc = "# Specification"] + #[doc = ""] + #[doc = $specification] + #[derive(Debug)] + pub struct $descriptor { + _private: (), + } + + impl Clone for $owned { + fn clone(&self) -> Self { + Self { + values: self.values.clone(), + } + } + } + + impl std::fmt::Debug for $owned { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct(stringify!($owned)) + .field("value_count", &self.values.len()) + .finish() + } + } + + impl Eq for $owned {} + + impl PartialEq for $owned { + fn eq(&self, other: &Self) -> bool { + self.values == other.values + } + } + + impl std::hash::Hash for $owned { + fn hash(&self, state: &mut H) { + self.values.hash(state); + } + } + + impl std::fmt::Debug for $view<'_> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct(stringify!($view)) + .field("value_count", &self.values.len()) + .finish() + } + } + + impl $owned { + /// Iterates raw list members in wire order. + /// + /// Empty RFC list members are ignored. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AcceptEncodingOwned; + /// + /// let value = AcceptEncodingOwned::try_from("gzip, br")?; + /// let items: Vec<&[u8]> = value.items().collect(); + /// assert_eq!(items, vec![b"gzip".as_slice(), b"br".as_slice()]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn items(&self) -> impl Iterator { + super::shared::owned_list_items!($item_style, self, $name) + } + + /// Iterates the preserved field lines. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AllowOwned; + /// + /// let value = AllowOwned::try_from("GET, HEAD")?; + /// let mut values = value.values(); + /// assert_eq!(values.len(), 1); + /// assert_eq!( + /// values.next().expect("one field line").as_bytes(), + /// b"GET, HEAD" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn values(&self) -> impl ExactSizeIterator> { + self.values.iter().map(FieldValue::as_field_value_ref) + } + } + + impl<'a> $view<'a> { + /// Iterates raw list members in wire order without allocating. + /// + /// Empty RFC list members are ignored. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{AcceptEncoding, AcceptEncodingView}; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source; + /// + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::AcceptEncoding).then(|| FieldLines::single(name, b"gzip, br")) + /// } + /// } + /// + /// let value: AcceptEncodingView<'_> = AcceptEncoding::view(&Source)?.expect("header is present"); + /// let items = value.items().collect::>(); + /// assert_eq!(items, vec![b"gzip".as_slice(), b"br".as_slice()]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn items(&self) -> impl Iterator + '_ { + super::shared::borrowed_list_items!($item_style, self) + } + + /// Iterates the original field lines. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{Allow, AllowView}; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source; + /// + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::Allow).then(|| FieldLines::single(name, b"GET, HEAD")) + /// } + /// } + /// + /// let value: AllowView<'_> = Allow::view(&Source)?.expect("header is present"); + /// let values: Vec<&[u8]> = value.values().map(|line| line.as_bytes()).collect(); + /// assert_eq!(values, vec![b"GET, HEAD".as_slice()]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn values(&self) -> impl Iterator> + '_ { + self.values.repeated() + } + } + + impl Field for $descriptor { + type View<'a> = $view<'a>; + type Owned = $owned; + + fn name() -> &'static FieldName { + $name + } + + #[inline] + fn view_with(source: &S, mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + super::shared::checked_view_values( + source.lines(Self::name()), + $name, + mode, + $validator, + $relaxed_validator, + $check_values, + ) + .map(|values| values.map(|values| $view { values })) + } + + #[inline] + fn owned_with(source: &S, mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + super::shared::decode_owned_list!( + $item_style, + source.lines(Self::name()), + $name, + mode, + $validator, + $relaxed_validator + ) + .map(|values| values.map(|values| $owned { values })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), super::shared::encoded_list(value.values)) + } + } + + impl TryFrom<&str> for $owned { + type Error = DecodeError; + + fn try_from(value: &str) -> Result { + let value = FieldValue::from_str(value).map_err(|_invalid| super::shared::invalid_syntax($name))?; + Self::try_from(value) + } + } + + impl TryFrom for $owned { + type Error = DecodeError; + + fn try_from(value: String) -> Result { + let value = FieldValue::try_from(value).map_err(|_invalid| super::shared::invalid_syntax($name))?; + Self::try_from(value) + } + } + + impl TryFrom for $owned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + $check_value($name, value.as_field_value_ref(), $validator)?; + Ok(Self { + values: super::shared::ListValues::One(value), + }) + } + } + }; +} + +/// Validates every member of a list whose members may contain quoted strings. +/// +/// The sole field line is never reported, because the delimited iteration this +/// performs does not reveal the field line boundaries it crosses. +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the header name and item validator into the scan" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn check_quoted_values<'a>( + name: &'static FieldName, + values: &FieldLines<'a>, + validator: ItemValidator, +) -> Result>, DecodeError> { + let mut lines = values.repeated(); + let first = lines.next().expect("untyped values always contain a field line"); + let single = lines.next().is_none(); + check_quoted_value(name, first, validator)?; + if single { + return Ok(Some(first)); + } + check_repeated_quoted_values(name, values, validator).map(|()| None) +} + +/// Scans the field lines after the first one, which only repeated headers have. +/// +/// A repeated negotiation header is rare, so keeping this cold and out of line +/// leaves the single-line path a short straight run. +#[cold] +#[inline(never)] +fn check_repeated_quoted_values(name: &'static FieldName, values: &FieldLines<'_>, validator: ItemValidator) -> Result<(), DecodeError> { + for line in values.repeated().skip(1) { + check_quoted_value(name, line, validator)?; + } + Ok(()) +} + +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the header name and item validator into the scan" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn check_quoted_value(name: &'static FieldName, value: FieldValueRef<'_>, validator: ItemValidator) -> Result<(), DecodeError> { + if is_well_known_negotiation_line(name, value.as_bytes()) { + return Ok(()); + } + if recognizes_whole_line(name, value.as_bytes()) { + return Ok(()); + } + if try_plain_items(value.as_bytes(), b',', true, validator)? { + return Ok(()); + } + check_quoted_members(name, value.as_bytes(), validator) +} + +/// Validates members of a line that carries quoted syntax. +/// +/// Quoting is rare in negotiation values, so this carries the quote and escape +/// state that [`try_plain_items`] deliberately leaves out. +#[cold] +#[inline(never)] +fn check_quoted_members(name: &'static FieldName, bytes: &[u8], validator: ItemValidator) -> Result<(), DecodeError> { + for item in QuotedItems::comma(bytes, name) { + validator(item?)?; + } + Ok(()) +} + +/// Recognizes a whole negotiation line without splitting it into members. +/// +/// The content negotiation headers each carry quality parameters, which makes +/// their member grammar rich enough to be worth a dedicated whole-line pass; +/// every other header falls straight through. A `false` answer always means +/// "not recognized", never "malformed", so the general parser still produces +/// every diagnostic. +fn recognizes_whole_line(name: &'static FieldName, bytes: &[u8]) -> bool { + match *name { + FieldName::Accept => super::accept_scan::scan_accept_line(bytes), + FieldName::AcceptEncoding => super::weighted_token_scan::scan_accept_encoding_line(bytes), + FieldName::AcceptLanguage => super::weighted_token_scan::scan_accept_language_line(bytes), + _ => false, + } +} + +/// The `Accept` lines recognized without walking the member grammar. +pub(super) const WELL_KNOWN_ACCEPT: &[&[u8]] = &[b"*/*", b"application/json", b"text/html"]; + +/// The `Accept-Encoding` lines recognized without walking the member grammar. +pub(super) const WELL_KNOWN_ACCEPT_ENCODING: &[&[u8]] = &[b"gzip", b"br", b"identity"]; + +/// The `Accept-Language` lines recognized without walking the member grammar. +pub(super) const WELL_KNOWN_ACCEPT_LANGUAGE: &[&[u8]] = &[b"en", b"en-US", b"*"]; + +/// Recognizes negotiation lines whose whole text is a single common token. +/// +/// A client that states one preference sends the same handful of lines over +/// and over, and matching the whole line settles them without walking the +/// member grammar. Lines carrying several members or quality parameters vary +/// too much between clients to be worth a literal, so they are parsed. +/// +/// Every recognized line also satisfies the grammar it skips, which the tests +/// beside each header's validator check. +fn is_well_known_negotiation_line(name: &'static FieldName, bytes: &[u8]) -> bool { + let recognized = if name == &FieldName::Accept { + WELL_KNOWN_ACCEPT + } else if name == &FieldName::AcceptEncoding { + WELL_KNOWN_ACCEPT_ENCODING + } else if name == &FieldName::AcceptLanguage { + WELL_KNOWN_ACCEPT_LANGUAGE + } else { + return false; + }; + recognized.contains(&bytes) +} + +/// Validates delimiter-separated items until quoted syntax requires fallback. +/// +/// Common negotiation values contain no quoted strings. Keeping that path to +/// delimiter checks and trimming avoids carrying quote and escape state +/// through every byte; encountering either marker leaves validation to +/// [`QuotedItems`] from the beginning. +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the item validator into the delimiter scan" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn try_plain_items( + bytes: &[u8], + delimiter: u8, + skip_empty: bool, + mut validator: impl FnMut(&[u8]) -> Result<(), DecodeError>, +) -> Result { + let mut start = 0; + for (position, byte) in bytes.iter().copied().enumerate() { + if matches!(byte, b'"' | b'\\') { + return Ok(false); + } + if byte == delimiter { + let item = validate::trim_ows(&bytes[start..position]); + if !skip_empty || !item.is_empty() { + validator(item)?; + } + start = position + 1; + } + } + let item = validate::trim_ows(&bytes[start..]); + if !skip_empty || !item.is_empty() { + return validator(item).map(|()| true); + } + Ok(true) +} + +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the header name and item validator into the scan" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn check_token_values<'a>( + name: &'static FieldName, + values: &FieldLines<'a>, + validator: ItemValidator, +) -> Result>, DecodeError> { + let mut lines = values.repeated(); + let first = lines.next().expect("untyped values always contain a field line"); + let single = lines.next().is_none(); + if !is_well_known_token_list(first.as_bytes()) { + return check_unusual_token_values(name, values, first, single, validator); + } + if single { + return Ok(Some(first)); + } + check_repeated_token_values(name, values, validator).map(|()| None) +} + +/// Finishes a decode whose first field line is not one of the well-known ones. +/// +/// Taking the whole tail out of line, rather than just the scan, keeps the +/// recognized-line path free of the register pressure a call site imposes. +#[inline(never)] +fn check_unusual_token_values<'a>( + name: &'static FieldName, + values: &FieldLines<'a>, + first: FieldValueRef<'a>, + single: bool, + validator: ItemValidator, +) -> Result>, DecodeError> { + if !scan_token_list_general(first.as_bytes()) { + return Err(token_list_error(name, first.as_bytes(), validator, Some(0))); + } + if single { + return Ok(Some(first)); + } + check_repeated_token_values(name, values, validator).map(|()| None) +} + +/// Scans the field lines after the first one, which only repeated headers have. +/// +/// A repeated list header is rare, so keeping this cold and out of line, and +/// rebuilding the walk from `values` rather than taking a half-consumed one, +/// leaves the single-line path a short straight run. +#[cold] +#[inline(never)] +fn check_repeated_token_values(name: &'static FieldName, values: &FieldLines<'_>, validator: ItemValidator) -> Result<(), DecodeError> { + for (offset, line) in values.repeated().enumerate().skip(1) { + if !scan_token_list(line.as_bytes()) { + return Err(token_list_error(name, line.as_bytes(), validator, Some(offset))); + } + } + Ok(()) +} + +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the header name and item validator into the scan" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn check_token_value(name: &'static FieldName, value: FieldValueRef<'_>, validator: ItemValidator) -> Result<(), DecodeError> { + if scan_token_list(value.as_bytes()) { + Ok(()) + } else { + Err(token_list_error(name, value.as_bytes(), validator, None)) + } +} + +/// Reads the eight bytes at `offset`, or zero when they are not all present. +#[expect(clippy::inline_always, reason = "the callers pass constant offsets that fold the bounds test away")] +#[cfg_attr(not(coverage_nightly), inline(always))] +fn word(bytes: &[u8], offset: usize) -> u64 { + let mut chunk = [0_u8; 8]; + if let Some(source) = bytes.get(offset..offset + 8) { + chunk.copy_from_slice(source); + } + u64::from_le_bytes(chunk) +} + +/// Returns whether `bytes` equals `expected`, comparing eight bytes at a time. +/// +/// The final word overlaps the one before it, so any length is covered by +/// `len / 8` rounded up comparisons rather than one comparison per byte. +#[expect( + clippy::inline_always, + reason = "expected is always a literal, which folds the loop and the words away" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +fn equals_word_at_a_time(bytes: &[u8], expected: &[u8]) -> bool { + if bytes.len() != expected.len() { + return false; + } + if bytes.len() < 8 { + return bytes == expected; + } + let mut offset = 0; + while offset + 8 < bytes.len() { + if word(bytes, offset) != word(expected, offset) { + return false; + } + offset += 8; + } + match bytes.len() - offset { + // One or two trailing bytes are cheaper read on their own than as a + // word overlapping the one before them. + 1 => bytes.get(offset) == expected.get(offset), + 2 => half_word(bytes, offset) == half_word(expected, offset), + _ => { + let last = bytes.len() - 8; + word(bytes, last) == word(expected, last) + } + } +} + +/// Reads the two bytes at `offset`, or zero when they are not both present. +#[expect(clippy::inline_always, reason = "the callers pass constant offsets that fold the bounds test away")] +#[cfg_attr(not(coverage_nightly), inline(always))] +fn half_word(bytes: &[u8], offset: usize) -> u16 { + let mut chunk = [0_u8; 2]; + if let Some(source) = bytes.get(offset..offset + 2) { + chunk.copy_from_slice(source); + } + u16::from_le_bytes(chunk) +} + +/// The field lines short enough that byte comparison is already cheap. +/// +/// This mirrors [`is_well_known_token_list`] so the tests can prove every +/// accepted line is valid. +#[cfg(test)] +const SHORT_WELL_KNOWN_TOKEN_LISTS: &[&[u8]] = &[ + b"*", b"GET", b"PUT", b"HEAD", b"POST", b"PATCH", b"DELETE", b"OPTIONS", b"accept", b"cookie", b"origin", +]; + +/// The field lines compared eight bytes at a time, grouped by length. +/// +/// This mirrors [`is_well_known_token_list`] so the tests can prove every +/// accepted line is valid. +#[cfg(test)] +const LONG_WELL_KNOWN_TOKEN_LISTS: &[&[u8]] = &[ + b"GET, HEAD", + b"GET, POST", + b"HEAD, GET", + b"user-agent", + b"authorization", + b"GET, HEAD, POST", + b"accept-encoding", + b"accept-language", + b"GET, HEAD, OPTIONS", + b"GET, POST, OPTIONS", + b"OPTIONS, GET, HEAD", + b"accept, accept-encoding", + b"accept-encoding, origin", + b"origin, accept-encoding", + b"accept-encoding, accept-language", +]; + +/// Returns whether a field line is one of the well-known valid token lists. +/// +/// These lines carry the overwhelming majority of real `Allow` and `Vary` +/// traffic. Recognizing one costs a load and a compare per eight bytes, where +/// the general scan costs several instructions per byte, and every line the +/// table accepts is a valid bare token list, so a hit may skip the scan. +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the table into a switch on the line length" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +fn is_well_known_token_list(bytes: &[u8]) -> bool { + if bytes.len() < 8 { + return matches!( + bytes, + b"*" | b"GET" | b"PUT" | b"HEAD" | b"POST" | b"PATCH" | b"DELETE" | b"OPTIONS" | b"accept" | b"cookie" | b"origin" + ); + } + match bytes.len() { + 9 => { + equals_word_at_a_time(bytes, b"GET, POST") + || equals_word_at_a_time(bytes, b"GET, HEAD") + || equals_word_at_a_time(bytes, b"HEAD, GET") + } + 10 => equals_word_at_a_time(bytes, b"user-agent"), + 13 => equals_word_at_a_time(bytes, b"authorization"), + 15 => { + equals_word_at_a_time(bytes, b"accept-encoding") + || equals_word_at_a_time(bytes, b"accept-language") + || equals_word_at_a_time(bytes, b"GET, HEAD, POST") + } + 18 => { + equals_word_at_a_time(bytes, b"GET, HEAD, OPTIONS") + || equals_word_at_a_time(bytes, b"GET, POST, OPTIONS") + || equals_word_at_a_time(bytes, b"OPTIONS, GET, HEAD") + } + 23 => { + equals_word_at_a_time(bytes, b"accept-encoding, origin") + || equals_word_at_a_time(bytes, b"origin, accept-encoding") + || equals_word_at_a_time(bytes, b"accept, accept-encoding") + } + 32 => equals_word_at_a_time(bytes, b"accept-encoding, accept-language"), + _ => false, + } +} + +/// Returns whether a field line is a comma-delimited list of bare tokens. +/// +/// Optional whitespace around members and empty members are ignored, so this +/// accepts exactly the lines whose members are all valid tokens. Rejected +/// lines are re-scanned by [`token_list_error`] to describe the failure. +#[inline] +pub(super) fn scan_token_list(bytes: &[u8]) -> bool { + is_well_known_token_list(bytes) || scan_token_list_general(bytes) +} + +/// Returns whether a field line is a comma-delimited list of bare tokens. +/// +/// This classifies every byte, so it accepts far more than the well-known +/// table and costs far more to do it. Lines long enough to be worth it are +/// folded a vector register at a time by [`http_headers_simd`]. +fn scan_token_list_general(bytes: &[u8]) -> bool { + http_headers_simd::scan_token_list(bytes, EmptyMembers::Skip) != TokenListScan::Rejected +} + +/// Describes why a field line is not a list of bare tokens. +/// +/// A line only reaches this function once [`scan_token_list`] has rejected it, +/// so the delimiter-aware scan it performs always fails as well. +#[cold] +fn token_list_error(name: &'static FieldName, bytes: &[u8], validator: ItemValidator, value_index: Option) -> DecodeError { + for item in QuotedItems::comma(bytes, name) { + match item { + Ok(item) => { + if let Err(error) = validator(item) { + return error; + } + } + Err(error) => { + return match value_index { + Some(index) => error.at_value(index), + None => error, + }; + } + } + } + invalid(name, DecodeErrorKind::InvalidToken) +} + +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the member grammar into the caller's item loop" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +pub(super) fn validate_weighted_token( + bytes: &[u8], + name: &'static FieldName, + item_validator: fn(&[u8]) -> bool, + relaxed: bool, +) -> Result<(), DecodeError> { + // A member without a parameter or quoting is just its own token, so the + // segment walk below would find one segment and validate exactly this. + if !carries_parameter_syntax(bytes) { + return if item_validator(validate::trim_ows(bytes)) { + Ok(()) + } else { + Err(invalid(name, DecodeErrorKind::InvalidToken)) + }; + } + if validate_weighted_token_plain(bytes, name, item_validator, relaxed)? { + return Ok(()); + } + validate_weighted_token_quoted(bytes, name, item_validator, relaxed) +} + +/// Reports whether a list member carries parameter or quoting syntax. +#[inline] +fn carries_parameter_syntax(bytes: &[u8]) -> bool { + bytes.iter().any(|byte| matches!(byte, b';' | b'"' | b'\\')) +} + +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the member grammar into the caller's item loop" +)] +#[cfg_attr(not(coverage_nightly), inline(always))] +fn validate_weighted_token_plain( + bytes: &[u8], + name: &'static FieldName, + item_validator: fn(&[u8]) -> bool, + relaxed: bool, +) -> Result { + let mut segment = 0_u8; + try_plain_items(bytes, b';', false, |bytes| { + segment = segment.saturating_add(1); + match segment { + 1 if item_validator(bytes) => Ok(()), + 1 => Err(invalid(name, DecodeErrorKind::InvalidToken)), + 2 if bytes.len() >= 2 && bytes[0].eq_ignore_ascii_case(&b'q') && bytes[1] == b'=' => { + validate_quality(&bytes[2..], name, relaxed) + } + 2 => { + let (parameter, value, compact) = parse_parameter(bytes, false, name)?; + if (!relaxed && !compact) || !validate::eq_ignore_ascii_case(parameter, b"q") { + return Err(invalid_syntax(name)); + } + validate_quality(value.expect("required parameters always contain a value"), name, relaxed) + } + _ => Err(invalid_syntax(name)), + } + }) +} + +#[cold] +#[inline(never)] +fn validate_weighted_token_quoted( + bytes: &[u8], + name: &'static FieldName, + item_validator: fn(&[u8]) -> bool, + relaxed: bool, +) -> Result<(), DecodeError> { + let mut segments = QuotedItems::semicolon(bytes, name); + let item = segments.next().expect("semicolon iteration always yields a first item")?; + if !item_validator(item) { + return Err(invalid(name, DecodeErrorKind::InvalidToken)); + } + if let Some(weight) = segments.next() { + let (parameter, value, compact) = parse_parameter(weight?, false, name)?; + if (!relaxed && !compact) || !validate::eq_ignore_ascii_case(parameter, b"q") { + return Err(invalid_syntax(name)); + } + let quality = value.expect("required parameters always contain a value"); + validate_quality(quality, name, relaxed)?; + } + if let Some(extra) = segments.next() { + extra?; + return Err(invalid_syntax(name)); + } + Ok(()) +} + +pub(super) type ParsedParameter<'a> = (&'a [u8], Option<&'a [u8]>, bool); + +pub(super) fn parse_parameter<'a>( + bytes: &'a [u8], + value_optional: bool, + header: &'static FieldName, +) -> Result, DecodeError> { + let equals = bytes.iter().position(|byte| *byte == b'='); + let (name, value, compact) = if let Some(equals) = equals { + let raw_name = &bytes[..equals]; + let raw_value = &bytes[equals + 1..]; + let name = validate::trim_ows(raw_name); + let value = validate::trim_ows(raw_value); + if !valid_parameter_value(value) { + return Err(invalid_syntax(header)); + } + (name, Some(value), raw_name.len() == name.len() && raw_value.len() == value.len()) + } else if value_optional { + (validate::trim_ows(bytes), None, true) + } else { + return Err(invalid_syntax(header)); + }; + if !validate::token(name) { + return Err(invalid(header, DecodeErrorKind::InvalidToken)); + } + Ok((name, value, compact)) +} + +pub(super) fn valid_parameter_value(bytes: &[u8]) -> bool { + validate::token(bytes) || valid_quoted_string(bytes) +} + +fn valid_quoted_string(bytes: &[u8]) -> bool { + if bytes.len() < 2 || bytes.first() != Some(&b'"') || bytes.last() != Some(&b'"') { + return false; + } + let mut escaped = false; + for byte in bytes[1..bytes.len() - 1].iter().copied() { + if escaped { + if !matches!(byte, b'\t' | b' '..=b'~' | 0x80..=0xff) { + return false; + } + escaped = false; + } else if byte == b'\\' { + escaped = true; + } else if !matches!(byte, b'\t' | b' ' | b'!' | b'#'..=b'[' | b']'..=b'~' | 0x80..=0xff) { + return false; + } + } + !escaped +} + +fn validate_qvalue(bytes: &[u8], header: &'static FieldName) -> Result<(), DecodeError> { + let valid = match bytes { + [b'0' | b'1'] => true, + [whole @ (b'0' | b'1'), b'.', fraction @ ..] if fraction.len() <= 3 => fraction + .iter() + .all(|byte| byte.is_ascii_digit() && (*whole == b'0' || *byte == b'0')), + _ => false, + }; + if valid { Ok(()) } else { Err(invalid_syntax(header)) } +} + +pub(super) fn validate_quality(bytes: &[u8], header: &'static FieldName, relaxed: bool) -> Result<(), DecodeError> { + if relaxed { + validate_qvalue_relaxed(bytes, header) + } else { + validate_qvalue(bytes, header) + } +} + +fn validate_qvalue_relaxed(bytes: &[u8], header: &'static FieldName) -> Result<(), DecodeError> { + let bytes = validate::trim_ows(bytes); + let valid = match bytes { + [b'0' | b'1'] => true, + [whole @ (b'0' | b'1'), b'.', fraction @ ..] => fraction + .iter() + .all(|byte| byte.is_ascii_digit() && (*whole == b'0' || *byte == b'0')), + [b'.', fraction @ ..] if !fraction.is_empty() => fraction.iter().all(u8::is_ascii_digit), + _ => false, + }; + if valid { Ok(()) } else { Err(invalid_syntax(header)) } +} + +#[derive(Debug)] +pub(super) struct QuotedItems<'a> { + bytes: &'a [u8], + header: &'static FieldName, + delimiter: u8, + position: usize, + start: usize, + finished: bool, + skip_empty: bool, +} + +impl<'a> QuotedItems<'a> { + pub(super) const fn comma(bytes: &'a [u8], header: &'static FieldName) -> Self { + Self::new(bytes, header, b',', true) + } + + pub(super) const fn semicolon(bytes: &'a [u8], header: &'static FieldName) -> Self { + Self::new(bytes, header, b';', false) + } + + const fn new(bytes: &'a [u8], header: &'static FieldName, delimiter: u8, skip_empty: bool) -> Self { + Self { + bytes, + header, + delimiter, + position: 0, + start: 0, + finished: false, + skip_empty, + } + } +} + +impl<'a> Iterator for QuotedItems<'a> { + type Item = Result<&'a [u8], DecodeError>; + + fn next(&mut self) -> Option { + if self.finished { + return None; + } + let remaining = &self.bytes[self.position..]; + if remaining.len() <= 1 && !remaining.first().is_some_and(|byte| *byte == self.delimiter || *byte == b'"') { + self.position = self.bytes.len(); + self.finished = true; + let item = validate::trim_ows(&self.bytes[self.start..]); + return (!self.skip_empty || !item.is_empty()).then_some(Ok(item)); + } + loop { + if self.finished { + return None; + } + let mut quoted = false; + let mut escaped = false; + while self.position < self.bytes.len() { + if !quoted { + let Some(skip) = http_headers_simd::find_either(&self.bytes[self.position..], self.delimiter, b'"') else { + self.position = self.bytes.len(); + break; + }; + self.position += skip; + } + let byte = self.bytes[self.position]; + if escaped { + escaped = false; + } else if quoted && byte == b'\\' { + escaped = true; + } else if byte == b'"' { + quoted = !quoted; + } else if !quoted && byte == self.delimiter { + let item = validate::trim_ows(&self.bytes[self.start..self.position]); + self.position += 1; + self.start = self.position; + if self.skip_empty && item.is_empty() { + continue; + } + return Some(Ok(item)); + } + self.position += 1; + } + self.finished = true; + if quoted || escaped { + return Some(Err(invalid(self.header, DecodeErrorKind::UnterminatedQuote))); + } + let item = validate::trim_ows(&self.bytes[self.start..]); + if !self.skip_empty || !item.is_empty() { + return Some(Ok(item)); + } + } + } +} + +/// Builds a decode failure. +/// +/// Every call sits on a path a well-formed field line never takes, so marking +/// it cold keeps the construction out of the straight-line decode. +#[cold] +pub(super) fn invalid(header: &'static FieldName, kind: DecodeErrorKind) -> DecodeError { + DecodeError::new(header, kind) +} + +pub(super) fn invalid_syntax(header: &'static FieldName) -> DecodeError { + invalid(header, DecodeErrorKind::InvalidSyntax) +} + +pub(super) use borrowed_list_items; +pub(super) use decode_owned_list; +pub(super) use list_header; +pub(super) use owned_list_items; + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod quoted_items_differential_tests { + use super::QuotedItems; + use crate::{DecodeErrorKind, FieldName}; + + type ScanItem<'a> = Result<&'a [u8], (DecodeErrorKind, Option)>; + + #[derive(Clone, Copy)] + enum State { + Outside, + Quoted, + Escaped, + } + + fn trim(mut bytes: &[u8]) -> &[u8] { + while bytes.first().is_some_and(|byte| matches!(*byte, b' ' | b'\t')) { + bytes = &bytes[1..]; + } + while bytes.last().is_some_and(|byte| matches!(*byte, b' ' | b'\t')) { + bytes = &bytes[..bytes.len() - 1]; + } + bytes + } + + fn scalar(bytes: &[u8], delimiter: u8) -> Vec> { + let mut items = Vec::new(); + let mut state = State::Outside; + let mut start = 0; + for (index, byte) in bytes.iter().copied().enumerate() { + state = match (state, byte) { + (State::Outside, b'"') | (State::Escaped, _) => State::Quoted, + (State::Quoted, b'\\') => State::Escaped, + (State::Quoted, b'"') => State::Outside, + (State::Outside, _) if byte == delimiter => { + let item = trim(&bytes[start..index]); + if delimiter != b',' || !item.is_empty() { + items.push(Ok(item)); + } + start = index + 1; + State::Outside + } + (state, _) => state, + }; + } + match state { + State::Outside => { + let item = trim(&bytes[start..]); + if delimiter != b',' || !item.is_empty() { + items.push(Ok(item)); + } + } + State::Quoted | State::Escaped => items.push(Err((DecodeErrorKind::UnterminatedQuote, None))), + } + items + } + + fn scanned(bytes: &[u8], delimiter: u8) -> Vec> { + let mut scanner = match delimiter { + b',' => QuotedItems::comma(bytes, &FieldName::Accept), + _ => QuotedItems::semicolon(bytes, &FieldName::Accept), + }; + let items = scanner + .by_ref() + .map(|item| item.map_err(|error| (error.kind(), error.value_index()))) + .collect(); + assert!(scanner.next().is_none()); + assert!(scanner.next().is_none()); + items + } + + #[test] + fn quote_escape_and_empty_member_results_are_explicit() { + assert_eq!( + scanned(b"br\\,gzip,\"x\\,y\",tail\\", b','), + vec![Ok(b"br\\".as_slice()), Ok(b"gzip"), Ok(b"\"x\\,y\""), Ok(b"tail\\")], + ); + assert_eq!( + scanned(b"a,\"unterminated\\", b','), + vec![Ok(b"a".as_slice()), Err((DecodeErrorKind::UnterminatedQuote, None))], + ); + assert_eq!( + scanned(b";\t;\"a;b\";;", b';'), + vec![Ok(b"".as_slice()), Ok(b""), Ok(b"\"a;b\""), Ok(b""), Ok(b"")], + ); + assert_eq!(scanned(b"\x80,\xff", b','), vec![Ok(b"\x80".as_slice()), Ok(b"\xff")]); + } + + #[test] + fn scalar_oracle_matches_empty_and_single_byte_tails() { + for delimiter in *b",;" { + assert_eq!(scanned(b"", delimiter), scalar(b"", delimiter)); + for byte in crate::test_support::byte_cases(delimiter) { + for prefix in [b"".as_slice(), b"x", b"\"quoted\""] { + let mut bytes = prefix.to_vec(); + if !prefix.is_empty() { + bytes.push(delimiter); + } + assert_eq!(scanned(&bytes, delimiter), scalar(&bytes, delimiter)); + bytes.push(byte); + assert_eq!( + scanned(&bytes, delimiter), + scalar(&bytes, delimiter), + "{bytes:?}, delimiter {delimiter}" + ); + } + } + } + } + + #[test] + fn scalar_oracle_matches_quote_and_delimiter_placements() { + const CASES: &[&[u8]] = &[ + b"", + b"|", + b"||", + b" \t|", + b"\\|tail", + b"\"a|b\"|tail", + b"\"a\\\"|b\"|tail", + b"\"a\\\\\"|tail", + b"\"a\\", + b"\"a", + b"\\\"a|tail", + b"\x80|\xff", + b"\"a\\\xff|b\"|tail", + b"\0|\r\n", + ]; + const BOUNDARIES: &[usize] = &[0, 1, 7, 8, 15, 16, 17, 31, 32, 33, 47, 48, 63, 64, 65, 95, 96, 127, 128, 129]; + for delimiter in *b",;" { + for (offset_index, offset) in [0, 1, 7, 15].into_iter().enumerate() { + for (boundary_index, &prefix) in BOUNDARIES.iter().enumerate() { + for suffix in [0, 65] { + // Miri pairs alignment and tail shape; every boundary/case keeps both tails. + if cfg!(miri) && (offset_index + boundary_index) % 2 != usize::from(suffix != 0) { + continue; + } + for case in CASES { + let mut backing = vec![b'x'; offset + prefix]; + backing.extend(case.iter().map(|byte| if *byte == b'|' { delimiter } else { *byte })); + backing.resize(backing.len() + suffix, b'y'); + let bytes = &backing[offset..]; + assert_eq!( + scanned(bytes, delimiter), + scalar(bytes, delimiter), + "{bytes:?}, delimiter {delimiter}" + ); + } + } + } + } + } + } + + #[test] + fn scalar_oracle_matches_each_byte_around_vector_boundaries() { + for delimiter in *b",;" { + for boundary in [15, 16, 17, 31, 32, 33, 63, 64, 65] { + for byte in crate::test_support::byte_cases(delimiter) { + let mut bytes = vec![b'x'; boundary]; + bytes.extend_from_slice(&[byte, delimiter, b'"', byte, b'\\', byte, b'"', delimiter, byte]); + assert_eq!( + scanned(&bytes, delimiter), + scalar(&bytes, delimiter), + "{bytes:?}, delimiter {delimiter}" + ); + } + } + } + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod weighted_token_tests { + // These tests pin the parameter-free fast path to the general segment walk. + use super::{DecodeError, validate_weighted_token, validate_weighted_token_plain, validate_weighted_token_quoted}; + use crate::{FieldName, validate}; + + /// Validates a member without the parameter-free fast path. + fn reference(bytes: &[u8], name: &'static FieldName, item_validator: fn(&[u8]) -> bool, relaxed: bool) -> Result<(), DecodeError> { + if validate_weighted_token_plain(bytes, name, item_validator, relaxed)? { + return Ok(()); + } + validate_weighted_token_quoted(bytes, name, item_validator, relaxed) + } + + fn members() -> Vec> { + let mut members: Vec> = [ + "", + " ", + " ", + "\t", + "gzip", + " gzip ", + "\tbr\t", + "en_US", + "*", + "identity", + "a;", + ";", + ";q=0.5", + "gzip;q=0.5", + "gzip ; q=0.5", + "gzip;;q=0.5", + "gzip;q=", + "\"x\"", + "\"gzip\"", + "a\\b", + "gzip;q=\"0.5\"", + "gzip;\"q\"=0.5", + "gzip,br", + "text/html", + "gzip;q=0.5;extra", + "\\", + "\"", + "\";\"", + "gzip\u{7f}", + ] + .iter() + .map(|member| member.as_bytes().to_vec()) + .collect(); + + members.extend(crate::test_support::byte_cases(b'g').map(|byte| vec![byte])); + members.extend(crate::test_support::byte_cases(b'g').map(|byte| vec![b'g', byte, b'z'])); + members + } + + #[test] + fn fast_path_matches_the_general_segment_walk() { + let validators: [fn(&[u8]) -> bool; 2] = [validate::token, |bytes| !bytes.is_empty()]; + + for member in members() { + for item_validator in validators { + for relaxed in [false, true] { + let fast = validate_weighted_token(&member, &FieldName::Accept, item_validator, relaxed); + let slow = reference(&member, &FieldName::Accept, item_validator, relaxed); + + #[cfg(miri)] + assert_eq!(fast, slow, "member {member:?}, relaxed: {relaxed}"); + #[cfg(not(miri))] + assert_eq!( + format!("{fast:?}"), + format!("{slow:?}"), + "member {:?} (relaxed: {relaxed}) must validate identically", + String::from_utf8_lossy(&member) + ); + } + } + } + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod token_list_tests { + use super::recognizes_whole_line; + + #[test] + fn only_the_negotiation_headers_are_recognized_whole() { + assert!(recognizes_whole_line(&FieldName::Accept, b"text/html;q=0.9")); + assert!(recognizes_whole_line(&FieldName::AcceptEncoding, b"gzip;q=0.9")); + assert!(recognizes_whole_line(&FieldName::AcceptLanguage, b"en-US;q=0.9")); + + for name in [&FieldName::Vary, &FieldName::Allow, &FieldName::Server] { + assert!( + !recognizes_whole_line(name, b"text/html;q=0.9"), + "{name} has no whole-line recognizer" + ); + } + } + + use super::{ + LONG_WELL_KNOWN_TOKEN_LISTS, QuotedItems, SHORT_WELL_KNOWN_TOKEN_LISTS, is_well_known_token_list, scan_token_list, + scan_token_list_general, + }; + use crate::{FieldName, validate}; + + /// Accepts exactly the lines the delimiter-aware fallback accepts. + fn delimited_scan(bytes: &[u8]) -> bool { + QuotedItems::comma(bytes, &FieldName::Allow).all(|item| match item { + Ok(item) => validate::token(item), + Err(_error) => false, + }) + } + + fn well_known_lines() -> impl Iterator { + SHORT_WELL_KNOWN_TOKEN_LISTS.iter().chain(LONG_WELL_KNOWN_TOKEN_LISTS).copied() + } + + #[test] + fn every_well_known_line_is_recognized_and_valid() { + for line in well_known_lines() { + let shown = line.escape_ascii().to_string(); + assert!(is_well_known_token_list(line), "{shown}"); + assert!(scan_token_list_general(line), "{shown}"); + assert!(delimited_scan(line), "{shown}"); + } + } + + #[cfg(test)] + #[expect( + clippy::assertions_on_result_states, + reason = "the tests classify many parser outcomes without needing their success values" + )] + mod behavior_tests { + use std::collections::hash_map::DefaultHasher; + use std::hash::{Hash, Hasher}; + + use super::super::{ + ListValues, QuotedItems, check_quoted_value, check_quoted_values, check_token_value, check_token_values, clone_checked_values, + equals_word_at_a_time, half_word, is_well_known_negotiation_line, parse_parameter, token_list_error, try_plain_items, + valid_quoted_string, validate_quality, validate_weighted_token, validate_weighted_token_quoted, word, + }; + use crate::headers::{ + Accept, AcceptEncoding, AcceptEncodingOwned, AcceptLanguage, AcceptLanguageOwned, AcceptOwned, Allow, AllowOwned, Vary, + VaryOwned, + }; + use crate::sink::{EncodedValues, FieldSink, InsertError}; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, validate}; + + #[derive(Default)] + struct Store { + name: Option<&'static FieldName>, + values: Vec, + } + + impl Store { + fn new(name: &'static FieldName, values: &[&'static str]) -> Self { + Self { + name: Some(name), + values: values.iter().copied().map(FieldValue::from_static).collect(), + } + } + } + + impl FieldSource for Store { + fn lines(&self, name: &'static FieldName) -> Option> { + (self.name == Some(name)) + .then(|| FieldLines::from_slice(name, &self.values)) + .flatten() + } + } + + impl FieldSink for Store { + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.name = Some(name); + self.values = values.into_iter().collect(); + Ok(()) + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + if self.name == Some(name) { + self.values.extend(values); + } else { + self.name = Some(name); + self.values = values.into_iter().collect(); + } + Ok(()) + } + + fn remove_values(&mut self, name: &'static FieldName) { + if self.name == Some(name) { + self.name = None; + self.values.clear(); + } + } + } + + fn validate_token(bytes: &[u8]) -> Result<(), DecodeError> { + if validate::token(bytes) { + Ok(()) + } else { + Err(DecodeError::new(&FieldName::Accept, DecodeErrorKind::InvalidToken)) + } + } + + #[test] + fn generated_list_headers_preserve_single_and_repeated_lines() { + let single = AcceptOwned::try_from(String::from("text/plain;level=\"one\", text/html")).expect("quoted media parameters"); + assert_eq!( + single.items().collect::>(), + [b"text/plain;level=\"one\"".as_slice(), b"text/html"] + ); + assert_eq!(single.values().len(), 1); + assert!(format!("{single:?}").contains("value_count: 1")); + assert_eq!(single, single.clone()); + let mut hasher = DefaultHasher::new(); + single.hash(&mut hasher); + assert_ne!(hasher.finish(), 0); + + let source = Store::new(&FieldName::Accept, &["text/plain;level=\"one\"", "application/json", "text/html"]); + let view = Accept::view(&source).expect("valid repeated list").expect("header present"); + assert_eq!(view.items().count(), 3); + assert_eq!(view.values().count(), 3); + assert!(format!("{view:?}").contains("value_count: 3")); + + let owned = Accept::owned(&source).expect("valid repeated list").expect("header present"); + assert_eq!(owned.values().len(), 3); + let mut sink = Store::default(); + Accept::insert(&mut sink, owned).expect("insert repeated values"); + assert_eq!(sink.values.len(), 3); + + let mut sink = Store::default(); + Accept::insert(&mut sink, single).expect("insert one value"); + assert_eq!(sink.values.len(), 1); + sink.remove_values(&FieldName::Accept); + assert!(sink.values.is_empty()); + assert_eq!(sink.name, None); + assert!(Accept::view(&Store::default()).expect("absent source").is_none()); + assert!(Accept::owned(&Store::default()).expect("absent source").is_none()); + } + + #[test] + fn generated_token_lists_validate_views_owned_values_and_conversions() { + let source = Store::new(&FieldName::Allow, &["GET, HEAD", "POST"]); + let view = Allow::view(&source).expect("valid methods").expect("header present"); + assert_eq!(view.items().collect::>(), [b"GET".as_slice(), b"HEAD", b"POST"]); + let owned = Allow::owned(&source).expect("valid methods").expect("header present"); + assert_eq!(owned.items().count(), 3); + + let mut sink = Store::default(); + Allow::insert(&mut sink, owned).expect("insert token list"); + assert_eq!(sink.values.len(), 2); + + assert_eq!( + AllowOwned::try_from(String::from("bad method")) + .expect_err("spaces are not tokens") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + VaryOwned::try_from(String::from("line\nbreak")) + .expect_err("invalid field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let relaxed = Store::new(&FieldName::Accept, &["text/plain; q = .75"]); + assert!(Accept::view_with(&relaxed, DecodeMode::Relaxed).expect("relaxed quality").is_some()); + assert!( + Accept::owned_with(&relaxed, DecodeMode::Relaxed) + .expect("relaxed quality") + .is_some() + ); + } + + #[test] + fn generated_owned_error_and_borrowed_conversion_paths_are_preserved() { + for source in [ + Store::new(&FieldName::Accept, &["text/plain", "\"unterminated"]), + Store::new(&FieldName::Accept, &["text/plain", "invalid"]), + ] { + assert!(Accept::owned(&source).is_err()); + } + let invalid_allow = Store::new(&FieldName::Allow, &["GET", "bad method"]); + assert!(Allow::owned(&invalid_allow).is_err()); + + assert!(AcceptOwned::try_from("text/plain").is_ok()); + assert!(AcceptEncodingOwned::try_from("gzip").is_ok()); + assert!(AcceptLanguageOwned::try_from("en").is_ok()); + assert!(AllowOwned::try_from("GET").is_ok()); + assert!(VaryOwned::try_from("accept").is_ok()); + + assert!(AcceptOwned::try_from("line\nbreak").is_err()); + assert!(AcceptEncodingOwned::try_from("line\nbreak").is_err()); + assert!(AcceptLanguageOwned::try_from("line\nbreak").is_err()); + assert!(AllowOwned::try_from("line\nbreak").is_err()); + assert!(VaryOwned::try_from("line\nbreak").is_err()); + assert!(AcceptOwned::try_from(String::from("line\nbreak")).is_err()); + assert!(AcceptEncodingOwned::try_from(String::from("line\nbreak")).is_err()); + assert!(AcceptLanguageOwned::try_from(String::from("line\nbreak")).is_err()); + assert!(AllowOwned::try_from(String::from("line\nbreak")).is_err()); + } + + #[test] + fn every_generated_list_type_exercises_its_owned_and_view_apis() { + macro_rules! exercise { + ($header:ty, $owned:ty, $name:expr, $line:expr, $items:expr) => {{ + let source = Store::new($name, $line); + let view = <$header>::view(&source).expect("valid view").expect("present view"); + assert_eq!(view.items().collect::>(), $items); + assert_eq!(view.values().count(), $line.len()); + assert!(!format!("{view:?}").is_empty()); + + let owned = <$header>::owned(&source) + .expect("valid owned value") + .expect("present owned value"); + assert_eq!(owned.items().collect::>(), $items); + assert_eq!(owned.values().len(), $line.len()); + assert_eq!(owned, owned.clone()); + assert!(!format!("{owned:?}").is_empty()); + let mut hasher = DefaultHasher::new(); + owned.hash(&mut hasher); + assert_ne!(hasher.finish(), 0); + + assert!(<$owned>::try_from(String::from($line[0])).is_ok()); + let mut sink = Store::default(); + <$header>::insert(&mut sink, owned).expect("insert generated list"); + assert_eq!(sink.values.len(), $line.len()); + }}; + } + + exercise!( + AcceptEncoding, + AcceptEncodingOwned, + &FieldName::AcceptEncoding, + &["gzip", "br"], + [b"gzip".as_slice(), b"br".as_slice()] + ); + exercise!( + AcceptLanguage, + AcceptLanguageOwned, + &FieldName::AcceptLanguage, + &["en", "fr"], + [b"en".as_slice(), b"fr".as_slice()] + ); + exercise!( + Allow, + AllowOwned, + &FieldName::Allow, + &["TRACE", "PATCH"], + [b"TRACE".as_slice(), b"PATCH".as_slice()] + ); + exercise!( + Vary, + VaryOwned, + &FieldName::Vary, + &["accept", "origin"], + [b"accept".as_slice(), b"origin".as_slice()] + ); + } + + #[test] + fn checked_copy_and_quoted_line_validation_cover_line_boundaries() { + let values = [ + FieldValue::from_static("alpha"), + FieldValue::from_static("beta"), + FieldValue::from_static("gamma"), + ]; + let borrowed = FieldLines::from_slice(&FieldName::Accept, &values).expect("nonempty values"); + let mut indices = Vec::new(); + let copied = clone_checked_values(&borrowed, |index, value| { + indices.push(index); + validate_token(value.as_bytes()) + }) + .expect("all lines are tokens"); + assert_eq!(indices, [0, 1, 2]); + assert!(matches!(copied, ListValues::Many(ref lines) if lines.len() == 3)); + + let one = [FieldValue::from_static("alpha")]; + let borrowed = FieldLines::from_slice(&FieldName::Accept, &one).expect("one value"); + assert!(matches!( + clone_checked_values(&borrowed, |_index, _value| Ok(())).expect("valid line"), + ListValues::One(_) + )); + + assert_eq!( + check_quoted_values(&FieldName::Accept, &borrowed, validate_token) + .expect("one quoted line") + .expect("single line") + .as_bytes(), + b"alpha" + ); + let invalid_first = FieldLines::single(&FieldName::Accept, b"bad token"); + assert!(check_quoted_values(&FieldName::Accept, &invalid_first, validate_token).is_err()); + assert!( + check_quoted_values( + &FieldName::Accept, + &FieldLines::from_slice(&FieldName::Accept, &values).expect("three values"), + validate_token, + ) + .expect("three valid lines") + .is_none() + ); + let invalid_third = [ + FieldValue::from_static("alpha"), + FieldValue::from_static("beta"), + FieldValue::from_static("bad token"), + ]; + assert!( + check_quoted_values( + &FieldName::Accept, + &FieldLines::from_slice(&FieldName::Accept, &invalid_third).expect("three values"), + validate_token, + ) + .is_err() + ); + assert_eq!( + check_quoted_value( + &FieldName::Accept, + FieldValue::from_static("\"unterminated").as_field_value_ref(), + validate_token, + ) + .expect_err("unterminated quote") + .kind(), + DecodeErrorKind::UnterminatedQuote + ); + + let borrowed_all = FieldLines::from_slice(&FieldName::Accept, &values).expect("three values"); + for failing_index in 0..3 { + let error = clone_checked_values(&borrowed_all, |index, _value| { + if index == failing_index { + Err(DecodeError::new(&FieldName::Accept, DecodeErrorKind::InvalidSyntax)) + } else { + Ok(()) + } + }) + .expect_err("the selected line fails validation"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + } + + let invalid_lines = [FieldValue::from_static("alpha"), FieldValue::from_static("\"unterminated")]; + let invalid = FieldLines::from_slice(&FieldName::Accept, &invalid_lines).expect("two field lines"); + assert_eq!( + check_quoted_values(&FieldName::Accept, &invalid, validate_token) + .expect_err("second line has an unterminated quote") + .kind(), + DecodeErrorKind::UnterminatedQuote + ); + } + + #[test] + fn token_list_private_paths_cover_unusual_and_invalid_lines() { + assert!( + check_quoted_value( + &FieldName::Accept, + FieldValue::from_static("\"quoted\"").as_field_value_ref(), + validate_token, + ) + .is_err() + ); + let token_lines = [FieldValue::from_static("GET"), FieldValue::from_static("bad method")]; + let tokens = FieldLines::from_slice(&FieldName::Allow, &token_lines).expect("two token lines"); + assert_eq!( + check_token_values(&FieldName::Allow, &tokens, validate_token) + .expect_err("second line is invalid") + .kind(), + DecodeErrorKind::InvalidToken + ); + + let unusual = FieldLines::single(&FieldName::Allow, b"TRACE"); + assert_eq!( + check_token_values(&FieldName::Allow, &unusual, validate_token) + .expect("valid uncommon token list") + .expect("single line") + .as_bytes(), + b"TRACE" + ); + let unusual_lines = [FieldValue::from_static("TRACE"), FieldValue::from_static("PATCH")]; + let unusual_repeated = FieldLines::from_slice(&FieldName::Allow, &unusual_lines).expect("two lines"); + assert!( + check_token_values(&FieldName::Allow, &unusual_repeated, validate_token) + .expect("valid repeated uncommon token list") + .is_none() + ); + let invalid_unusual = FieldLines::single(&FieldName::Allow, b"bad method"); + assert!(check_token_values(&FieldName::Allow, &invalid_unusual, validate_token).is_err()); + assert_eq!( + check_token_value( + &FieldName::Allow, + FieldValue::from_static("bad method").as_field_value_ref(), + validate_token, + ) + .expect_err("single invalid token line") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert!( + check_token_value( + &FieldName::Allow, + FieldValue::from_static("TRACE").as_field_value_ref(), + validate_token, + ) + .is_ok() + ); + } + + #[test] + fn weighted_tokens_parameters_and_qvalues_cover_strict_and_relaxed_grammar() { + for valid in [b"gzip".as_slice(), b"gzip;q=0.5"] { + assert!( + validate_weighted_token(valid, &FieldName::AcceptEncoding, validate::token, false).is_ok(), + "{valid:?}" + ); + } + assert!(validate_weighted_token(b"gzip; q = .1234", &FieldName::AcceptEncoding, validate::token, true,).is_ok()); + for invalid in [ + b"".as_slice(), + b"bad token", + b"gzip;level=1", + b"gzip;q=0.5;level=1", + b"gzip;q=\"unterminated", + ] { + assert!( + validate_weighted_token(invalid, &FieldName::AcceptEncoding, validate::token, false,).is_err(), + "{invalid:?}" + ); + } + + assert_eq!( + parse_parameter(b" name = \"a b\" ", false, &FieldName::Accept).expect("quoted parameter"), + (b"name".as_slice(), Some(b"\"a b\"".as_slice()), false) + ); + assert_eq!( + parse_parameter(b"flag", true, &FieldName::Accept).expect("optional value"), + (b"flag".as_slice(), None, true) + ); + assert!(parse_parameter(b"flag", false, &FieldName::Accept).is_err()); + assert!(parse_parameter(b"bad name=x", false, &FieldName::Accept).is_err()); + assert!(parse_parameter(b"name=\"bad\\\x01\"", false, &FieldName::Accept).is_err()); + + for valid in [b"\"\"".as_slice(), b"\"a b\"", b"\"a\\\"b\"", b"\"\xff\""] { + assert!(valid_quoted_string(valid), "{valid:?}"); + } + assert!(valid_quoted_string(b"\"a\\\xff\"")); + for invalid in [ + b"".as_slice(), + b"token", + b"\"unterminated", + b"\"bad\\\x01\"", + b"\"bad\x01\"", + b"\"bad\\\"", + ] { + assert!(!valid_quoted_string(invalid), "{invalid:?}"); + } + + for valid in [b"0".as_slice(), b"1", b"0.123", b"1.000"] { + assert!(validate_quality(valid, &FieldName::Accept, false).is_ok()); + } + for invalid in [b"2".as_slice(), b"1.001", b"0.1234", b".5"] { + assert!(validate_quality(invalid, &FieldName::Accept, false).is_err()); + } + for valid in [b" .5 ".as_slice(), b"0.1234", b"1.0000"] { + assert!(validate_quality(valid, &FieldName::Accept, true).is_ok()); + } + assert!(validate_quality(b"1.0001", &FieldName::Accept, true).is_err()); + assert!(validate_quality(b"2", &FieldName::Accept, true).is_err()); + assert!(validate_quality(b"1", &FieldName::Accept, true).is_ok()); + assert!(validate_weighted_token(b"gzip;bad", &FieldName::AcceptEncoding, validate::token, false,).is_err()); + } + + #[test] + fn direct_quoted_weight_and_plain_list_paths_cover_fallbacks() { + validate_weighted_token_quoted(b"gzip;q=0.5", &FieldName::AcceptEncoding, validate::token, false) + .expect("direct quoted-path validation"); + for invalid in [ + b"bad token;q=0.5".as_slice(), + b"gzip;level=1", + b"gzip;bad", + b"gzip;q=2", + b"gzip;q=0.5;extra=1", + b"gzip;q=\"unterminated", + b"gzip;q=0.5;\"unterminated", + b"\"unterminated", + ] { + assert!( + validate_weighted_token_quoted(invalid, &FieldName::AcceptEncoding, validate::token, false,).is_err(), + "{invalid:?}" + ); + } + + assert!(is_well_known_negotiation_line(&FieldName::AcceptEncoding, b"gzip")); + assert!(is_well_known_negotiation_line(&FieldName::AcceptLanguage, b"en-US")); + // Multi-member lines vary too much between clients to be worth a + // literal, so they reach the member grammar instead. + for line in [&b"gzip, br"[..], b"gzip, br;q=0.8"] { + assert!(!is_well_known_negotiation_line(&FieldName::AcceptEncoding, line)); + } + assert!(!is_well_known_negotiation_line(&FieldName::AcceptLanguage, b"en, en-US;q=0.8")); + assert!(!is_well_known_negotiation_line(&FieldName::Vary, b"accept")); + + let mut items = Vec::new(); + assert!( + try_plain_items(b", alpha, , beta", b',', true, |item| { + items.push(item.to_vec()); + Ok(()) + }) + .expect("plain list") + ); + assert_eq!(items, [b"alpha".to_vec(), b"beta".to_vec()]); + assert!(!try_plain_items(b"alpha,\"beta\"", b',', true, |_item| Ok(())).expect("quoted fallback")); + assert!(try_plain_items(b"alpha,", b',', true, validate_token).expect("trailing empty item")); + assert!(try_plain_items(b"bad method,GET", b',', true, validate_token).is_err()); + validate_weighted_token_quoted(b"gzip", &FieldName::AcceptEncoding, validate::token, false) + .expect("quoted parser without a weight"); + } + + #[test] + fn quoted_iterator_and_low_level_word_helpers_cover_edge_shapes() { + let items = QuotedItems::comma(b", alpha, \"beta,gamma\",,", &FieldName::Accept) + .collect::, _>>() + .expect("quoted commas stay within an item"); + assert_eq!(items, [b"alpha".as_slice(), b"\"beta,gamma\""]); + + let items = QuotedItems::semicolon(b"alpha;\"beta\\\"gamma\";", &FieldName::Accept) + .collect::, _>>() + .expect("escaped quote"); + assert_eq!(items, [b"alpha".as_slice(), b"\"beta\\\"gamma\"", b""]); + assert_eq!( + QuotedItems::comma(b"\"unfinished", &FieldName::Accept) + .next() + .expect("one error") + .expect_err("quote is unfinished") + .kind(), + DecodeErrorKind::UnterminatedQuote + ); + + assert_eq!(word(b"abcdefgh", 0), u64::from_le_bytes(*b"abcdefgh")); + assert_eq!(word(b"short", 0), 0); + assert_eq!(half_word(b"ab", 0), u16::from_le_bytes(*b"ab")); + assert_eq!(half_word(b"a", 0), 0); + assert!(equals_word_at_a_time(b"short", b"short")); + assert!(!equals_word_at_a_time(b"short", b"longer")); + assert!(equals_word_at_a_time(b"123456789", b"123456789")); + assert!(!equals_word_at_a_time(b"123456789", b"123456780")); + assert!(equals_word_at_a_time(b"1234567890", b"1234567890")); + assert!(equals_word_at_a_time(b"12345678901", b"12345678901")); + + assert_eq!( + token_list_error(&FieldName::Allow, b"alpha,\"unfinished", validate_token, Some(2),).kind(), + DecodeErrorKind::UnterminatedQuote + ); + assert_eq!( + token_list_error(&FieldName::Allow, b"\"unfinished", validate_token, None,).kind(), + DecodeErrorKind::UnterminatedQuote + ); + assert_eq!( + token_list_error(&FieldName::Allow, b"bad token", validate_token, None).kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + token_list_error(&FieldName::Allow, b"GET", validate_token, None).kind(), + DecodeErrorKind::InvalidToken + ); + } + } + + #[test] + fn the_well_known_table_never_accepts_an_invalid_line() { + let mut mutated = Vec::new(); + for line in well_known_lines() { + for position in 0..line.len() { + for replacement in crate::test_support::substitution_bytes(line[position], position, line.len()) { + mutated.clear(); + mutated.extend_from_slice(line); + mutated[position] = replacement; + assert!( + !is_well_known_token_list(&mutated) || scan_token_list_general(&mutated), + "{:?}", + mutated.escape_ascii().to_string() + ); + } + } + for position in 0..line.len() { + mutated.clear(); + mutated.extend_from_slice(line); + mutated.remove(position); + assert!( + !is_well_known_token_list(&mutated) || scan_token_list_general(&mutated), + "{:?}", + mutated.escape_ascii().to_string() + ); + } + } + } + + #[test] + fn fast_scan_agrees_with_the_delimited_scan() { + let alphabet: &[u8] = b"a,; \t\"\\!\x00\x80"; + let maximum_length = if cfg!(miri) { 2 } else { 4 }; + for length in 0..=maximum_length { + let mut counters = vec![0_usize; length]; + let mut input = vec![0_u8; length]; + loop { + for (slot, counter) in input.iter_mut().zip(counters.iter()) { + *slot = alphabet[*counter]; + } + assert_eq!( + scan_token_list(&input), + delimited_scan(&input), + "{:?}", + input.escape_ascii().to_string() + ); + let mut position = 0; + loop { + if position == counters.len() { + break; + } + counters[position] += 1; + if counters[position] < alphabet.len() { + break; + } + counters[position] = 0; + position += 1; + } + if position == counters.len() { + break; + } + } + } + + #[cfg(miri)] + for value in [ + b"a,a".as_slice(), + b"a,,", + b",a,", + b" a ", + b"a,\t", + b"a\\a", + b"a;!", + b"a\0a", + b"\"a\"", + b"a,\"", + b"a\"a", + b"a,\x80", + b"a, a", + b"a,,a", + b"a ,a", + b"a,\\a", + ] { + assert_eq!(scan_token_list(value), delimited_scan(value), "{value:?}"); + } + + for value in [ + &b"GET, POST"[..], + b"accept-encoding, origin", + b"a-very-long-token-name, another-token, third", + b"abcdefgh", + b"abcdefg h", + b"abcdefgh,", + b"abcdefghi", + b"abcdefgh ,i", + b"abc\x7fdefgh", + b"abcdefg\x7f", + ] { + assert_eq!( + scan_token_list(value), + delimited_scan(value), + "{:?}", + value.escape_ascii().to_string() + ); + } + } + + #[test] + fn fast_scan_agrees_on_lines_long_enough_to_vectorize() { + let base: &[u8] = b"alpha, beta, gamma, delta, epsilon, zeta, eta, theta, iota, kappa, mu"; + let alphabet: &[u8] = b"a,; \t\"\\!\x00\x80\x7f"; + let mut mutated = Vec::new(); + for length in 0..=base.len() { + if cfg!(miri) && length != base.len() && !matches!(length, 0..=2 | 7..=9 | 15..=17 | 31..=33 | 47..=49 | 63..=65) { + continue; + } + let line = &base[..length]; + assert_eq!(scan_token_list(line), delimited_scan(line), "{:?}", line.escape_ascii().to_string()); + for position in 0..length { + for replacement in alphabet { + mutated.clear(); + mutated.extend_from_slice(line); + mutated[position] = *replacement; + assert_eq!( + scan_token_list(&mutated), + delimited_scan(&mutated), + "{:?}", + mutated.escape_ascii().to_string() + ); + } + } + } + + let mut state = 0x2545_f491_4f6c_dd1d_u64; + let cases = if cfg!(miri) { 256 } else { 200_000 }; + for _ in 0..cases { + state = state + .wrapping_mul(6_364_136_223_846_793_005) + .wrapping_add(1_442_695_040_888_963_407); + let length = usize::try_from(state >> 57).expect("a 7-bit length fits a usize"); + mutated.clear(); + for _ in 0..length { + state = state + .wrapping_mul(6_364_136_223_846_793_005) + .wrapping_add(1_442_695_040_888_963_407); + let pick = usize::try_from(state >> 40).expect("a 24-bit index fits a usize"); + mutated.push(alphabet[pick % alphabet.len()]); + } + assert_eq!( + scan_token_list(&mutated), + delimited_scan(&mutated), + "{:?}", + mutated.escape_ascii().to_string() + ); + } + } +} diff --git a/crates/http_headers/src/headers/negotiation/vary.rs b/crates/http_headers/src/headers/negotiation/vary.rs new file mode 100644 index 000000000..4154775e9 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/vary.rs @@ -0,0 +1,173 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::shared::{ListValues, check_token_value, check_token_values, invalid}; +use super::{FieldNameView, VaryEntryView}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Owned value for the `Vary` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 12.5.5]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::VaryOwned::try_from("accept, origin")?; +/// assert_eq!(value.items().count(), 2); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Vary: Accept-Encoding, Accept-Language` names selection fields, while +/// `Vary: *` means that other aspects of the request influenced selection. +/// +/// [RFC 9110 section 12.5.5]: https://www.rfc-editor.org/rfc/rfc9110#section-12.5.5 +pub struct VaryOwned { + values: ListValues, +} + +/// Borrowed value for the `Vary` header. +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{Vary, VaryView}; +/// use http_headers::source::{FieldLines, FieldSource}; +/// use http_headers::{Field, FieldName}; +/// +/// struct Source; +/// +/// impl FieldSource for Source { +/// fn lines(&self, name: &'static FieldName) -> Option> { +/// (name == &FieldName::Vary).then(|| FieldLines::single(name, b"Accept-Encoding, Origin")) +/// } +/// } +/// +/// let value: VaryView<'_> = Vary::view(&Source)?.expect("header is present"); +/// let items = value.items().collect::>(); +/// assert_eq!(items, [&b"Accept-Encoding"[..], &b"Origin"[..]]); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct VaryView<'a> { + values: FieldLines<'a>, +} + +super::shared::list_header!( + Vary, + VaryOwned, + VaryView, + "Vary", + "Defined by [RFC 9110 section 12.5.5](https://www.rfc-editor.org/rfc/rfc9110#section-12.5.5).", + &FieldName::Vary, + validate_vary_item, + validate_vary_item, + check_token_values, + check_token_value, + token +); + +impl VaryOwned { + /// Iterates wildcard or validated-name members in wire order. + /// + /// Name comparisons are case-insensitive; original spelling is retained. + /// Empty list members are ignored, but wildcard members are never hidden. + #[inline] + pub fn entries(&self) -> impl Iterator> { + self.items().map(VaryEntryView::from_validated) + } + + /// Reports whether any member is `*`, including mixed wildcard/name lists. + /// + /// This allocation-free query scans until the first wildcard. + /// + /// # Examples + /// + /// ``` + /// use http_headers::headers::VaryOwned; + /// + /// assert!(VaryOwned::try_from("Origin, *")?.contains_wildcard()); + /// assert!(!VaryOwned::try_from("Origin, X-*")?.contains_wildcard()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[must_use] + #[inline] + pub fn contains_wildcard(&self) -> bool { + self.items().any(|item| item == b"*") + } + + /// Constructs the wildcard field value. + #[must_use] + pub fn wildcard() -> Self { + Self { + values: ListValues::One(FieldValue::from_static("*")), + } + } + + /// Constructs one field line from validated names. + /// + /// Order, duplicates, and spelling are preserved. A name equal to `*` + /// has wildcard semantics. An empty iterator produces an empty field. + #[must_use] + pub fn from_field_names<'a>(names: impl IntoIterator>) -> Self { + Self::from_entries(names.into_iter().map(VaryEntryView::from_field_name)) + } + + /// Constructs one field line from typed wildcard and name members. + /// + /// Mixed wildcard/name lists are retained without sorting or deduplication. + #[must_use] + pub fn from_entries<'a>(entries: impl IntoIterator>) -> Self { + let mut wire = String::new(); + for entry in entries { + if !wire.is_empty() { + wire.push_str(", "); + } + wire.push_str(entry.as_str()); + } + Self { + values: ListValues::One(FieldValue::from_validated_owned_bytes(wire.into_bytes(), false)), + } + } +} + +impl<'a> VaryView<'a> { + /// Iterates wildcard or validated-name members without allocating. + #[inline] + pub fn entries(&self) -> impl Iterator> + '_ { + self.items().map(VaryEntryView::from_validated) + } + + /// Reports whether any member is `*`, scanning until the first wildcard. + #[must_use] + #[inline] + pub fn contains_wildcard(&self) -> bool { + self.items().any(|item| item == b"*") + } +} + +fn validate_vary_item(bytes: &[u8]) -> Result<(), DecodeError> { + if validate::token(bytes) { + Ok(()) + } else { + Err(invalid(&FieldName::Vary, DecodeErrorKind::InvalidToken)) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::validate_vary_item; + use crate::DecodeErrorKind; + + #[test] + fn vary_items_are_bare_field_name_tokens_or_wildcard() { + validate_vary_item(b"*").expect("wildcard token"); + validate_vary_item(b"x-selection-input").expect("extension field name"); + assert_eq!( + validate_vary_item(b"bad field").expect_err("spaces are not token bytes").kind(), + DecodeErrorKind::InvalidToken + ); + } +} diff --git a/crates/http_headers/src/headers/negotiation/vary_entry_view.rs b/crates/http_headers/src/headers/negotiation/vary_entry_view.rs new file mode 100644 index 000000000..bb4318c1a --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/vary_entry_view.rs @@ -0,0 +1,70 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; + +use super::FieldNameView; + +/// A Vary member: either a wildcard or a validated field name. +/// +/// Wildcards have no field name. Named members compare and hash +/// case-insensitively while retaining their original spelling. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct VaryEntryView<'a>(Option>); + +impl<'a> VaryEntryView<'a> { + /// The wildcard selection member. + pub const WILDCARD: Self = Self(None); + + /// Constructs a member from a validated name. + /// + /// The name `*` becomes a wildcard, never a literal selection field. + #[must_use] + #[inline] + pub fn from_field_name(name: FieldNameView<'a>) -> Self { + if name.as_bytes() == b"*" { + Self::WILDCARD + } else { + Self(Some(name)) + } + } + + #[inline] + pub(super) fn from_validated(bytes: &'a [u8]) -> Self { + if bytes == b"*" { + Self::WILDCARD + } else { + Self(Some(FieldNameView::from_validated(bytes))) + } + } + + /// Returns whether this member is the wildcard. + #[must_use] + #[inline] + pub const fn is_wildcard(self) -> bool { + self.0.is_none() + } + + /// Returns the validated field name, or `None` for the wildcard. + #[must_use] + #[inline] + pub const fn field_name(self) -> Option> { + self.0 + } + + /// Returns the member's original spelling. + #[must_use] + #[inline] + pub fn as_str(self) -> &'a str { + match self.0 { + Some(name) => name.as_str(), + None => "*", + } + } +} + +impl fmt::Display for VaryEntryView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(self.as_str()) + } +} diff --git a/crates/http_headers/src/headers/negotiation/weighted_token_scan.rs b/crates/http_headers/src/headers/negotiation/weighted_token_scan.rs new file mode 100644 index 000000000..c6b82ff29 --- /dev/null +++ b/crates/http_headers/src/headers/negotiation/weighted_token_scan.rs @@ -0,0 +1,394 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Whole-line recognizer for the weighted-token negotiation field lines. +//! +//! `Accept-Encoding` and `Accept-Language` share a member grammar: an item +//! followed by at most one quality parameter. Only the item differs, so one +//! machine serves both and the item shape selects which table is built. + +/// The item a member starts with. +#[derive(Clone, Copy)] +pub(super) enum Item { + /// Any RFC 9110 token, which is what a content coding is. + Token, + /// `*`, or an alphabetic primary tag and alphanumeric subtags of one to + /// eight bytes each. + LanguageRange, +} + +const CLASS_OTHER: u8 = 0; +const CLASS_TCHAR: u8 = 1; +const CLASS_Q: u8 = 2; +const CLASS_ALPHA: u8 = 3; +const CLASS_ZERO: u8 = 4; +const CLASS_ONE: u8 = 5; +const CLASS_DIGIT: u8 = 6; +const CLASS_DOT: u8 = 7; +const CLASS_STAR: u8 = 8; +const CLASS_HYPHEN: u8 = 9; +const CLASS_SEMICOLON: u8 = 10; +const CLASS_COMMA: u8 = 11; +const CLASS_EQUALS: u8 = 12; +const CLASS_OWS: u8 = 13; + +const CLASSES: usize = 16; +const STATES: usize = 64; + +/// The line left the recognized subset. The state absorbs every later byte. +const REJECTED: u8 = 0; +/// A member is due: the line just started or a comma just closed one. +const MEMBER_DUE: u8 = 1; +/// Whitespace closed a complete item, so only `;`, `,`, or the end may follow. +const MEMBER_OWS: u8 = 2; +/// A semicolon opened the quality parameter. +const PARAM_DUE: u8 = 3; +/// The parameter name `q` has been read. +const NAME_Q: u8 = 4; +/// An equals sign closed the `q` name, so a quality value is due. +const QUALITY_DUE: u8 = 5; +/// The quality value read so far is `0`. +const QUALITY_ZERO: u8 = 6; +/// The quality value read so far is `1`. +const QUALITY_ONE: u8 = 7; +/// `0.` with no fraction digit yet. +const QUALITY_ZERO_DOT: u8 = 8; +/// `0.` followed by one fraction digit. +const QUALITY_ZERO_ONE_DIGIT: u8 = 9; +/// `0.` followed by two fraction digits. +const QUALITY_ZERO_TWO_DIGITS: u8 = 10; +/// `0.` followed by three fraction digits, the most the grammar allows. +const QUALITY_ZERO_THREE_DIGITS: u8 = 11; +/// `1.` with no fraction zero yet. +const QUALITY_ONE_DOT: u8 = 12; +/// `1.` followed by one fraction zero. +const QUALITY_ONE_ONE_DIGIT: u8 = 13; +/// `1.` followed by two fraction zeros. +const QUALITY_ONE_TWO_DIGITS: u8 = 14; +/// `1.` followed by three fraction zeros, the most the grammar allows. +const QUALITY_ONE_THREE_DIGITS: u8 = 15; +/// Whitespace closed a member that already carries its quality. +const MEMBER_OWS_WEIGHED: u8 = 16; + +/// A token item is in progress. +const TOKEN: u8 = 17; +/// The language range read so far is exactly `*`. +const LANGUAGE_STAR: u8 = 17; +/// The primary tag is one byte long. Seven more states follow it. +const PRIMARY_ONE: u8 = 18; +/// The primary tag has reached the eight-byte limit. +const PRIMARY_LAST: u8 = 25; +/// A hyphen closed a tag, so a subtag is due. +const SUBTAG_DUE: u8 = 26; +/// A subtag is one byte long. Seven more states follow it. +const SUBTAG_ONE: u8 = 27; +/// A subtag has reached the eight-byte limit. +const SUBTAG_LAST: u8 = 34; + +/// The states that end a line inside the recognized subset, other than the +/// item states, which each shape contributes itself. +const SHARED_ACCEPTING: u64 = (1 << MEMBER_DUE) + | (1 << MEMBER_OWS) + | (1 << MEMBER_OWS_WEIGHED) + | (1 << QUALITY_ZERO) + | (1 << QUALITY_ONE) + | (1 << QUALITY_ZERO_DOT) + | (1 << QUALITY_ZERO_ONE_DIGIT) + | (1 << QUALITY_ZERO_TWO_DIGITS) + | (1 << QUALITY_ZERO_THREE_DIGITS) + | (1 << QUALITY_ONE_DOT) + | (1 << QUALITY_ONE_ONE_DIGIT) + | (1 << QUALITY_ONE_TWO_DIGITS) + | (1 << QUALITY_ONE_THREE_DIGITS); + +/// Maps each byte to its role inside a weighted-token field line. +static CLASS: [u8; 256] = class_table(); + +/// Maps a row and a byte class to the next row. +/// +/// A row is a state already multiplied by [`CLASSES`], so a step indexes the +/// table with `row | class` and stores the entry back unchanged. Keeping the +/// multiply in the table removes a shift from every byte of the scan. +static TOKEN_TRANSITION: [u16; STATES * CLASSES] = transition_table(Item::Token); +static LANGUAGE_TRANSITION: [u16; STATES * CLASSES] = transition_table(Item::LanguageRange); + +/// The row a state occupies in a transition table. +const fn row(state: u8) -> u16 { + (state as u16) << 4 +} + +const TOKEN_ACCEPTING: u64 = SHARED_ACCEPTING | (1 << TOKEN); +const LANGUAGE_ACCEPTING: u64 = + SHARED_ACCEPTING | (1 << LANGUAGE_STAR) | range_mask(PRIMARY_ONE, PRIMARY_LAST) | range_mask(SUBTAG_ONE, SUBTAG_LAST); + +/// Builds a mask holding every state from `first` to `last` inclusive. +const fn range_mask(first: u8, last: u8) -> u64 { + let mut mask = 0_u64; + let mut state = first; + while state <= last { + mask |= 1 << state; + state += 1; + } + mask +} + +const fn class_table() -> [u8; 256] { + let mut table = [CLASS_OTHER; 256]; + let mut byte = 0_usize; + while byte < 256 { + #[expect(clippy::cast_possible_truncation, reason = "the loop bound keeps the index inside a byte")] + let value = byte as u8; + table[byte] = match value { + b'q' | b'Q' => CLASS_Q, + b'0' => CLASS_ZERO, + b'1' => CLASS_ONE, + b'2'..=b'9' => CLASS_DIGIT, + b'A'..=b'Z' | b'a'..=b'z' => CLASS_ALPHA, + b'.' => CLASS_DOT, + b'*' => CLASS_STAR, + b'-' => CLASS_HYPHEN, + b';' => CLASS_SEMICOLON, + b',' => CLASS_COMMA, + b'=' => CLASS_EQUALS, + b' ' | b'\t' => CLASS_OWS, + b'!' | b'#' | b'$' | b'%' | b'&' | b'\'' | b'+' | b'^' | b'_' | b'`' | b'|' | b'~' => CLASS_TCHAR, + _ => CLASS_OTHER, + }; + byte += 1; + } + table +} + +/// Returns whether a class is one of the `tchar` classes. +const fn is_token_class(class: u8) -> bool { + matches!( + class, + CLASS_TCHAR | CLASS_Q | CLASS_ALPHA | CLASS_ZERO | CLASS_ONE | CLASS_DIGIT | CLASS_DOT | CLASS_STAR | CLASS_HYPHEN + ) +} + +/// Returns whether a class is an ASCII letter. +const fn is_alpha_class(class: u8) -> bool { + matches!(class, CLASS_Q | CLASS_ALPHA) +} + +/// Returns whether a class is an ASCII letter or digit. +const fn is_alphanumeric_class(class: u8) -> bool { + matches!(class, CLASS_Q | CLASS_ALPHA | CLASS_ZERO | CLASS_ONE | CLASS_DIGIT) +} + +/// Closes a complete item, which may carry a quality parameter next. +const fn item_close(class: u8) -> u8 { + match class { + CLASS_OWS => MEMBER_OWS, + CLASS_SEMICOLON => PARAM_DUE, + CLASS_COMMA => MEMBER_DUE, + _ => REJECTED, + } +} + +/// Closes a member whose quality parameter is complete. +/// +/// A second parameter is always an error in this grammar, so a semicolon +/// leaves the subset rather than opening one. +const fn weighed_close(class: u8) -> u8 { + match class { + CLASS_OWS => MEMBER_OWS_WEIGHED, + CLASS_COMMA => MEMBER_DUE, + _ => REJECTED, + } +} + +/// Steps the item half of the machine for one shape. +const fn item_step(item: Item, state: u8, class: u8) -> u8 { + match item { + Item::Token => match state { + MEMBER_DUE | TOKEN if is_token_class(class) => TOKEN, + TOKEN => item_close(class), + _ => REJECTED, + }, + Item::LanguageRange => match state { + MEMBER_DUE if class == CLASS_STAR => LANGUAGE_STAR, + MEMBER_DUE if is_alpha_class(class) => PRIMARY_ONE, + LANGUAGE_STAR => item_close(class), + _ if state >= PRIMARY_ONE && state < PRIMARY_LAST && is_alpha_class(class) => state + 1, + _ if state >= PRIMARY_ONE && state <= PRIMARY_LAST && class == CLASS_HYPHEN => SUBTAG_DUE, + _ if state >= PRIMARY_ONE && state <= PRIMARY_LAST => item_close(class), + SUBTAG_DUE if is_alphanumeric_class(class) => SUBTAG_ONE, + _ if state >= SUBTAG_ONE && state < SUBTAG_LAST && is_alphanumeric_class(class) => state + 1, + _ if state >= SUBTAG_ONE && state <= SUBTAG_LAST && class == CLASS_HYPHEN => SUBTAG_DUE, + _ if state >= SUBTAG_ONE && state <= SUBTAG_LAST => item_close(class), + _ => REJECTED, + }, + } +} + +const fn transition_table(item: Item) -> [u16; STATES * CLASSES] { + let mut table = [row(REJECTED); STATES * CLASSES]; + let mut state = 0_usize; + while state < STATES { + let mut class = 0_usize; + while class < CLASSES { + #[expect(clippy::cast_possible_truncation, reason = "the loop bounds keep both indices inside a byte")] + let (state_value, class_value) = (state as u8, class as u8); + table[(state << 4) | class] = row(match state_value { + MEMBER_DUE => match class_value { + CLASS_OWS | CLASS_COMMA => MEMBER_DUE, + _ => item_step(item, MEMBER_DUE, class_value), + }, + MEMBER_OWS => item_close(class_value), + PARAM_DUE => match class_value { + CLASS_OWS => PARAM_DUE, + CLASS_Q => NAME_Q, + _ => REJECTED, + }, + NAME_Q => match class_value { + CLASS_EQUALS => QUALITY_DUE, + _ => REJECTED, + }, + QUALITY_DUE => match class_value { + CLASS_ZERO => QUALITY_ZERO, + CLASS_ONE => QUALITY_ONE, + _ => REJECTED, + }, + QUALITY_ZERO => match class_value { + CLASS_DOT => QUALITY_ZERO_DOT, + _ => weighed_close(class_value), + }, + QUALITY_ONE => match class_value { + CLASS_DOT => QUALITY_ONE_DOT, + _ => weighed_close(class_value), + }, + QUALITY_ZERO_DOT => match class_value { + CLASS_ZERO | CLASS_ONE | CLASS_DIGIT => QUALITY_ZERO_ONE_DIGIT, + _ => weighed_close(class_value), + }, + QUALITY_ZERO_ONE_DIGIT => match class_value { + CLASS_ZERO | CLASS_ONE | CLASS_DIGIT => QUALITY_ZERO_TWO_DIGITS, + _ => weighed_close(class_value), + }, + QUALITY_ZERO_TWO_DIGITS => match class_value { + CLASS_ZERO | CLASS_ONE | CLASS_DIGIT => QUALITY_ZERO_THREE_DIGITS, + _ => weighed_close(class_value), + }, + QUALITY_ONE_DOT => match class_value { + CLASS_ZERO => QUALITY_ONE_ONE_DIGIT, + _ => weighed_close(class_value), + }, + QUALITY_ONE_ONE_DIGIT => match class_value { + CLASS_ZERO => QUALITY_ONE_TWO_DIGITS, + _ => weighed_close(class_value), + }, + QUALITY_ONE_TWO_DIGITS => match class_value { + CLASS_ZERO => QUALITY_ONE_THREE_DIGITS, + _ => weighed_close(class_value), + }, + QUALITY_ZERO_THREE_DIGITS | QUALITY_ONE_THREE_DIGITS | MEMBER_OWS_WEIGHED => weighed_close(class_value), + _ => item_step(item, state_value, class_value), + }); + class += 1; + } + state += 1; + } + table +} + +/// Advances the table one byte. +/// +/// Every row is already a multiple of `CLASSES` and every class is smaller +/// than `CLASSES`, so the mask changes no index it is given and exists only to +/// prove the bound, which is what lets the step compile to two loads and no +/// branch. +#[expect(clippy::inline_always, reason = "the step is the unrolled loop body and must not become a call")] +#[inline(always)] +fn step(row: u16, byte: u8, transition: &[u16; STATES * CLASSES]) -> u16 { + let class = CLASS[usize::from(byte)]; + let index = (usize::from(row) | usize::from(class)) & (STATES * CLASSES - 1); + transition[index] +} + +/// Runs one table to completion over a whole field line. +/// +/// Stepping eight bytes per iteration amortizes the loop counter and branch, +/// which otherwise cost about as much as the two table loads they carry. +fn scan(bytes: &[u8], transition: &[u16; STATES * CLASSES], accepting: u64) -> bool { + let mut row = row(MEMBER_DUE); + + let mut blocks = bytes.chunks_exact(8); + for block in &mut blocks { + for byte in block { + row = step(row, *byte, transition); + } + } + for byte in blocks.remainder() { + row = step(row, *byte, transition); + } + + accepting & (1 << (row >> 4)) != 0 +} + +/// Reports whether a whole `Accept-Encoding` field line is valid. +/// +/// A `false` answer means "not recognized" rather than "malformed": quoted +/// strings and whitespace around an equals sign are left to the general +/// parser, which also produces the diagnostic when the line really is +/// malformed. +pub(super) fn scan_accept_encoding_line(bytes: &[u8]) -> bool { + scan(bytes, &TOKEN_TRANSITION, TOKEN_ACCEPTING) +} + +/// Reports whether a whole `Accept-Language` field line is valid. +/// +/// A `false` answer means "not recognized" rather than "malformed", exactly as +/// for [`scan_accept_encoding_line`]. +pub(super) fn scan_accept_language_line(bytes: &[u8]) -> bool { + scan(bytes, &LANGUAGE_TRANSITION, LANGUAGE_ACCEPTING) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::hint::black_box; + + use super::{ + CLASS, CLASSES, Item, LANGUAGE_ACCEPTING, LANGUAGE_STAR, LANGUAGE_TRANSITION, PRIMARY_LAST, PRIMARY_ONE, REJECTED, + SHARED_ACCEPTING, STATES, SUBTAG_LAST, SUBTAG_ONE, TOKEN, TOKEN_ACCEPTING, TOKEN_TRANSITION, class_table, range_mask, row, + transition_table, + }; + + #[test] + fn runtime_tables_match_the_static_tables() { + assert_eq!(black_box(class_table()), CLASS); + + for (item, table) in [(Item::Token, &TOKEN_TRANSITION), (Item::LanguageRange, &LANGUAGE_TRANSITION)] { + let generated = black_box(transition_table(item)); + assert_eq!(&generated, table); + + for class in 0..CLASSES { + assert_eq!( + generated[usize::from(row(REJECTED)) | class], + row(REJECTED), + "rejection must absorb every later byte" + ); + } + for entry in generated { + assert!(usize::from(entry >> 4) < STATES, "every transition must land on a defined state"); + assert_eq!(entry & 0xf, 0, "every entry must be a state already multiplied by the class count"); + } + } + } + + #[test] + fn runtime_accepting_masks_match_the_constants() { + let token = SHARED_ACCEPTING | (1 << TOKEN); + let language = SHARED_ACCEPTING + | (1 << LANGUAGE_STAR) + | black_box(range_mask(PRIMARY_ONE, PRIMARY_LAST)) + | black_box(range_mask(SUBTAG_ONE, SUBTAG_LAST)); + + assert_eq!(token, TOKEN_ACCEPTING); + assert_eq!(language, LANGUAGE_ACCEPTING); + assert_eq!(TOKEN_ACCEPTING & (1 << REJECTED), 0); + assert_eq!(LANGUAGE_ACCEPTING & (1 << REJECTED), 0); + } +} diff --git a/crates/http_headers/src/headers/range/accept_ranges.rs b/crates/http_headers/src/headers/range/accept_ranges.rs new file mode 100644 index 000000000..5c42a04d9 --- /dev/null +++ b/crates/http_headers/src/headers/range/accept_ranges.rs @@ -0,0 +1,891 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::{fmt, str}; + +use super::super::shared::FieldLinesIter; +use super::super::{invalid_syntax, normalized_comma_value, trim_ows}; +use super::shared::validate_range_unit_for; +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Defines the `Accept-Ranges` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 14.3](https://www.rfc-editor.org/rfc/rfc9110#section-14.3). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{AcceptRanges, AcceptRangesOwned}; +/// +/// let mut map = HeaderMap::new(); +/// AcceptRanges::insert(&mut map, AcceptRangesOwned::from_units(["bytes"])?)?; +/// assert!(AcceptRanges::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct AcceptRanges { + _private: (), +} + +/// Owned value for the `Accept-Ranges` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 14.3]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::AcceptRangesOwned::from_units(["bytes"])?; +/// assert_eq!(value.units().collect::>(), ["bytes"]); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Accept-Ranges: bytes` advertises byte ranges, `Accept-Ranges: none` +/// advertises no supported unit, and `Accept-Ranges: custom-unit` preserves an +/// extension range unit. +/// +/// [RFC 9110 section 14.3]: https://www.rfc-editor.org/rfc/rfc9110#section-14.3 +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct AcceptRangesOwned { + values: Option, + none: bool, +} + +/// Borrowed value for the `Accept-Ranges` header. +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{AcceptRanges, AcceptRangesView}; +/// use http_headers::source::{FieldLines, FieldSource}; +/// use http_headers::{Field, FieldName}; +/// +/// struct Source; +/// impl FieldSource for Source { +/// fn lines(&self, name: &'static FieldName) -> Option> { +/// (name == &FieldName::AcceptRanges) +/// .then_some(FieldLines::single(&FieldName::AcceptRanges, b"bytes")) +/// } +/// } +/// +/// let source = Source; +/// let view: AcceptRangesView<'_> = AcceptRanges::view(&source)?.expect("present"); +/// assert_eq!(view.units().collect::>(), ["bytes"]); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct AcceptRangesView<'a> { + values: Option>, + none: bool, +} + +impl fmt::Debug for AcceptRangesOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("AcceptRangesOwned") + .field("value_count", &self.values.as_ref().map_or(1, FieldLinesIter::len)) + .field("none", &self.none) + .finish() + } +} + +impl fmt::Display for AcceptRangesOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut units = self.units(); + let unit = units.next().expect("constructors guarantee at least one range unit"); + f.write_str(unit)?; + for unit in units { + f.write_str(", ")?; + f.write_str(unit)?; + } + Ok(()) + } +} + +impl fmt::Debug for AcceptRangesView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("AcceptRangesView") + .field("value_count", &self.values.as_ref().map_or(1, FieldLines::len)) + .field("none", &self.none) + .finish() + } +} + +impl AcceptRangesOwned { + #[cfg(all(feature = "serde", feature = "headers-range"))] + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + let canonical = self + .values + .is_none() + .then_some(FieldValueRef::new(if self.none { &b"none"[..] } else { &b"bytes"[..] })); + canonical.into_iter().chain( + self.values + .iter() + .flat_map(FieldLinesIter::iter) + .map(FieldValue::as_field_value_ref), + ) + } + + pub(crate) fn encoded_value(&self) -> Option { + normalized_comma_value(&FieldName::AcceptRanges, self.units().map(str::as_bytes)) + .expect("validated range units remain valid when comma-joined") + } + + /// Constructs `Accept-Ranges: bytes`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AcceptRangesOwned; + /// + /// let value = AcceptRangesOwned::bytes(); + /// assert!(!value.is_none()); + /// assert_eq!(value.units().collect::>(), ["bytes"]); + /// ``` + pub fn bytes() -> Self { + Self { values: None, none: false } + } + + /// Constructs `Accept-Ranges: none`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AcceptRangesOwned; + /// + /// let value = AcceptRangesOwned::none(); + /// assert!(value.is_none()); + /// assert_eq!(value.units().collect::>(), ["none"]); + /// + /// let bytes = AcceptRangesOwned::bytes(); + /// assert!(!bytes.is_none()); + /// ``` + pub fn none() -> Self { + Self { values: None, none: true } + } + + /// Constructs a canonical comma-separated unit list. + /// + /// # Errors + /// + /// Returns an error for an empty list, invalid token, or `none` combined + /// with another unit. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AcceptRangesOwned; + /// + /// let value = AcceptRangesOwned::from_units(["bytes", "items"])?; + /// assert_eq!(value.units().collect::>(), ["bytes", "items"]); + /// assert!(AcceptRangesOwned::from_units(["none", "bytes"]).is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_units(units: I) -> Result + where + I: IntoIterator, + S: AsRef, + { + let units = units.into_iter(); + let mut wire = String::with_capacity(units.size_hint().0.saturating_mul(8)); + let mut count = 0_usize; + let mut has_none = false; + for unit in units { + let unit = unit.as_ref(); + if !validate::token(unit.as_bytes()) { + return Err(DecodeError::new(&FieldName::AcceptRanges, DecodeErrorKind::InvalidToken)); + } + if !wire.is_empty() { + wire.push_str(", "); + } + wire.push_str(unit); + count += 1; + has_none |= unit.eq_ignore_ascii_case("none"); + } + if wire.is_empty() { + return Err(DecodeError::new(&FieldName::AcceptRanges, DecodeErrorKind::MissingValue)); + } + if has_none && count != 1 { + return Err(invalid_syntax(&FieldName::AcceptRanges)); + } + if wire == "bytes" { + return Ok(Self::bytes()); + } + if wire == "none" { + return Ok(Self::none()); + } + let value = validated_units_value(wire); + Ok(Self { + values: Some(FieldLinesIter::one(value)), + none: has_none, + }) + } + + /// Returns whether the field exclusively contains `none`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AcceptRangesOwned; + /// + /// let none = AcceptRangesOwned::none(); + /// assert!(none.is_none()); + /// + /// let bytes = AcceptRangesOwned::bytes(); + /// assert!(!bytes.is_none()); + /// ``` + pub const fn is_none(&self) -> bool { + self.none + } + + /// Iterates advertised range units in wire order. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AcceptRangesOwned; + /// + /// let value = AcceptRangesOwned::from_units(["bytes", "items"])?; + /// let units = value.units().collect::>(); + /// assert_eq!(units, ["bytes", "items"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn units(&self) -> impl Iterator { + let canonical = self.values.is_none().then_some(if self.none { "none" } else { "bytes" }); + canonical + .into_iter() + .chain(self.values.iter().flat_map(FieldLinesIter::iter).flat_map(|value| { + value + .as_bytes() + .split(|byte| *byte == b',') + .map(trim_ows) + .filter_map(|unit| str::from_utf8(unit).ok()) + .filter(|unit| !unit.is_empty()) + })) + } +} + +impl<'a> AcceptRangesView<'a> { + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + let canonical = self + .values + .is_none() + .then_some(FieldValueRef::new(if self.none { &b"none"[..] } else { &b"bytes"[..] })); + canonical.into_iter().chain(self.values.iter().flat_map(FieldLines::repeated)) + } + + /// Returns whether the field exclusively contains `none`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AcceptRanges; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source(&'static [u8]); + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::AcceptRanges) + /// .then_some(FieldLines::single(&FieldName::AcceptRanges, self.0)) + /// } + /// } + /// + /// let none_source = Source(b"none"); + /// let none = AcceptRanges::view(&none_source)?.expect("present"); + /// assert!(none.is_none()); + /// + /// let bytes_source = Source(b"bytes"); + /// let bytes = AcceptRanges::view(&bytes_source)?.expect("present"); + /// assert!(!bytes.is_none()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn is_none(&self) -> bool { + self.none + } + + /// Iterates advertised range units in wire order. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::AcceptRanges; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source(&'static [u8]); + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::AcceptRanges) + /// .then_some(FieldLines::single(&FieldName::AcceptRanges, self.0)) + /// } + /// } + /// + /// let source = Source(b"bytes, items"); + /// let view = AcceptRanges::view(&source)?.expect("present"); + /// let units = view.units().collect::>(); + /// assert_eq!(units, ["bytes", "items"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn units(&self) -> impl Iterator + '_ { + let canonical = self.values.is_none().then_some(if self.none { "none" } else { "bytes" }); + canonical.into_iter().chain( + self.values + .iter() + .flat_map(FieldLines::comma_items) + .filter_map(|item| item.ok().and_then(|item| str::from_utf8(item).ok())), + ) + } +} + +impl Field for AcceptRanges { + type View<'a> = AcceptRangesView<'a>; + type Owned = AcceptRangesOwned; + + fn name() -> &'static FieldName { + &FieldName::AcceptRanges + } + + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + accept_ranges_view(source.lines(Self::name())) + } + + /// Validates and clones in a single pass over the field values. + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + accept_ranges_owned(source.lines(Self::name())) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + let encoded = value.encoded_value().map_or_else(EncodedValues::new, EncodedValues::single); + sink.set_values(Self::name(), encoded) + } +} + +fn accept_ranges_view(values: Option>) -> Result>, DecodeError> { + let Some(values) = values else { + return Ok(None); + }; + values.validate_list_item_limit(b',', true)?; + if let Some(none) = lone_canonical_unit(&values) { + return Ok(Some(AcceptRangesView { values: None, none })); + } + let none = validate_accept_ranges(&values)?; + Ok(Some(AcceptRangesView { + values: Some(values), + none, + })) +} + +fn accept_ranges_owned(values: Option>) -> Result, DecodeError> { + let Some(values) = values else { + return Ok(None); + }; + values.validate_list_item_limit(b',', true)?; + let mut repeated = values.repeated(); + if let Some(first) = repeated.next() + && repeated.next().is_none() + && let Some(none) = canonical_unit(first.as_bytes()) + { + return Ok(Some(AcceptRangesOwned { values: None, none })); + } + + let mut repeated = values.repeated(); + if let Some(first) = repeated.next() + && repeated.next().is_none() + && let Some((units, none)) = scan_units(first.as_bytes()) + { + validate_none_cardinality(units, none)?; + let first_owned = first + .try_to_field_value() + .map_err(|_invalid| invalid_syntax(&FieldName::AcceptRanges))?; + return Ok(Some(AcceptRangesOwned { + values: Some(FieldLinesIter::one(first_owned)), + none, + })); + } + + let none = validate_accept_ranges(&values)?; + let mut repeated = values.repeated_owned()?; + let mut collected = repeated.next().map(|(_, first_owned)| FieldLinesIter::one(first_owned)); + for (_value, owned) in repeated { + collected.as_mut().expect("a later value has a preceding first value").push(owned); + } + Ok(Some(AcceptRangesOwned { values: collected, none })) +} + +super::super::shared::impl_string_conversions!(AcceptRangesOwned, &FieldName::AcceptRanges, invalid_syntax, wire); + +impl TryFrom for AcceptRangesOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + if let Some(none) = canonical_unit(value.as_bytes()) { + return Ok(Self { values: None, none }); + } + let none = if let Some((units, none)) = scan_units(value.as_bytes()) { + validate_none_cardinality(units, none)?; + none + } else { + validate_units_slow(value.as_bytes())? + }; + Ok(Self { + values: Some(FieldLinesIter::one(value)), + none, + }) + } +} + +fn validated_units_value(wire: String) -> FieldValue { + FieldValue::try_from(wire).expect("validated range units remain valid when comma-joined") +} + +/// Reports whether the field carries exactly one value that is just `bytes`. +fn lone_canonical_unit(values: &FieldLines<'_>) -> Option { + let mut repeated = values.repeated(); + let first = repeated.next().expect("untyped values always contain a field line"); + repeated.next().is_none().then(|| canonical_unit(first.as_bytes())).flatten() +} + +fn canonical_unit(bytes: &[u8]) -> Option { + match bytes { + b"bytes" => Some(false), + b"none" => Some(true), + _ => None, + } +} + +fn validate_accept_ranges(values: &FieldLines<'_>) -> Result { + let mut repeated = values.repeated(); + let first = repeated.next().expect("untyped values always contain a field line"); + if repeated.next().is_none() + && let Some((units, none)) = scan_units(first.as_bytes()) + { + validate_none_cardinality(units, none)?; + return Ok(none); + } + validate_accept_ranges_repeated(values) +} + +/// Validates a field spread over several lines. +fn validate_accept_ranges_repeated(values: &FieldLines<'_>) -> Result { + let mut units = 0_usize; + let mut none = false; + for value in values.repeated() { + let Some((scanned, has_none)) = scan_units(value.as_bytes()) else { + return validate_accept_ranges_slow(values); + }; + units += scanned; + none |= has_none; + } + validate_none_cardinality(units, none)?; + Ok(none) +} + +#[cold] +#[inline(never)] +fn validate_accept_ranges_slow(values: &FieldLines<'_>) -> Result { + let mut units = 0_usize; + let mut none = false; + for item in values.comma_items() { + let item = item?; + validate_range_unit(item)?; + units += 1; + if validate::eq_ignore_ascii_case(item, b"none") { + none = true; + } + } + validate_none_cardinality(units, none)?; + Ok(none) +} + +#[cold] +#[inline(never)] +pub(super) fn validate_units_slow(bytes: &[u8]) -> Result { + let mut units = 0_usize; + let mut none = false; + for item in bytes.split(|byte| *byte == b',') { + let item = trim_ows(item); + if item.is_empty() { + continue; + } + validate_range_unit(item)?; + units += 1; + if validate::eq_ignore_ascii_case(item, b"none") { + none = true; + } + } + validate_none_cardinality(units, none)?; + Ok(none) +} + +/// Counts the units of a well-formed token list in one pass. +/// +/// Returns `None` for anything unusual so the general implementation can +/// decide between acceptance and the precise error it reports. +pub(super) fn scan_units(bytes: &[u8]) -> Option<(usize, bool)> { + if bytes == b"bytes" { + return Some((1, false)); + } + if bytes == b"none" { + return Some((1, true)); + } + let mut rest = bytes; + let mut units = 0_usize; + let mut none = false; + loop { + while let [b' ' | b'\t', tail @ ..] = rest { + rest = tail; + } + let end = rest.iter().position(|byte| !TOKEN_BYTE[*byte as usize]).unwrap_or(rest.len()); + let (item, tail) = rest.split_at(end); + rest = tail; + while let [b' ' | b'\t', tail @ ..] = rest { + rest = tail; + } + if !item.is_empty() { + units += 1; + none |= is_none_unit(item); + } + match rest { + [] => return Some((units, none)), + [b',', tail @ ..] => rest = tail, + _ => return None, + } + } +} + +fn is_none_unit(item: &[u8]) -> bool { + let Ok(unit) = <[u8; 4]>::try_from(item) else { + return false; + }; + [unit[0] | 0x20, unit[1] | 0x20, unit[2] | 0x20, unit[3] | 0x20] == *b"none" +} + +/// Maps every byte to whether it is permitted in a token. +static TOKEN_BYTE: [bool; 256] = { + let mut table = [false; 256]; + let mut index = 0_u8; + loop { + table[index as usize] = validate::token_byte(index); + if index == u8::MAX { + break; + } + index += 1; + } + table +}; + +pub(super) fn validate_none_cardinality(units: usize, none: bool) -> Result<(), DecodeError> { + if units == 0 { + Err(DecodeError::new(&FieldName::AcceptRanges, DecodeErrorKind::MissingValue)) + } else if none && units != 1 { + Err(invalid_syntax(&FieldName::AcceptRanges)) + } else { + Ok(()) + } +} + +fn validate_range_unit(bytes: &[u8]) -> Result<(), DecodeError> { + validate_range_unit_for(bytes, &FieldName::AcceptRanges) +} + +#[cfg(test)] +#[expect( + clippy::assertions_on_result_states, + reason = "the tests classify many parser outcomes without needing their success values" +)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{ + AcceptRanges, AcceptRangesOwned, accept_ranges_view, canonical_unit, is_none_unit, scan_units, validate_accept_ranges, + validate_accept_ranges_slow, validate_none_cardinality, validate_units_slow, + }; + use crate::sink::{EncodedValues, FieldSink, InsertError}; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeErrorKind, FieldName, FieldValue}; + + #[derive(Default)] + struct Store { + values: Vec, + } + + impl Store { + fn new(values: &[&'static str]) -> Self { + Self { + values: values.iter().copied().map(FieldValue::from_static).collect(), + } + } + } + + impl FieldSource for Store { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::AcceptRanges) + .then(|| FieldLines::from_slice(name, &self.values)) + .flatten() + } + } + + impl FieldSink for Store { + fn set_values(&mut self, _name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.values = values.into_iter().collect(); + Ok(()) + } + + fn append_values(&mut self, _name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.values.extend(values); + Ok(()) + } + + fn remove_values(&mut self, _name: &'static FieldName) { + self.values.clear(); + } + } + + #[test] + fn singleton_probes_preserve_exact_units_and_canonical_storage() { + for (lines, none, units, canonical) in [ + (&["bytes"][..], false, &["bytes"][..], true), + (&["none"][..], true, &["none"][..], true), + (&["Bytes"][..], false, &["Bytes"][..], false), + (&[" \tbytes \t"][..], false, &["bytes"][..], false), + (&["bytes, items"][..], false, &["bytes", "items"][..], false), + (&["bytes", "items", "records"][..], false, &["bytes", "items", "records"][..], false), + ] { + let store = Store::new(lines); + let borrowed = AcceptRanges::view(&store).unwrap().unwrap(); + let owned = AcceptRanges::owned(&store).unwrap().unwrap(); + assert_eq!(borrowed.is_none(), none, "{lines:?}"); + assert_eq!(owned.is_none(), none, "{lines:?}"); + assert_eq!(borrowed.units().collect::>(), units, "{lines:?}"); + assert_eq!(owned.units().collect::>(), units, "{lines:?}"); + assert_eq!(borrowed.values.is_none(), canonical, "{lines:?}"); + assert_eq!(owned.values.is_none(), canonical, "{lines:?}"); + } + } + + #[test] + fn borrowed_singleton_probe_matches_owned_values_and_errors() { + for lines in [ + &[][..], + &["bytes"], + &["none"], + &["Bytes"], + &[" \tbytes \t"], + &["bytes, items"], + &["bytes", "items", "records"], + &["none", "bytes"], + &["bytes", "bad/unit"], + &["\"bytes"], + &[""], + ] { + let store = Store::new(lines); + let borrowed = AcceptRanges::view(&store) + .map(|view| view.map(|view| (view.is_none(), view.units().map(str::to_owned).collect::>()))); + let owned = AcceptRanges::owned(&store) + .map(|value| value.map(|value| (value.is_none(), value.units().map(str::to_owned).collect::>()))); + assert_eq!(borrowed, owned, "{lines:?}"); + } + } + + #[test] + fn constructors_cover_canonical_extension_and_rejection_paths() { + let bytes = AcceptRangesOwned::bytes(); + assert!(!bytes.is_none()); + assert_eq!(bytes.units().collect::>(), ["bytes"]); + assert!(format!("{bytes:?}").contains("none: false")); + + let none = AcceptRangesOwned::none(); + assert!(none.is_none()); + assert_eq!(none.units().collect::>(), ["none"]); + assert_eq!( + AcceptRangesOwned::from_units(vec!["bytes"]) + .expect("canonical bytes") + .units() + .collect::>(), + ["bytes"] + ); + assert!(AcceptRangesOwned::from_units(vec!["none"]).expect("canonical none").is_none()); + + let extensions = AcceptRangesOwned::from_units(vec!["bytes", "items"]).expect("two valid range units"); + assert_eq!(extensions.units().collect::>(), ["bytes", "items"]); + assert_eq!( + AcceptRangesOwned::from_units(Vec::<&str>::new()).expect_err("empty list").kind(), + DecodeErrorKind::MissingValue + ); + assert_eq!( + AcceptRangesOwned::from_units(vec!["bad unit"]).expect_err("invalid token").kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + AcceptRangesOwned::from_units(vec!["none", "bytes"]) + .expect_err("none must stand alone") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + AcceptRangesOwned::try_from(String::from("line\nbreak")) + .expect_err("invalid field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + AcceptRangesOwned::try_from("line\nbreak") + .expect_err("invalid borrowed field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + AcceptRangesOwned::try_from("Bytes, items") + .expect("valid borrowed string") + .units() + .collect::>(), + ["Bytes", "items"] + ); + assert_eq!( + AcceptRangesOwned::try_from(String::from("Bytes, items")) + .expect("valid owned string") + .units() + .collect::>(), + ["Bytes", "items"] + ); + + let canonical = AcceptRangesOwned::try_from(FieldValue::from_static("bytes")).expect("canonical"); + assert!(canonical.values.is_none()); + let scanned = AcceptRangesOwned::try_from(FieldValue::from_static("Bytes, items")).expect("fast scanned list"); + assert!(scanned.values.is_some()); + let fallback = AcceptRangesOwned::try_from(FieldValue::from_static("bytes ,\titems")).expect("fallback list"); + assert_eq!(fallback.units().collect::>(), ["bytes", "items"]); + assert_eq!( + AcceptRangesOwned::try_from(FieldValue::from_static("bad unit")) + .expect_err("slow validation reports invalid syntax") + .kind(), + DecodeErrorKind::InvalidToken + ); + } + + #[test] + fn header_views_and_owned_values_cover_canonical_and_repeated_storage() { + assert!(AcceptRanges::view(&Store::default()).expect("absent").is_none()); + assert!(AcceptRanges::owned(&Store::default()).expect("absent").is_none()); + + for (wire, expected_none) in [("bytes", false), ("none", true)] { + let store = Store::new(&[wire]); + let view = AcceptRanges::view(&store).expect("valid canonical value").expect("present"); + assert_eq!(view.is_none(), expected_none); + assert_eq!(view.units().collect::>(), [wire]); + let fields = view.field_values().collect::>(); + assert_eq!(fields.len(), 1); + assert_eq!(fields[0].as_bytes(), wire.as_bytes()); + assert!(format!("{view:?}").contains("value_count: 1")); + + let owned = AcceptRanges::owned(&store).expect("valid canonical value").expect("present"); + assert_eq!(owned.is_none(), expected_none); + } + + let store = Store::new(&["bytes, items", "records"]); + let view = AcceptRanges::view(&store).expect("valid repeated values").expect("present"); + assert_eq!(view.units().collect::>(), ["bytes", "items", "records"]); + assert_eq!(view.field_values().count(), 2); + let owned = AcceptRanges::owned(&store).expect("valid repeated values").expect("present"); + assert_eq!(owned.units().collect::>(), ["bytes", "items", "records"]); + + let single = Store::new(&["Bytes, items"]); + let owned = AcceptRanges::owned(&single).expect("fast single-line validation").expect("present"); + assert_eq!(owned.units().collect::>(), ["Bytes", "items"]); + + assert!(AcceptRanges::view(&Store::new(&["bad unit"])).is_err()); + assert!(AcceptRanges::owned(&Store::new(&["none, bytes"])).is_err()); + assert!(AcceptRanges::owned(&Store::new(&["bad unit"])).is_err()); + + let mut sink = Store::default(); + AcceptRanges::insert(&mut sink, owned).expect("normalized insert"); + assert_eq!(sink.values.len(), 1); + assert_eq!(sink.values[0].as_bytes(), b"Bytes, items"); + sink.remove_values(&FieldName::AcceptRanges); + assert!(sink.values.is_empty()); + } + + #[test] + fn scanners_and_fallback_validation_agree_on_edge_lists() { + assert_eq!(canonical_unit(b"bytes"), Some(false)); + assert_eq!(canonical_unit(b"none"), Some(true)); + assert_eq!(canonical_unit(b"Bytes"), None); + assert!(!is_none_unit(b"byte")); + assert!(is_none_unit(b"NoNe")); + + for (wire, expected) in [ + (b"bytes".as_slice(), Some((1, false))), + (b"none", Some((1, true))), + (b" bytes , items ", Some((2, false))), + (b"bytes,,items", Some((2, false))), + (b"", Some((0, false))), + (b"bad unit", None), + ] { + assert_eq!(scan_units(wire), expected, "{wire:?}"); + } + assert_eq!( + validate_units_slow(b", ,").expect_err("empty members do not supply a unit").kind(), + DecodeErrorKind::MissingValue + ); + assert!(!validate_units_slow(b"bytes, items").expect("valid list")); + assert_eq!( + validate_units_slow(b"none, bytes").expect_err("none cannot be combined").kind(), + DecodeErrorKind::InvalidSyntax + ); + + assert_eq!( + validate_none_cardinality(0, false).expect_err("missing units").kind(), + DecodeErrorKind::MissingValue + ); + assert_eq!( + validate_none_cardinality(2, true).expect_err("none cardinality").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert!(validate_none_cardinality(1, true).is_ok()); + + let none = FieldLines::single(&FieldName::AcceptRanges, b"none"); + assert!(validate_accept_ranges_slow(&none).expect("direct slow-path validation")); + assert!(AcceptRangesOwned::try_from(FieldValue::from_static("none, bytes")).is_err()); + let invalid_single = FieldLines::single(&FieldName::AcceptRanges, b"none, bytes"); + assert!(validate_accept_ranges(&invalid_single).is_err()); + let repeated = [FieldValue::from_static("none"), FieldValue::from_static("bytes")]; + let invalid_repeated = FieldLines::from_slice(&FieldName::AcceptRanges, &repeated).expect("two field lines"); + assert!(validate_accept_ranges(&invalid_repeated).is_err()); + assert!(validate_accept_ranges_slow(&invalid_repeated).is_err()); + + let values = [FieldValue::from_static("bytes"), FieldValue::from_static("bad unit")]; + let repeated = FieldLines::from_slice(&FieldName::AcceptRanges, &values).expect("values"); + assert_eq!( + validate_accept_ranges(&repeated).expect_err("invalid repeated unit").kind(), + DecodeErrorKind::InvalidToken + ); + + let one = [FieldValue::from_static("bytes")]; + let values = FieldLines::from_slice(&FieldName::AcceptRanges, &one).expect("one value"); + let view = accept_ranges_view(Some(values)).unwrap().unwrap(); + assert!(!view.is_none()); + assert!(view.values.is_none()); + + let invalid = [FieldValue::from_static("\"unterminated")]; + let values = FieldLines::from_slice(&FieldName::AcceptRanges, &invalid).expect("one value"); + assert_eq!( + validate_accept_ranges(&values).expect_err("unterminated quoted member").kind(), + DecodeErrorKind::UnterminatedQuote + ); + } +} diff --git a/crates/http_headers/src/headers/range/content_range.rs b/crates/http_headers/src/headers/range/content_range.rs new file mode 100644 index 000000000..d85f44ddc --- /dev/null +++ b/crates/http_headers/src/headers/range/content_range.rs @@ -0,0 +1,1077 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt::{self, Write as _}; +use std::ops::{Range as IndexRange, RangeBounds}; +use std::str; + +use super::super::invalid_syntax; +use super::range::ByteRangeSpec; +use super::shared::{parse_number, validate_range_unit_for}; +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef, SingleValueField, validate}; + +/// The complete representation length in a satisfied byte content range. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub enum CompleteLength { + /// The complete representation length is known. + Known(u64), + /// The complete representation length is unknown and is serialized as `*`. + Unknown, +} + +impl CompleteLength { + const fn into_option(self) -> Option { + match self { + Self::Known(length) => Some(length), + Self::Unknown => None, + } + } +} + +impl From for CompleteLength { + fn from(length: u64) -> Self { + Self::Known(length) + } +} + +/// The byte-specific interpretation of a `Content-Range` value. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{ByteContentRange, ContentRangeOwned}; +/// +/// let satisfied = ContentRangeOwned::try_from("bytes 0-99/200")?; +/// assert_eq!( +/// satisfied.byte_range(), +/// Some(ByteContentRange::Satisfied { +/// first: 0, +/// last: 99, +/// complete_length: Some(200), +/// }) +/// ); +/// +/// let unsatisfied = ContentRangeOwned::try_from("bytes */200")?; +/// assert_eq!( +/// unsatisfied.byte_range(), +/// Some(ByteContentRange::Unsatisfied { +/// complete_length: 200, +/// }) +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub enum ByteContentRange { + /// A satisfied inclusive range. + Satisfied { + /// The first byte position. + first: u64, + /// The last byte position. + last: u64, + /// The complete representation length, or `None` when unknown. + complete_length: Option, + }, + /// An unsatisfied range response with the current complete length. + Unsatisfied { + /// The complete representation length. + complete_length: u64, + }, +} + +/// Defines the `Content-Range` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 14.4](https://www.rfc-editor.org/rfc/rfc9110#section-14.4). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{ContentRange, ContentRangeOwned}; +/// +/// let mut map = HeaderMap::new(); +/// ContentRange::insert(&mut map, ContentRangeOwned::try_from("bytes 0-99/200")?)?; +/// assert!(ContentRange::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct ContentRange { + _private: (), +} + +/// Owned value for the `Content-Range` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 14.4]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::ContentRangeOwned::try_from("bytes 0-99/200")?; +/// assert_eq!( +/// value.byte_range(), +/// Some(http_headers::headers::ByteContentRange::Satisfied { +/// first: 0, +/// last: 99, +/// complete_length: Some(200), +/// }) +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Content-Range: bytes 0-499/1234` is satisfied, +/// `Content-Range: bytes 500-999/*` has an unknown complete length, and +/// `Content-Range: bytes */1234` is unsatisfied. Extension forms such as +/// `Content-Range: custom opaque-payload` are preserved. +/// +/// [RFC 9110 section 14.4]: https://www.rfc-editor.org/rfc/rfc9110#section-14.4 +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct ContentRangeOwned { + value: FieldValue, + unit: IndexRange, + payload: IndexRange, + parsed: Option, +} + +/// Borrowed value for the `Content-Range` header. +#[derive(Clone, Copy, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{ContentRange, ContentRangeView}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let view: ContentRangeView<'_> = +/// ContentRange::decode_view(FieldValueRef::new(b"bytes 0-99/200"))?; +/// assert_eq!(view.unit(), "bytes"); +/// assert_eq!(view.as_field_value().as_bytes(), b"bytes 0-99/200"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct ContentRangeView<'a> { + value: FieldValueRef<'a>, + unit: &'a str, + payload: &'a [u8], + parsed: Option, +} + +impl fmt::Debug for ContentRangeOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("ContentRangeOwned") + .field("unit", &self.unit()) + .field("byte_range", &self.parsed) + .finish_non_exhaustive() + } +} + +impl fmt::Debug for ContentRangeView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("ContentRangeView") + .field("unit", &self.unit) + .field("byte_range", &self.parsed) + .finish_non_exhaustive() + } +} + +impl ContentRangeOwned { + /// Constructs a satisfied byte content range from explicit positions. + /// + /// Prefer [`Self::bytes_range`] when the caller already holds a Rust + /// range. + /// + /// # Errors + /// + /// Returns an error for an inverted range or a complete length not + /// greater than the last position. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{CompleteLength, ContentRangeOwned}; + /// + /// let value = ContentRangeOwned::bytes(0, 9, CompleteLength::Known(10))?; + /// assert_eq!(value.as_field_value().as_bytes(), b"bytes 0-9/10"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn bytes(first: u64, last: u64, complete_length: CompleteLength) -> Result { + Self::bytes_range(first..=last, complete_length) + } + + /// Constructs a satisfied byte content range. + /// + /// Included bounds map directly to HTTP's inclusive positions. An + /// excluded start is incremented and an excluded end is decremented, so + /// both `0..100` and `0..=99` represent bytes 0 through 99. + /// + /// # Errors + /// + /// Returns an error for an unbounded bound, an empty or inverted range, a + /// bound adjustment that overflows, or a complete length not greater than + /// the last position. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{CompleteLength, ContentRangeOwned}; + /// + /// // Half-open and inclusive ranges describe the same bytes. + /// let half_open = ContentRangeOwned::bytes_range(0..100, CompleteLength::Known(1000))?; + /// let inclusive = ContentRangeOwned::bytes_range(0..=99, CompleteLength::Known(1000))?; + /// assert_eq!(half_open.as_field_value().as_bytes(), b"bytes 0-99/1000"); + /// assert_eq!( + /// half_open.as_field_value().as_bytes(), + /// inclusive.as_field_value().as_bytes() + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn bytes_range(range: impl RangeBounds, complete_length: CompleteLength) -> Result { + Self::bytes_range_from_spec(ByteRangeSpec::from_range(range), complete_length) + } + + fn bytes_range_from_spec(spec: Result, complete_length: CompleteLength) -> Result { + let ByteRangeSpec::FromTo { first, last } = spec? else { + return Err(invalid_syntax(&FieldName::ContentRange)); + }; + let complete_length = complete_length.into_option(); + validate_satisfied_content_range(first, last, complete_length)?; + let mut wire = format!("bytes {first}-{last}/"); + if let Some(value) = complete_length { + write!(&mut wire, "{value}").map_err(|_| invalid_syntax(&FieldName::ContentRange))?; + } else { + wire.push('*'); + } + Ok(content_range_from_parts( + wire, + 5, + Some(ByteContentRange::Satisfied { + first, + last, + complete_length, + }), + )) + } + + /// Constructs an unsatisfied byte content range. + /// + /// This represents the wire form `bytes */complete_length`. + /// + /// # Errors + /// + /// The result type is retained for constructor API consistency; every + /// `u64` complete length has a valid representation. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ByteContentRange, ContentRangeOwned}; + /// + /// let value = ContentRangeOwned::unsatisfied_bytes(1234)?; + /// assert_eq!(value.as_field_value().as_bytes(), b"bytes */1234"); + /// assert_eq!( + /// value.byte_range(), + /// Some(ByteContentRange::Unsatisfied { + /// complete_length: 1234, + /// }) + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::unnecessary_wraps, + reason = "content-range constructors consistently report validation through DecodeError" + )] + pub fn unsatisfied_bytes(complete_length: u64) -> Result { + Ok(content_range_from_parts( + format!("bytes */{complete_length}"), + 5, + Some(ByteContentRange::Unsatisfied { complete_length }), + )) + } + + /// Constructs an extension content range. + /// + /// The extension payload is preserved; this type does not claim to + /// normalize extension range semantics. + /// + /// # Errors + /// + /// Returns an error when the unit is not a token, is the reserved + /// case-insensitive `bytes` unit, or the payload contains bytes outside + /// the extension content-range grammar. Use [`Self::bytes`] or + /// [`Self::unsatisfied_bytes`] for byte content ranges. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ContentRangeOwned; + /// + /// let value = ContentRangeOwned::extension("items", "0-9/100")?; + /// assert_eq!(value.unit()?, "items"); + /// assert_eq!(value.byte_range(), None); + /// assert_eq!(value.extension_payload(), Some(&b"0-9/100"[..])); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn extension(unit: impl AsRef, payload: impl AsRef) -> Result { + let unit = unit.as_ref(); + let payload = payload.as_ref(); + validate_range_unit_for(unit.as_bytes(), &FieldName::ContentRange)?; + if validate::eq_ignore_ascii_case(unit.as_bytes(), b"bytes") { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + if !valid_content_range_extension(payload.as_bytes()) { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + Ok(content_range_from_parts(format!("{unit} {payload}"), unit.len(), None)) + } + + /// Returns the range unit exactly as received. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the wire value. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ContentRangeOwned; + /// + /// let bytes = ContentRangeOwned::try_from("bytes 0-99/200")?; + /// assert_eq!(bytes.unit()?, "bytes"); + /// + /// let items = ContentRangeOwned::extension("items", "0-9/100")?; + /// assert_eq!(items.unit()?, "items"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn unit(&self) -> Result<&str, DecodeError> { + let bytes = self + .value + .as_bytes() + .get(self.unit.clone()) + .ok_or_else(|| invalid_syntax(&FieldName::ContentRange))?; + str::from_utf8(bytes).map_err(|_invalid| invalid_syntax(&FieldName::ContentRange)) + } + + /// Returns the byte-specific interpretation, if the unit is `bytes`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ByteContentRange, ContentRangeOwned}; + /// + /// let satisfied = ContentRangeOwned::try_from("bytes 0-99/200")?; + /// assert_eq!( + /// satisfied.byte_range(), + /// Some(ByteContentRange::Satisfied { + /// first: 0, + /// last: 99, + /// complete_length: Some(200), + /// }) + /// ); + /// + /// let extension = ContentRangeOwned::extension("items", "0-9/100")?; + /// assert_eq!(extension.byte_range(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn byte_range(&self) -> Option { + self.parsed + } + + /// Returns an extension payload without normalization. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ContentRangeOwned; + /// + /// let extension = ContentRangeOwned::extension("items", "0-9/100")?; + /// assert_eq!(extension.extension_payload(), Some(&b"0-9/100"[..])); + /// + /// let bytes = ContentRangeOwned::try_from("bytes 0-99/200")?; + /// assert_eq!(bytes.extension_payload(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn extension_payload(&self) -> Option<&[u8]> { + if self.parsed.is_some() { + None + } else { + self.value.as_bytes().get(self.payload.clone()) + } + } + + /// Returns the preserved field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ContentRangeOwned; + /// + /// let value = ContentRangeOwned::try_from("bytes 0-99/200")?; + /// assert_eq!(value.as_field_value().as_bytes(), b"bytes 0-99/200"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Consumes the header and returns its field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ContentRangeOwned; + /// + /// let value = ContentRangeOwned::try_from("bytes */1234")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"bytes */1234"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(ContentRangeOwned, |value| value.value); + +impl<'a> ContentRangeView<'a> { + /// Returns the range unit exactly as received. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ContentRange; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = ContentRange::decode_view(FieldValueRef::new(b"items 0-9/100"))?; + /// assert_eq!(view.unit(), "items"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn unit(self) -> &'a str { + self.unit + } + + /// Returns the byte-specific interpretation, if the unit is `bytes`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ByteContentRange, ContentRange}; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = ContentRange::decode_view(FieldValueRef::new(b"bytes 0-99/*"))?; + /// assert_eq!( + /// view.byte_range(), + /// Some(ByteContentRange::Satisfied { + /// first: 0, + /// last: 99, + /// complete_length: None, + /// }) + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn byte_range(self) -> Option { + self.parsed + } + + /// Returns an extension payload without normalization. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ContentRange; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = ContentRange::decode_view(FieldValueRef::new(b"items 0-9/100"))?; + /// assert_eq!(view.extension_payload(), Some(&b"0-9/100"[..])); + /// + /// let bytes = ContentRange::decode_view(FieldValueRef::new(b"bytes 0-99/200"))?; + /// assert_eq!(bytes.extension_payload(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn extension_payload(self) -> Option<&'a [u8]> { + if self.parsed.is_some() { None } else { Some(self.payload) } + } + + /// Returns the borrowed field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ContentRange; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = ContentRange::decode_view(FieldValueRef::new(b"bytes */1234"))?; + /// assert_eq!(view.as_field_value().as_bytes(), b"bytes */1234"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } +} + +impl SingleValueField for ContentRange { + type View<'a> = ContentRangeView<'a>; + type Owned = ContentRangeOwned; + + fn name() -> &'static FieldName { + &FieldName::ContentRange + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + let parsed = parse_content_range(value.as_bytes())?; + Ok(ContentRangeView { + value, + unit: parsed.unit, + payload: parsed.payload, + parsed: parsed.byte_range, + }) + } + + fn decode_owned(value: FieldValue) -> Result { + ContentRangeOwned::try_from(value) + } + + fn decode_view_with(value: FieldValueRef<'_>, mode: crate::DecodeMode) -> Result, DecodeError> { + let parsed = parse_content_range_with(value.as_bytes(), mode)?; + Ok(ContentRangeView { + value, + unit: parsed.unit, + payload: parsed.payload, + parsed: parsed.byte_range, + }) + } + + fn decode_owned_with(value: FieldValue, mode: crate::DecodeMode) -> Result { + content_range_from_value_with(value, mode) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +fn content_range_from_parts(wire: String, unit_end: usize, parsed: Option) -> ContentRangeOwned { + let payload_start = unit_end + 1; + let payload_end = wire.len(); + let value = FieldValue::try_from(wire).expect("validated content-range parts form a valid field value"); + ContentRangeOwned { + value, + unit: 0..unit_end, + payload: payload_start..payload_end, + parsed, + } +} + +super::super::shared::impl_string_conversions!(ContentRangeOwned, &FieldName::ContentRange, invalid_syntax, wire); + +impl TryFrom for ContentRangeOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + content_range_from_value_with(value, crate::DecodeMode::Strict) + } +} + +fn content_range_from_value_with(value: FieldValue, mode: crate::DecodeMode) -> Result { + let (unit_end, payload_end, byte_range) = { + let parsed = parse_content_range_with(value.as_bytes(), mode)?; + // Every parser preserves the unit prefix followed by one SP. + let unit_end = parsed.unit.len(); + (unit_end, unit_end + 1 + parsed.payload.len(), parsed.byte_range) + }; + Ok(ContentRangeOwned { + value, + unit: 0..unit_end, + payload: unit_end + 1..payload_end, + parsed: byte_range, + }) +} + +struct ParsedContentRange<'a> { + unit: &'a str, + payload: &'a [u8], + byte_range: Option, +} + +fn parse_content_range(bytes: &[u8]) -> Result, DecodeError> { + if let Some(payload) = bytes.strip_prefix(b"bytes ") { + let unit = http_headers_simd::ascii_str(&bytes[..5]).expect("the matched bytes unit is valid ASCII"); + return Ok(ParsedContentRange { + unit, + payload, + byte_range: Some(parse_byte_content_range(payload)?), + }); + } + parse_content_range_extension(bytes) +} + +fn parse_content_range_with(bytes: &[u8], mode: crate::DecodeMode) -> Result, DecodeError> { + if mode == crate::DecodeMode::Strict { + return parse_content_range(bytes); + } + if let Ok(parsed) = parse_content_range(bytes) { + return Ok(parsed); + } + parse_content_range_relaxed(bytes) +} + +fn parse_content_range_relaxed(bytes: &[u8]) -> Result, DecodeError> { + if !validate::field_value(bytes) { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + let Some(separator) = bytes.iter().position(|byte| *byte == b' ') else { + return Err(invalid_syntax(&FieldName::ContentRange)); + }; + let unit_bytes = &bytes[..separator]; + let payload = &bytes[separator + 1..]; + if payload.first().is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + validate_range_unit_for(unit_bytes, &FieldName::ContentRange)?; + let unit = http_headers_simd::ascii_str(unit_bytes).expect("validated range units are ASCII"); + let byte_range = if validate::eq_ignore_ascii_case(unit_bytes, b"bytes") { + Some(parse_byte_content_range_relaxed(payload)?) + } else { + if !valid_content_range_extension(payload) { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + None + }; + Ok(ParsedContentRange { unit, payload, byte_range }) +} + +#[cold] +#[inline(never)] +fn parse_content_range_extension(bytes: &[u8]) -> Result, DecodeError> { + let Some(separator) = bytes.iter().position(|byte| *byte == b' ') else { + return Err(invalid_syntax(&FieldName::ContentRange)); + }; + let unit_bytes = &bytes[..separator]; + let payload = &bytes[separator + 1..]; + validate_range_unit_for(unit_bytes, &FieldName::ContentRange)?; + let unit = http_headers_simd::ascii_str(unit_bytes).expect("validated range units are ASCII"); + let byte_range = if validate::eq_ignore_ascii_case(unit_bytes, b"bytes") { + Some(parse_byte_content_range(payload)?) + } else { + if !valid_content_range_extension(payload) { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + None + }; + Ok(ParsedContentRange { unit, payload, byte_range }) +} + +fn parse_byte_content_range(bytes: &[u8]) -> Result { + if let Some(complete) = bytes.strip_prefix(b"*/") { + return parse_number(complete, &FieldName::ContentRange).map(|complete_length| ByteContentRange::Unsatisfied { complete_length }); + } + let Some(slash) = bytes.iter().position(|byte| *byte == b'/') else { + return Err(invalid_syntax(&FieldName::ContentRange)); + }; + if bytes[slash + 1..].contains(&b'/') { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + let included = &bytes[..slash]; + let Some(dash) = included.iter().position(|byte| *byte == b'-') else { + return Err(invalid_syntax(&FieldName::ContentRange)); + }; + if included[dash + 1..].contains(&b'-') { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + let first = parse_number(&included[..dash], &FieldName::ContentRange)?; + let last = parse_number(&included[dash + 1..], &FieldName::ContentRange)?; + let complete = &bytes[slash + 1..]; + let complete_length = match complete { + b"*" => None, + _ => Some(parse_number(complete, &FieldName::ContentRange)?), + }; + validate_satisfied_content_range(first, last, complete_length)?; + Ok(ByteContentRange::Satisfied { + first, + last, + complete_length, + }) +} + +fn parse_byte_content_range_relaxed(bytes: &[u8]) -> Result { + let Some(slash) = bytes.iter().position(|byte| *byte == b'/') else { + return Err(invalid_syntax(&FieldName::ContentRange)); + }; + if bytes[slash + 1..].contains(&b'/') { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + let included = trim_ows_end(&bytes[..slash]); + let complete = trim_ows_start(&bytes[slash + 1..]); + if included == b"*" { + return parse_number(complete, &FieldName::ContentRange).map(|complete_length| ByteContentRange::Unsatisfied { complete_length }); + } + let Some(dash) = included.iter().position(|byte| *byte == b'-') else { + return Err(invalid_syntax(&FieldName::ContentRange)); + }; + if included[dash + 1..].contains(&b'-') { + return Err(invalid_syntax(&FieldName::ContentRange)); + } + let first = parse_number(trim_ows_end(&included[..dash]), &FieldName::ContentRange)?; + let last = parse_number(trim_ows_start(&included[dash + 1..]), &FieldName::ContentRange)?; + let complete_length = if complete == b"*" { + None + } else { + Some(parse_number(complete, &FieldName::ContentRange)?) + }; + validate_satisfied_content_range(first, last, complete_length)?; + Ok(ByteContentRange::Satisfied { + first, + last, + complete_length, + }) +} + +fn trim_ows_start(mut bytes: &[u8]) -> &[u8] { + while bytes.first().is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + bytes = &bytes[1..]; + } + bytes +} + +fn trim_ows_end(mut bytes: &[u8]) -> &[u8] { + while bytes.last().is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + bytes = &bytes[..bytes.len() - 1]; + } + bytes +} + +fn validate_satisfied_content_range(first: u64, last: u64, complete_length: Option) -> Result<(), DecodeError> { + if last < first || complete_length.is_some_and(|length| last >= length) { + Err(invalid_syntax(&FieldName::ContentRange)) + } else { + Ok(()) + } +} + +fn valid_content_range_extension(bytes: &[u8]) -> bool { + bytes.iter().all(|byte| matches!(byte, b'\t' | 0x20..=0x7e)) +} + +#[cfg(test)] +#[expect( + clippy::assertions_on_result_states, + reason = "the tests classify many parser outcomes without needing their success values" +)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::ops::Bound; + + use super::{ + ByteContentRange, CompleteLength, ContentRange, ContentRangeOwned, content_range_from_value_with, parse_byte_content_range, + parse_byte_content_range_relaxed, parse_content_range, parse_content_range_relaxed, parse_content_range_with, + valid_content_range_extension, validate_satisfied_content_range, + }; + use crate::{DecodeErrorKind, DecodeMode, FieldName, FieldValue, FieldValueRef, SingleValueField}; + + #[test] + fn relaxed_whitespace_trimming_is_directional_and_preserves_other_bytes() { + assert_eq!(super::trim_ows_start(b" \t1 \t"), b"1 \t"); + assert_eq!(super::trim_ows_end(b" \t1 \t"), b" \t1"); + for bytes in [b"".as_slice(), b" \t"] { + assert_eq!(super::trim_ows_start(bytes), b""); + assert_eq!(super::trim_ows_end(bytes), b""); + } + for byte in u8::MIN..=u8::MAX { + if !matches!(byte, b' ' | b'\t') { + let bytes = [byte, b'1', byte]; + assert_eq!(super::trim_ows_start(&bytes), bytes); + assert_eq!(super::trim_ows_end(&bytes), bytes); + } + } + } + + #[test] + fn owned_offsets_preserve_unit_and_payload_across_storage() { + for (wire, mode, unit, payload) in [ + ("bytes 0-9/10", DecodeMode::Strict, "bytes", b"0-9/10".as_slice()), + ("Bytes 0 - 9 / 10", DecodeMode::Relaxed, "Bytes", b"0 - 9 / 10"), + ("items ", DecodeMode::Strict, "items", b""), + ("items first-last/complete", DecodeMode::Strict, "items", b"first-last/complete"), + ] { + for value in [FieldValue::from_static(wire), FieldValue::try_from(wire.to_owned()).unwrap()] { + let decoded = content_range_from_value_with(value, mode).unwrap(); + assert_eq!(decoded.unit().unwrap(), unit); + assert_eq!(&decoded.value.as_bytes()[decoded.payload.clone()], payload); + assert_eq!(decoded.value.as_bytes(), wire.as_bytes()); + } + } + } + + #[test] + fn constructors_and_accessors_cover_byte_and_extension_forms() { + let bytes = ContentRangeOwned::bytes(0, 9, CompleteLength::Known(10)).expect("satisfied range"); + assert_eq!(bytes.unit(), Ok("bytes")); + assert_eq!( + bytes.byte_range(), + Some(ByteContentRange::Satisfied { + first: 0, + last: 9, + complete_length: Some(10), + }) + ); + assert_eq!(bytes.extension_payload(), None); + assert_eq!(bytes.as_field_value().as_bytes(), b"bytes 0-9/10"); + assert_eq!(bytes.clone().into_field_value().as_bytes(), b"bytes 0-9/10"); + assert!(format!("{bytes:?}").contains("byte_range")); + assert_eq!( + ContentRangeOwned::try_from("bytes 0-1/2") + .expect("valid borrowed content range") + .byte_range(), + Some(ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: Some(2), + }) + ); + + let open_length = ContentRangeOwned::bytes_range(0..10, CompleteLength::Unknown).expect("unknown length"); + assert_eq!(open_length.as_field_value().as_bytes(), b"bytes 0-9/*"); + let unsatisfied = ContentRangeOwned::unsatisfied_bytes(42).expect("unsatisfied range"); + assert_eq!( + unsatisfied.byte_range(), + Some(ByteContentRange::Unsatisfied { complete_length: 42 }) + ); + + let extension = ContentRangeOwned::extension("items", "first-last/complete").expect("extension"); + assert_eq!(extension.unit(), Ok("items")); + assert_eq!(extension.byte_range(), None); + assert_eq!(extension.extension_payload(), Some(b"first-last/complete".as_slice())); + assert_eq!( + ContentRangeOwned::extension("bad unit", "payload") + .expect_err("invalid unit") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + ContentRangeOwned::extension("items", "\u{7f}").expect_err("invalid payload").kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn range_bounds_and_complete_lengths_reject_empty_or_impossible_ranges() { + for result in [ + ContentRangeOwned::bytes(9, 0, CompleteLength::Unknown), + ContentRangeOwned::bytes(0, 9, CompleteLength::Known(9)), + ContentRangeOwned::bytes_range((Bound::Excluded(u64::MAX), Bound::Unbounded), CompleteLength::Unknown), + ContentRangeOwned::bytes_range((Bound::Included(0), Bound::Excluded(0)), CompleteLength::Unknown), + ContentRangeOwned::bytes_range((Bound::Unbounded, Bound::Included(1)), CompleteLength::Unknown), + ContentRangeOwned::bytes_range((Bound::Included(0), Bound::Unbounded), CompleteLength::Unknown), + ] { + assert_eq!(result.expect_err("invalid content range").kind(), DecodeErrorKind::InvalidSyntax); + } + assert!(validate_satisfied_content_range(0, 0, Some(1)).is_ok()); + assert!(validate_satisfied_content_range(1, 0, None).is_err()); + } + + #[test] + fn strict_parsing_borrows_units_and_rejects_malformed_byte_ranges() { + let wire = b"bytes 0-1/2"; + let parsed = parse_content_range(wire).expect("strict byte range"); + assert_eq!(parsed.unit, "bytes"); + assert_eq!(parsed.unit.as_ptr(), wire.as_ptr()); + assert_eq!(parsed.payload, b"0-1/2"); + + assert_eq!( + parse_content_range(b"Bytes */10").expect("case-insensitive bytes unit").byte_range, + Some(ByteContentRange::Unsatisfied { complete_length: 10 }) + ); + assert_eq!(parse_content_range(b"items opaque").expect("extension range").byte_range, None); + + for malformed in [ + b"bytes".as_slice(), + b"bytes 0-1", + b"bytes 0-1/2/3", + b"bytes 01/2", + b"bytes 0--1/2", + b"bytes 2-1/3", + b"bytes 0-2/2", + b"bytes 0-x/2", + b"bytes 0-1/x", + b"bytes */*", + b"bad/unit payload", + b"items \x7f", + ] { + assert!(parse_content_range(malformed).is_err(), "{malformed:?}"); + } + } + + #[test] + fn relaxed_parsing_accepts_only_whitespace_deviations() { + let parsed = parse_content_range_with(b"Bytes 0 - 9 / 10", DecodeMode::Relaxed).expect("relaxed byte whitespace"); + assert_eq!( + parsed.byte_range, + Some(ByteContentRange::Satisfied { + first: 0, + last: 9, + complete_length: Some(10), + }) + ); + assert_eq!( + parse_byte_content_range_relaxed(b"* / 10").expect("relaxed unsatisfied form"), + ByteContentRange::Unsatisfied { complete_length: 10 } + ); + assert!(parse_content_range_with(b"bytes 0-1/2", DecodeMode::Relaxed).is_err()); + assert!(parse_content_range_with(b"bytes", DecodeMode::Relaxed).is_err()); + assert!(parse_byte_content_range_relaxed(b"0-1/2/3").is_err()); + assert!(parse_byte_content_range_relaxed(b"0--1/2").is_err()); + assert!(parse_byte_content_range_relaxed(b"01/2").is_err()); + assert!(parse_byte_content_range_relaxed(b"2-1/3").is_err()); + assert!(parse_byte_content_range_relaxed(b"0-x/2").is_err()); + assert!(parse_byte_content_range_relaxed(b"x-1/2").is_err()); + assert!(parse_byte_content_range_relaxed(b"0-1/x").is_err()); + assert!(parse_byte_content_range_relaxed(b"*/x").is_err()); + assert!(parse_byte_content_range_relaxed(b"0-1").is_err()); + assert_eq!( + parse_byte_content_range_relaxed(b"0-1/*").expect("unknown complete length"), + ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: None, + } + ); + + assert_eq!( + parse_content_range_relaxed(b"items opaque") + .expect("direct relaxed extension") + .byte_range, + None + ); + assert_eq!( + parse_content_range_relaxed(b"items ") + .expect("empty extension payload is preserved") + .byte_range, + None + ); + assert!(parse_content_range_relaxed(b"items \x7f").is_err()); + assert!(parse_content_range_relaxed(b"items \xff").is_err()); + assert!(parse_content_range_relaxed(b"bad/unit payload").is_err()); + assert!(parse_content_range_relaxed(b"bytes 0-x/2").is_err()); + + let strict_fast_path = + parse_content_range_with(b"bytes 0-1/2", DecodeMode::Relaxed).expect("strict syntax remains valid in relaxed mode"); + assert_eq!(strict_fast_path.unit, "bytes"); + let extension = parse_content_range_with(b"items opaque", DecodeMode::Relaxed).expect("relaxed extension"); + assert_eq!(extension.byte_range, None); + assert!(parse_content_range_with(b"items \x7f", DecodeMode::Relaxed).is_err()); + } + + #[test] + fn owned_and_borrowed_decoding_preserve_payloads() { + let value = FieldValue::from_static("items opaque payload"); + let view = ::decode_view(value.as_field_value_ref()).expect("borrowed extension"); + assert_eq!(view.unit(), "items"); + assert_eq!(view.byte_range(), None); + assert_eq!(view.extension_payload(), Some(b"opaque payload".as_slice())); + assert_eq!(view.as_field_value().as_bytes(), b"items opaque payload"); + assert!(format!("{view:?}").contains("items")); + + let bytes = FieldValue::from_static("bytes 0-1/2"); + let bytes_view = ::decode_view(bytes.as_field_value_ref()).expect("borrowed byte content range"); + assert_eq!(bytes_view.extension_payload(), None); + assert!(::decode_view(FieldValueRef::new(b"invalid")).is_err()); + assert!(::decode_view_with(FieldValueRef::new(b"invalid"), DecodeMode::Relaxed,).is_err()); + assert!(::decode_owned_with(FieldValue::from_static("invalid"), DecodeMode::Relaxed,).is_err()); + + let owned = ::decode_owned(value).expect("owned extension"); + assert_eq!(owned.extension_payload(), Some(b"opaque payload".as_slice())); + assert_eq!( + ::as_field_value(&owned).as_bytes(), + b"items opaque payload" + ); + assert_eq!( + ::into_field_value(owned).as_bytes(), + b"items opaque payload" + ); + assert_eq!( + ContentRangeOwned::try_from(String::from("line\nbreak")) + .expect_err("invalid field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ContentRangeOwned::try_from("line\nbreak") + .expect_err("invalid borrowed field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + let owned = ContentRangeOwned::try_from(String::from("bytes 0-1/2")).expect("owned string conversion"); + assert_eq!( + owned.byte_range(), + Some(ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: Some(2), + }) + ); + let relaxed = + ::decode_owned_with(FieldValue::from_static("bytes 0 - 1 / 2"), DecodeMode::Relaxed) + .expect("relaxed owned decode"); + assert_eq!( + relaxed.byte_range(), + Some(ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: Some(2), + }) + ); + assert_eq!( + ::decode_view_with(FieldValueRef::new(b"bytes 0 - 1 / 2"), DecodeMode::Relaxed,) + .expect("relaxed view") + .byte_range(), + Some(ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: Some(2), + }) + ); + assert_eq!(::name(), &FieldName::ContentRange); + assert!(valid_content_range_extension(b"")); + assert!(!valid_content_range_extension(b"\x7f")); + assert!(parse_byte_content_range(b"0-1/*").is_ok()); + } + + #[test] + fn defensive_owned_accessors_revalidate_private_ranges() { + let out_of_bounds = ContentRangeOwned { + value: FieldValue::from_static("items payload"), + unit: 0..99, + payload: 6..13, + parsed: None, + }; + assert_eq!( + out_of_bounds.unit().expect_err("unit range is outside storage").kind(), + DecodeErrorKind::InvalidSyntax + ); + + let invalid_utf8 = FieldValue::from_bytes(b"\xff payload").expect("obs-text field value"); + let invalid_utf8 = ContentRangeOwned { + value: invalid_utf8, + unit: 0..1, + payload: 2..9, + parsed: None, + }; + assert_eq!( + invalid_utf8.unit().expect_err("unit bytes are not UTF-8").kind(), + DecodeErrorKind::InvalidSyntax + ); + + let missing_payload = ContentRangeOwned { + value: FieldValue::from_static("items payload"), + unit: 0..5, + payload: 99..100, + parsed: None, + }; + assert_eq!(missing_payload.extension_payload(), None); + } +} diff --git a/crates/http_headers/src/headers/range/mod.rs b/crates/http_headers/src/headers/range/mod.rs new file mode 100644 index 000000000..4c65b43c9 --- /dev/null +++ b/crates/http_headers/src/headers/range/mod.rs @@ -0,0 +1,20 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Byte range request, response, and capability headers. + +mod accept_ranges; +mod content_range; +#[expect( + clippy::module_inception, + reason = "the file is named after the `Range` header it defines, matching this family's one-file-per-header-concept convention" +)] +mod range; +mod shared; + +#[doc(inline)] +pub use accept_ranges::{AcceptRanges, AcceptRangesOwned, AcceptRangesView}; +#[doc(inline)] +pub use content_range::{ByteContentRange, CompleteLength, ContentRange, ContentRangeOwned, ContentRangeView}; +#[doc(inline)] +pub use range::{ByteRangeSpec, Range, RangeOwned, RangeView}; diff --git a/crates/http_headers/src/headers/range/range.rs b/crates/http_headers/src/headers/range/range.rs new file mode 100644 index 000000000..b04c3d902 --- /dev/null +++ b/crates/http_headers/src/headers/range/range.rs @@ -0,0 +1,1346 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt::Write as _; +use std::ops::{Bound, RangeBounds}; +use std::{fmt, str}; + +use http_headers_simd::ascii_str; + +use super::super::{invalid_syntax, trim_ows}; +use super::shared::{parse_number, validate_range_unit_for}; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, SingleValueField, validate}; + +/// One byte-range specification from a `Range` field. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{ByteRangeSpec, RangeOwned}; +/// +/// let value = RangeOwned::try_from("bytes=0-9, 20-, -5")?; +/// let specs = value +/// .byte_ranges() +/// .expect("byte ranges") +/// .collect::>(); +/// assert_eq!( +/// specs, +/// [ +/// ByteRangeSpec::FromTo { first: 0, last: 9 }, +/// ByteRangeSpec::From { first: 20 }, +/// ByteRangeSpec::Suffix { length: 5 }, +/// ] +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub enum ByteRangeSpec { + /// A closed inclusive range. + /// + /// Direct construction can represent `last < first`; encoding through + /// [`RangeOwned::bytes`] rejects that state. + FromTo { + /// The first byte position. + first: u64, + /// The last byte position. + last: u64, + }, + /// A range from a position through the end of the representation. + From { + /// The first byte position. + first: u64, + }, + /// The final `length` bytes of the representation. + Suffix { + /// The requested suffix length. + length: u64, + }, +} + +fn parse_range_with(bytes: &[u8], mode: crate::DecodeMode) -> Result, DecodeError> { + if mode == crate::DecodeMode::Strict { + return parse_range(bytes); + } + if let Ok(parsed) = parse_range(bytes) { + return Ok(parsed); + } + parse_range_relaxed(bytes) +} + +fn parse_range_relaxed(bytes: &[u8]) -> Result, DecodeError> { + let Some(separator) = separator_index(bytes) else { + return Err(invalid_syntax(&FieldName::Range)); + }; + let unit_bytes = trim_ows(&bytes[..separator]); + let payload = trim_ows(&bytes[separator + 1..]); + validate_range_unit_for(unit_bytes, &FieldName::Range)?; + let unit = ascii_str(unit_bytes).expect("validated range units are ASCII"); + let is_bytes = validate::eq_ignore_ascii_case(unit_bytes, b"bytes"); + if is_bytes { + validate_byte_range_set_relaxed(payload)?; + } else if !valid_extension_payload(payload) { + return Err(invalid_syntax(&FieldName::Range)); + } + Ok(ParsedRange { + unit, + payload, + bytes: is_bytes, + }) +} + +fn parse_byte_spec_relaxed(bytes: &[u8]) -> Result { + let bytes = trim_ows(bytes); + let Some(separator) = bytes.iter().position(|byte| *byte == b'-') else { + return Err(invalid_syntax(&FieldName::Range)); + }; + if bytes[separator + 1..].contains(&b'-') { + return Err(invalid_syntax(&FieldName::Range)); + } + let first = trim_ows(&bytes[..separator]); + let last = trim_ows(&bytes[separator + 1..]); + if first.is_empty() { + return parse_number(last, &FieldName::Range).map(|length| ByteRangeSpec::Suffix { length }); + } + let first = parse_number(first, &FieldName::Range)?; + if last.is_empty() { + Ok(ByteRangeSpec::From { first }) + } else { + let last = parse_number(last, &FieldName::Range)?; + ByteRangeSpec::from_range(first..=last) + } +} + +fn validate_byte_range_set_relaxed(bytes: &[u8]) -> Result<(), DecodeError> { + let mut count = 0_usize; + for item in bytes.split(|byte| *byte == b',') { + let item = trim_ows(item); + if item.is_empty() { + continue; + } + parse_byte_spec_relaxed(item)?; + count += 1; + } + if count == 0 { + Err(DecodeError::new(&FieldName::Range, DecodeErrorKind::MissingValue)) + } else { + Ok(()) + } +} + +impl ByteRangeSpec { + /// Constructs a closed inclusive range from explicit positions. + /// + /// Prefer [`Self::from_range`] when the caller already holds a Rust range. + /// + /// # Errors + /// + /// Returns an error when `last` precedes `first`. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ByteRangeSpec, RangeOwned}; + /// + /// let spec = ByteRangeSpec::from_to(0, 99)?; + /// assert_eq!(spec, ByteRangeSpec::FromTo { first: 0, last: 99 }); + /// + /// let value = RangeOwned::bytes([spec])?; + /// assert_eq!(value.as_field_value().as_bytes(), b"bytes=0-99"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_to(first: u64, last: u64) -> Result { + Self::from_range(first..=last) + } + + /// Constructs a bounded or open-ended byte range. + /// + /// Included bounds map directly to HTTP's inclusive positions. An + /// excluded start is incremented and an excluded end is decremented, so + /// both `0..100` and `0..=99` represent bytes 0 through 99. An unbounded + /// end constructs an open-ended range. + /// + /// # Errors + /// + /// Returns an error for an unbounded start, an empty or inverted range, or + /// a bound adjustment that overflows. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ByteRangeSpec; + /// + /// let half_open = ByteRangeSpec::from_range(0..100)?; + /// let inclusive = ByteRangeSpec::from_range(0..=99)?; + /// assert_eq!(half_open, ByteRangeSpec::FromTo { first: 0, last: 99 }); + /// assert_eq!(half_open, inclusive); + /// + /// let open_ended = ByteRangeSpec::from_range(100..)?; + /// assert_eq!(open_ended, ByteRangeSpec::From { first: 100 }); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_range(range: impl RangeBounds) -> Result { + Self::from_bounds(range.start_bound().cloned(), range.end_bound().cloned()) + } + + fn from_bounds(start: Bound, end: Bound) -> Result { + let first = match start { + Bound::Included(first) => first, + Bound::Excluded(first) => first.checked_add(1).ok_or_else(|| invalid_syntax(&FieldName::Range))?, + Bound::Unbounded => return Err(invalid_syntax(&FieldName::Range)), + }; + let last = match end { + Bound::Included(last) => Some(last), + Bound::Excluded(last) => Some(last.checked_sub(1).ok_or_else(|| invalid_syntax(&FieldName::Range))?), + Bound::Unbounded => None, + }; + match last { + Some(last) if last >= first => Ok(Self::FromTo { first, last }), + Some(_last) => Err(invalid_syntax(&FieldName::Range)), + None => Ok(Self::From { first }), + } + } + + /// Constructs an open-ended range. + #[must_use] + #[deprecated(since = "0.1.0", note = "use ByteRangeSpec::starting_at")] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::RangeOwned::try_from("bytes=0-99")?; + /// assert!(value.is_bytes()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn from(first: u64) -> Self { + Self::starting_at(first) + } + + /// Constructs an open-ended range without resembling `From::from`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ByteRangeSpec, RangeOwned}; + /// + /// let spec = ByteRangeSpec::starting_at(100); + /// assert_eq!(spec, ByteRangeSpec::From { first: 100 }); + /// + /// let value = RangeOwned::bytes([spec])?; + /// assert_eq!(value.as_field_value().as_bytes(), b"bytes=100-"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn starting_at(first: u64) -> Self { + Self::From { first } + } + + /// Constructs a suffix range. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ByteRangeSpec, RangeOwned}; + /// + /// let spec = ByteRangeSpec::suffix(500); + /// assert_eq!(spec, ByteRangeSpec::Suffix { length: 500 }); + /// + /// let value = RangeOwned::bytes([spec])?; + /// assert_eq!(value.as_field_value().as_bytes(), b"bytes=-500"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn suffix(length: u64) -> Self { + Self::Suffix { length } + } +} + +/// Defines the `Range` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 14.2](https://www.rfc-editor.org/rfc/rfc9110#section-14.2). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{Range, RangeOwned}; +/// +/// let mut map = HeaderMap::new(); +/// Range::insert(&mut map, RangeOwned::try_from("bytes=0-99")?)?; +/// assert!(Range::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct Range { + _private: (), +} + +/// Owned value for the `Range` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 14.2]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::RangeOwned::try_from("bytes=0-99")?; +/// assert!(value.is_bytes()); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Range: bytes=0-499` is closed, `Range: bytes=500-` is open-ended, and +/// `Range: bytes=-500` is a suffix range. Multiple ranges can be sent as +/// `Range: bytes=0-499, 1000-1499`; extension units such as +/// `Range: custom=opaque-set` are preserved. +/// +/// [RFC 9110 section 14.2]: https://www.rfc-editor.org/rfc/rfc9110#section-14.2 +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct RangeOwned { + value: FieldValue, + bytes: bool, +} + +/// Borrowed value for the `Range` header. +#[derive(Clone, Copy, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{Range, RangeView}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let view: RangeView<'_> = Range::decode_view(FieldValueRef::new(b"bytes=0-99"))?; +/// assert_eq!(view.unit(), "bytes"); +/// assert!(view.is_bytes()); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct RangeView<'a> { + value: FieldValueRef<'a>, + unit: &'a str, + range_set: &'a [u8], + bytes: bool, +} + +impl fmt::Debug for RangeOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("RangeOwned") + .field("unit", &self.unit()) + .field("is_bytes", &self.bytes) + .finish_non_exhaustive() + } +} + +impl fmt::Debug for RangeView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("RangeView") + .field("unit", &self.unit) + .field("is_bytes", &self.bytes) + .finish_non_exhaustive() + } +} + +impl RangeOwned { + /// Constructs a canonical byte range-set. + /// + /// # Errors + /// + /// Returns an error for an empty set or an inverted closed range. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::RangeOwned::try_from("bytes=0-99")?; + /// assert!(value.is_bytes()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn bytes(specs: I) -> Result + where + I: IntoIterator, + { + let specs = specs.into_iter(); + let mut wire = String::with_capacity(6_usize.saturating_add(specs.size_hint().0.saturating_mul(8))); + wire.push_str("bytes="); + let mut count = 0_usize; + for spec in specs { + if count != 0 { + wire.push_str(", "); + } + append_byte_spec(&mut wire, spec)?; + count += 1; + } + if count == 0 { + return Err(DecodeError::new(&FieldName::Range, DecodeErrorKind::MissingValue)); + } + let value = validated_byte_range_value(wire); + Ok(Self { value, bytes: true }) + } + + /// Constructs an extension range-unit and opaque range-set. + /// + /// The extension payload is preserved; this type does not claim to + /// normalize extension range semantics. + /// + /// # Errors + /// + /// Returns an error when the unit is not a token, is the reserved + /// case-insensitive `bytes` unit, or the payload is not nonempty visible + /// ASCII. Use [`Self::bytes`] for byte ranges. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::RangeOwned; + /// + /// let value = RangeOwned::extension("items", "1-5")?; + /// assert_eq!(value.unit()?, "items"); + /// assert!(!value.is_bytes()); + /// assert_eq!(value.extension_range_set(), Some(&b"1-5"[..])); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn extension(unit: impl AsRef, range_set: impl AsRef) -> Result { + let unit = unit.as_ref(); + let range_set = range_set.as_ref(); + validate_range_unit_for(unit.as_bytes(), &FieldName::Range)?; + if validate::eq_ignore_ascii_case(unit.as_bytes(), b"bytes") { + return Err(invalid_syntax(&FieldName::Range)); + } + if !valid_extension_payload(range_set.as_bytes()) { + return Err(invalid_syntax(&FieldName::Range)); + } + let mut wire = String::with_capacity(unit.len() + 1 + range_set.len()); + wire.push_str(unit); + wire.push('='); + wire.push_str(range_set); + let value = validated_extension_range_value(wire); + Ok(Self { value, bytes: false }) + } + + /// Returns the range unit exactly as received. + /// + /// # Errors + /// + /// Returns an error if stored metadata does not match the wire value. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::RangeOwned; + /// + /// let bytes = RangeOwned::try_from("bytes=0-99")?; + /// assert_eq!(bytes.unit()?, "bytes"); + /// + /// let items = RangeOwned::extension("items", "1-5")?; + /// assert_eq!(items.unit()?, "items"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn unit(&self) -> Result<&str, DecodeError> { + let bytes = self.value.as_bytes(); + let unit = separator_index(bytes) + .and_then(|separator| bytes.get(..separator)) + .ok_or_else(|| invalid_syntax(&FieldName::Range))?; + str::from_utf8(trim_ows(unit)).map_err(|_invalid| invalid_syntax(&FieldName::Range)) + } + + /// Returns whether the range unit is `bytes`, case-insensitively. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::RangeOwned::try_from("bytes=0-99")?; + /// assert!(value.is_bytes()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn is_bytes(&self) -> bool { + self.bytes + } + + /// Iterates byte range specifications, or returns `None` for an extension unit. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ByteRangeSpec, RangeOwned}; + /// + /// let value = RangeOwned::try_from("bytes=0-49, 100-, -500")?; + /// let specs = value + /// .byte_ranges() + /// .expect("byte ranges") + /// .collect::>(); + /// assert_eq!( + /// specs, + /// [ + /// ByteRangeSpec::FromTo { first: 0, last: 49 }, + /// ByteRangeSpec::From { first: 100 }, + /// ByteRangeSpec::Suffix { length: 500 }, + /// ] + /// ); + /// + /// let extension = RangeOwned::extension("items", "1-5")?; + /// assert!(extension.byte_ranges().is_none()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn byte_ranges(&self) -> Option + '_> { + self.range_set().map(ByteRangeIter::new).filter(|_iterator| self.bytes) + } + + /// Returns an extension range-set without normalization. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::RangeOwned; + /// + /// let extension = RangeOwned::extension("items", "1-5")?; + /// assert_eq!(extension.extension_range_set(), Some(&b"1-5"[..])); + /// + /// let bytes = RangeOwned::try_from("bytes=0-99")?; + /// assert_eq!(bytes.extension_range_set(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn extension_range_set(&self) -> Option<&[u8]> { + if self.bytes { None } else { self.range_set() } + } + + #[inline] + fn range_set(&self) -> Option<&[u8]> { + let bytes = self.value.as_bytes(); + bytes.get(separator_index(bytes)?.saturating_add(1)..).map(trim_ows) + } + + /// Returns the preserved field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::RangeOwned; + /// + /// let value = RangeOwned::try_from("bytes=0-99")?; + /// assert_eq!(value.as_field_value().as_bytes(), b"bytes=0-99"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Consumes the header and returns its field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::RangeOwned; + /// + /// let value = RangeOwned::try_from("bytes=-500")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"bytes=-500"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(RangeOwned, |value| value.value); + +impl<'a> RangeView<'a> { + /// Returns the range unit exactly as received. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::Range; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = Range::decode_view(FieldValueRef::new(b"items=1-5"))?; + /// assert_eq!(view.unit(), "items"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn unit(self) -> &'a str { + self.unit + } + + /// Returns whether the range unit is `bytes`, case-insensitively. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::RangeOwned::try_from("bytes=0-99")?; + /// assert!(value.is_bytes()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn is_bytes(self) -> bool { + self.bytes + } + + /// Iterates byte range specifications, or returns `None` for an extension unit. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ByteRangeSpec, Range}; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = Range::decode_view(FieldValueRef::new(b"bytes=0-49, 100-, -500"))?; + /// let specs = view.byte_ranges().expect("byte ranges").collect::>(); + /// assert_eq!( + /// specs, + /// [ + /// ByteRangeSpec::FromTo { first: 0, last: 49 }, + /// ByteRangeSpec::From { first: 100 }, + /// ByteRangeSpec::Suffix { length: 500 }, + /// ] + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn byte_ranges(self) -> Option + 'a> { + self.bytes.then(|| ByteRangeIter::new(self.range_set)) + } + + /// Returns the extension range-set bytes without normalization. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::Range; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let extension = Range::decode_view(FieldValueRef::new(b"items=1-5"))?; + /// assert_eq!(extension.extension_range_set(), Some(&b"1-5"[..])); + /// + /// let bytes = Range::decode_view(FieldValueRef::new(b"bytes=0-99"))?; + /// assert_eq!(bytes.extension_range_set(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn extension_range_set(self) -> Option<&'a [u8]> { + if self.bytes { None } else { Some(self.range_set) } + } + + /// Returns the borrowed field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::Range; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = Range::decode_view(FieldValueRef::new(b"bytes=100-"))?; + /// assert_eq!(view.as_field_value().as_bytes(), b"bytes=100-"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } +} + +impl SingleValueField for Range { + type View<'a> = RangeView<'a>; + type Owned = RangeOwned; + + fn name() -> &'static FieldName { + &FieldName::Range + } + + #[inline] + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + let parsed = parse_range(value.as_bytes())?; + Ok(RangeView { + value, + unit: parsed.unit, + range_set: parsed.payload, + bytes: parsed.bytes, + }) + } + + #[expect( + clippy::inline_always, + reason = "measured: Criterion otherwise outlines this conversion while Callgrind inlines it" + )] + #[inline(always)] + fn decode_owned(value: FieldValue) -> Result { + RangeOwned::try_from(value) + } + + fn decode_view_with(value: FieldValueRef<'_>, mode: crate::DecodeMode) -> Result, DecodeError> { + let parsed = parse_range_with(value.as_bytes(), mode)?; + Ok(RangeView { + value, + unit: parsed.unit, + range_set: parsed.payload, + bytes: parsed.bytes, + }) + } + + fn decode_owned_with(value: FieldValue, mode: crate::DecodeMode) -> Result { + let bytes = parse_range_with(value.as_bytes(), mode)?.bytes; + Ok(RangeOwned { value, bytes }) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +super::super::shared::impl_string_conversions!(RangeOwned, &FieldName::Range, invalid_syntax, wire); + +impl TryFrom for RangeOwned { + type Error = DecodeError; + + #[expect( + clippy::inline_always, + reason = "measured: outlining this conversion leaves large Result traffic in Criterion" + )] + #[inline(always)] + fn try_from(value: FieldValue) -> Result { + let bytes = parse_range(value.as_bytes())?.bytes; + Ok(Self { value, bytes }) + } +} + +/// Locates the `=` that separates the range unit from the range set. +fn separator_index(bytes: &[u8]) -> Option { + bytes.iter().position(|byte| *byte == b'=') +} + +pub(super) struct ParsedRange<'a> { + pub(super) unit: &'a str, + pub(super) payload: &'a [u8], + pub(super) bytes: bool, +} + +#[expect( + clippy::inline_always, + reason = "measured: Callgrind inlines this parser while Criterion otherwise emits a real call" +)] +#[inline(always)] +pub(super) fn parse_range(bytes: &[u8]) -> Result, DecodeError> { + if let Some(payload) = bytes.strip_prefix(b"bytes=") { + if !http_headers_simd::scan_byte_range_set(bytes, 6) { + validate_byte_range_set_wide(bytes, 6)?; + } + return Ok(ParsedRange { + unit: "bytes", + payload, + bytes: true, + }); + } + parse_range_slow(bytes) +} + +#[cold] +#[inline(never)] +pub(super) fn parse_range_slow(bytes: &[u8]) -> Result, DecodeError> { + let Some(separator) = bytes.iter().position(|byte| *byte == b'=') else { + return Err(invalid_syntax(&FieldName::Range)); + }; + let unit_bytes = &bytes[..separator]; + let payload = &bytes[separator + 1..]; + validate_range_unit_for(unit_bytes, &FieldName::Range)?; + let unit = ascii_str(unit_bytes).expect("validated range units are ASCII"); + let is_bytes = validate::eq_ignore_ascii_case(unit_bytes, b"bytes"); + if is_bytes { + validate_byte_range_set(bytes, separator + 1)?; + } else if !valid_extension_payload(payload) { + return Err(invalid_syntax(&FieldName::Range)); + } + Ok(ParsedRange { + unit, + payload, + bytes: is_bytes, + }) +} + +struct ByteRangeIter<'a> { + remaining: &'a [u8], +} + +impl<'a> ByteRangeIter<'a> { + const fn new(remaining: &'a [u8]) -> Self { + Self { remaining } + } +} + +impl Iterator for ByteRangeIter<'_> { + type Item = ByteRangeSpec; + + fn next(&mut self) -> Option { + loop { + if self.remaining.is_empty() { + return None; + } + let (item, remaining) = if let Some(separator) = self.remaining.iter().position(|byte| *byte == b',') { + (&self.remaining[..separator], &self.remaining[separator + 1..]) + } else { + (self.remaining, &[] as &[u8]) + }; + self.remaining = remaining; + let item = trim_ows(item); + if item.is_empty() { + continue; + } + return parse_byte_spec_relaxed(item).ok(); + } + } +} + +#[inline] +pub(super) fn validate_byte_range_set(bytes: &[u8], start: usize) -> Result<(), DecodeError> { + // The accelerated scanner settles any set that fits one classification + // window; the word-at-a-time scanner covers the longer ones. + if http_headers_simd::scan_byte_range_set(bytes, start) { + return Ok(()); + } + validate_byte_range_set_wide(bytes, start) +} + +/// Validates a range-set that the one-window scanner could not settle. +#[cold] +#[inline(never)] +fn validate_byte_range_set_wide(bytes: &[u8], start: usize) -> Result<(), DecodeError> { + if scan_byte_range_set(bytes, start) { + return Ok(()); + } + validate_byte_range_set_slow(bytes.get(start..).unwrap_or_default()) +} + +/// Marks every byte of `word` that is not an ASCII digit. +/// +/// The low seven bits are biased so that only digits stay below `0x80`, and +/// the original word is folded back in so that `obs-text` bytes never borrow +/// into their neighbor. +#[inline] +const fn nondigit_bits(word: u64) -> u64 { + // Repeated-byte SWAR masks: clear high bits, align ASCII zero, then bias + // values outside `0..=9` into each lane's high bit. + const LOW: u64 = 0x7f7f_7f7f_7f7f_7f7f; + const ZEROS: u64 = 0x3030_3030_3030_3030; + const BIAS: u64 = 0x7676_7676_7676_7676; + + (((word & LOW) ^ ZEROS).wrapping_add(BIAS) | word) & !LOW +} + +/// Counts the leading ASCII digits marked by `nondigit`, saturating at eight. +#[inline] +const fn digit_run(nondigit: u64) -> usize { + (nondigit.trailing_zeros() >> 3) as usize +} + +/// Reads the eight bytes at `at`, zero-padded past the end of `bytes`. +/// +/// Callers guarantee `8 <= bytes.len()` and `at < bytes.len()`, so a window +/// that would overrun is taken from the end of `bytes` and shifted into place. +#[inline] +fn word_at(bytes: &[u8], at: usize) -> u64 { + if let Some(chunk) = bytes.get(at..at.saturating_add(8)) { + return u64::from_le_bytes(<[u8; 8]>::try_from(chunk).unwrap_or_default()); + } + let offset = bytes.len().saturating_sub(8); + let tail = <[u8; 8]>::try_from(bytes.get(offset..).unwrap_or_default()).unwrap_or_default(); + u64::from_le_bytes(tail) >> ((at.wrapping_sub(offset) & 7) * 8) +} + +/// Validates the range-set of `bytes` that starts at `start`. +/// +/// Returns `false` for anything unusual so the general implementation can +/// decide between acceptance and the precise error it reports. +#[expect( + clippy::cast_possible_truncation, + reason = "item lengths and single bytes are extracted from a machine word" +)] +#[inline] +fn scan_byte_range_set(bytes: &[u8], start: usize) -> bool { + if bytes.len() < 8 { + return false; + } + let mut at = start; + loop { + if at >= bytes.len() { + return false; + } + let word = word_at(bytes, at); + let Some(item) = scan_byte_range_spec(word) else { + return false; + }; + let end = at + item; + if end == bytes.len() { + return true; + } + + // The delimiter fits the word, but following OWS may lie beyond it. + let rest = word.rotate_right((item as u32) << 3); + if rest as u8 != b',' { + return false; + } + at = end + 1; + if matches!(bytes.get(at), Some(b' ' | b'\t')) { + at += 1; + } + } +} + +/// Returns the length of the byte-range-spec that starts `word`. +/// +/// Specs whose digits could continue past the eight byte window, that carry a +/// leading zero on a compared position, or that are inverted return `None`. +#[expect( + clippy::cast_possible_truncation, + reason = "digit runs never exceed eight and single bytes are extracted from a word" +)] +#[inline] +fn scan_byte_range_spec(word: u64) -> Option { + /// Longest digit run that always leaves room for a terminator. + const MAX_RUN: usize = 6; + + let nondigit = nondigit_bits(word); + let first_len = digit_run(nondigit); + if first_len == 0 { + if word as u8 != b'-' { + return None; + } + let suffix_len = digit_run(nondigit.rotate_right(8)); + if suffix_len == 0 || suffix_len > MAX_RUN { + return None; + } + return Some(1 + suffix_len); + } + if first_len > MAX_RUN { + return None; + } + + // Rotating keeps the vacated bytes outside every window this function + // inspects, so an over-long run always trips the `MAX_RUN` guard. + let shift = (first_len as u32) << 3; + if word.rotate_right(shift) as u8 != b'-' { + return None; + } + let last = word.rotate_right(shift + 8); + let last_len = digit_run(nondigit.rotate_right(shift + 8)); + if first_len + last_len > MAX_RUN { + return None; + } + if last_len != 0 { + // A leading zero would let the shorter run compare as the smaller + // number even when it is not, so leave those to the general path. + if last_len > 1 && last as u8 == b'0' { + return None; + } + if first_len > last_len { + return None; + } + if first_len == last_len { + let mask = u64::MAX >> ((8 - first_len) * 8); + if (word & mask).swap_bytes() > (last & mask).swap_bytes() { + return None; + } + } + } + Some(first_len + 1 + last_len) +} + +#[cold] +#[inline(never)] +pub(super) fn validate_byte_range_set_slow(bytes: &[u8]) -> Result<(), DecodeError> { + let mut count = 0_usize; + for item in bytes.split(|byte| *byte == b',') { + let item = trim_ows(item); + if item.is_empty() { + continue; + } + let _spec = parse_byte_spec(item)?; + count += 1; + } + if count == 0 { + Err(DecodeError::new(&FieldName::Range, DecodeErrorKind::MissingValue)) + } else { + Ok(()) + } +} + +fn parse_byte_spec(bytes: &[u8]) -> Result { + if let Some(suffix) = bytes.strip_prefix(b"-") { + return parse_number(suffix, &FieldName::Range).map(|length| ByteRangeSpec::Suffix { length }); + } + let Some(separator) = bytes.iter().position(|byte| *byte == b'-') else { + return Err(invalid_syntax(&FieldName::Range)); + }; + if bytes[separator + 1..].contains(&b'-') { + return Err(invalid_syntax(&FieldName::Range)); + } + let first = parse_number(&bytes[..separator], &FieldName::Range)?; + let last = &bytes[separator + 1..]; + if last.is_empty() { + Ok(ByteRangeSpec::From { first }) + } else { + let last = parse_number(last, &FieldName::Range)?; + ByteRangeSpec::from_range(first..=last) + } +} + +fn append_byte_spec(wire: &mut String, spec: ByteRangeSpec) -> Result<(), DecodeError> { + match spec { + ByteRangeSpec::FromTo { first, last } => { + if last < first { + return Err(invalid_syntax(&FieldName::Range)); + } + write!(wire, "{first}").expect("writing to a String is infallible"); + wire.push('-'); + write!(wire, "{last}").expect("writing to a String is infallible"); + } + ByteRangeSpec::From { first } => { + write!(wire, "{first}").expect("writing to a String is infallible"); + wire.push('-'); + } + ByteRangeSpec::Suffix { length } => { + wire.push('-'); + write!(wire, "{length}").expect("writing to a String is infallible"); + } + } + Ok(()) +} + +fn validated_byte_range_value(wire: String) -> FieldValue { + FieldValue::try_from(wire).expect("formatted byte ranges form a valid field value") +} + +fn validated_extension_range_value(wire: String) -> FieldValue { + FieldValue::try_from(wire).expect("validated extension ranges form a valid field value") +} + +fn valid_extension_payload(bytes: &[u8]) -> bool { + !bytes.is_empty() && bytes.iter().all(|byte| matches!(byte, 0x21..=0x7e)) +} + +#[cfg(test)] +#[expect( + clippy::assertions_on_result_states, + reason = "the tests classify many parser outcomes without needing their success values" +)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::ops::Bound; + + use super::{ + ByteRangeIter, ByteRangeSpec, Range, RangeOwned, digit_run, nondigit_bits, parse_byte_spec, parse_byte_spec_relaxed, parse_range, + parse_range_with, scan_byte_range_set, scan_byte_range_spec, valid_extension_payload, validate_byte_range_set_relaxed, + validate_byte_range_set_slow, word_at, + }; + use crate::{DecodeErrorKind, DecodeMode, FieldName, FieldValue, FieldValueRef, SingleValueField}; + + fn word(bytes: &[u8]) -> u64 { + u64::from_le_bytes(bytes.try_into().expect("test words contain eight bytes")) + } + + #[test] + fn byte_range_bounds_cover_closed_open_suffix_and_error_cases() { + assert_eq!( + ByteRangeSpec::from_to(0, 9).expect("closed range"), + ByteRangeSpec::FromTo { first: 0, last: 9 } + ); + assert_eq!( + ByteRangeSpec::from_range(0..10).expect("exclusive end"), + ByteRangeSpec::FromTo { first: 0, last: 9 } + ); + assert_eq!(ByteRangeSpec::from_range(5..).expect("open end"), ByteRangeSpec::From { first: 5 }); + assert_eq!( + ByteRangeSpec::from_range((Bound::Excluded(4), Bound::Included(9))).expect("excluded start"), + ByteRangeSpec::FromTo { first: 5, last: 9 } + ); + assert_eq!(ByteRangeSpec::starting_at(7), ByteRangeSpec::From { first: 7 }); + assert_eq!(ByteRangeSpec::suffix(12), ByteRangeSpec::Suffix { length: 12 }); + + for result in [ + ByteRangeSpec::from_to(9, 0), + ByteRangeSpec::from_range((Bound::Unbounded, Bound::Included(1))), + ByteRangeSpec::from_range((Bound::Excluded(u64::MAX), Bound::Unbounded)), + ByteRangeSpec::from_range((Bound::Included(0), Bound::Excluded(0))), + ] { + assert_eq!(result.expect_err("invalid bounds").kind(), DecodeErrorKind::InvalidSyntax); + } + } + + #[test] + #[expect(deprecated, reason = "the compatibility constructor remains covered")] + fn deprecated_open_range_constructor_delegates_to_starting_at() { + assert_eq!(ByteRangeSpec::from(7), ByteRangeSpec::starting_at(7)); + } + + #[test] + fn owned_range_constructors_and_accessors_preserve_semantics() { + let owned = RangeOwned::bytes(vec![ + ByteRangeSpec::FromTo { first: 0, last: 9 }, + ByteRangeSpec::From { first: 20 }, + ByteRangeSpec::Suffix { length: 5 }, + ]) + .expect("valid byte ranges"); + assert_eq!(owned.unit(), Ok("bytes")); + assert!(owned.is_bytes()); + assert_eq!( + owned.byte_ranges().expect("byte iterator").collect::>(), + [ + ByteRangeSpec::FromTo { first: 0, last: 9 }, + ByteRangeSpec::From { first: 20 }, + ByteRangeSpec::Suffix { length: 5 }, + ] + ); + assert_eq!(owned.extension_range_set(), None); + assert_eq!(owned.as_field_value().as_bytes(), b"bytes=0-9, 20-, -5"); + assert_eq!(owned.clone().into_field_value().as_bytes(), b"bytes=0-9, 20-, -5"); + assert!(format!("{owned:?}").contains("is_bytes: true")); + + assert_eq!( + RangeOwned::bytes(Vec::::new()).expect_err("empty set").kind(), + DecodeErrorKind::MissingValue + ); + assert_eq!( + RangeOwned::bytes(vec![ByteRangeSpec::FromTo { first: 2, last: 1 }]) + .expect_err("inverted range") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let extension = RangeOwned::extension("items", "opaque-set").expect("extension range"); + assert_eq!(extension.unit(), Ok("items")); + assert!(!extension.is_bytes()); + assert!(extension.byte_ranges().is_none()); + assert_eq!(extension.extension_range_set(), Some(b"opaque-set".as_slice())); + assert_eq!( + RangeOwned::extension("bad unit", "opaque").expect_err("invalid unit").kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + RangeOwned::extension("items", "").expect_err("empty payload").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + RangeOwned::try_from(String::from("line\nbreak")) + .expect_err("invalid field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn strict_and_relaxed_parsers_cover_byte_and_extension_paths() { + assert!(RangeOwned::try_from("bytes=0-1").expect("valid borrowed range").is_bytes()); + let strict = parse_range(b"bytes=0-9, 20-, -5").expect("strict byte set"); + assert_eq!(strict.unit, "bytes"); + assert_eq!(strict.payload, b"0-9, 20-, -5"); + assert!(strict.bytes); + + let extension = parse_range(b"items=opaque").expect("extension range"); + assert_eq!(extension.unit, "items"); + assert!(!extension.bytes); + assert!(parse_range(b"missing-separator").is_err()); + assert!(parse_range(b"bad unit=opaque").is_err()); + assert!(parse_range(b"items=").is_err()); + assert!(parse_range(b"bytes=").is_err()); + assert!(parse_range(b"bytes=1--2").is_err()); + assert!(parse_range(b"bytes=2-1").is_err()); + let strict_mode = parse_range_with(b"bytes=0-1", DecodeMode::Strict).expect("explicit strict mode"); + assert!(strict_mode.bytes); + let relaxed_fast_path = parse_range_with(b"bytes=0-1", DecodeMode::Relaxed).expect("strict syntax remains valid in relaxed mode"); + assert!(relaxed_fast_path.bytes); + + let relaxed = parse_range_with(b" Bytes = 0 - 9 , - 5 ", DecodeMode::Relaxed).expect("relaxed whitespace"); + assert_eq!(relaxed.unit, "Bytes"); + assert_eq!(relaxed.payload, b"0 - 9 , - 5"); + assert!(relaxed.bytes); + assert!(parse_range_with(b" = 0-1", DecodeMode::Relaxed).is_err()); + assert!(parse_range_with(b"bytes 0-1", DecodeMode::Relaxed).is_err()); + assert!(parse_range_with(b"bytes = , ,", DecodeMode::Relaxed).is_err()); + assert!(parse_range_with(b"items = ", DecodeMode::Relaxed).is_err()); + assert!(validate_byte_range_set_relaxed(b"0 - 1, - 2, 3 -").is_ok()); + assert!(validate_byte_range_set_relaxed(b"x - 1").is_err()); + assert!(parse_byte_spec_relaxed(b"x - 1").is_err()); + assert!(parse_byte_spec_relaxed(b"1 - x").is_err()); + assert!(::decode_view(FieldValueRef::new(b"invalid")).is_err()); + assert!(::decode_view_with(FieldValueRef::new(b"invalid"), DecodeMode::Relaxed,).is_err()); + assert!(::decode_owned_with(FieldValue::from_static("invalid"), DecodeMode::Relaxed,).is_err()); + } + + #[test] + fn byte_spec_parsers_and_iterator_skip_empty_members_and_stop_at_invalid_ones() { + assert_eq!(parse_byte_spec(b"-10").expect("suffix"), ByteRangeSpec::Suffix { length: 10 }); + assert_eq!(parse_byte_spec(b"10-").expect("open range"), ByteRangeSpec::From { first: 10 }); + assert_eq!( + parse_byte_spec(b"10-20").expect("closed range"), + ByteRangeSpec::FromTo { first: 10, last: 20 } + ); + for malformed in [b"10".as_slice(), b"1--2", b"-", b"2-1"] { + assert!(parse_byte_spec(malformed).is_err(), "{malformed:?}"); + } + assert_eq!( + parse_byte_spec_relaxed(b" 10 - 20 ").expect("relaxed closed range"), + ByteRangeSpec::FromTo { first: 10, last: 20 } + ); + assert_eq!( + parse_byte_spec_relaxed(b" - 5 ").expect("relaxed suffix"), + ByteRangeSpec::Suffix { length: 5 } + ); + assert!(parse_byte_spec_relaxed(b"1--2").is_err()); + assert!(parse_byte_spec_relaxed(b"not-a-range").is_err()); + + assert_eq!( + ByteRangeIter::new(b", 0-1, , -5, 9-, invalid").collect::>(), + [ + ByteRangeSpec::FromTo { first: 0, last: 1 }, + ByteRangeSpec::Suffix { length: 5 }, + ByteRangeSpec::From { first: 9 }, + ] + ); + } + + #[test] + fn borrowed_views_cover_byte_and_extension_accessors() { + let byte_view = ::decode_view(FieldValueRef::new(b"bytes=0-1")).expect("borrowed byte range"); + assert_eq!(byte_view.unit(), "bytes"); + assert!(byte_view.is_bytes()); + assert_eq!( + byte_view.byte_ranges().expect("byte ranges").collect::>(), + [ByteRangeSpec::FromTo { first: 0, last: 1 }] + ); + assert_eq!(byte_view.extension_range_set(), None); + assert_eq!(byte_view.as_field_value().as_bytes(), b"bytes=0-1"); + assert!(format!("{byte_view:?}").contains("is_bytes: true")); + + let value = FieldValue::from_static("items=opaque"); + let extension = ::decode_owned(value).expect("owned extension"); + assert_eq!(extension.extension_range_set(), Some(b"opaque".as_slice())); + let extension_view = ::decode_view_with(FieldValueRef::new(b" items = opaque "), DecodeMode::Relaxed) + .expect("relaxed extension"); + assert_eq!(extension_view.unit(), "items"); + assert_eq!(extension_view.extension_range_set(), Some(b"opaque".as_slice())); + assert!(extension_view.byte_ranges().is_none()); + + let decoded = ::decode_owned(FieldValue::from_static("bytes=0-1")).expect("owned byte range"); + assert_eq!(::as_field_value(&decoded).as_bytes(), b"bytes=0-1"); + assert_eq!(::into_field_value(decoded).as_bytes(), b"bytes=0-1"); + let relaxed_owned = ::decode_owned_with(FieldValue::from_static("bytes = 0 - 1"), DecodeMode::Relaxed) + .expect("relaxed owned range"); + assert!(relaxed_owned.is_bytes()); + assert_eq!( + RangeOwned::try_from("line\nbreak") + .expect_err("invalid borrowed field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + let from_string = RangeOwned::try_from(String::from("bytes=0-1")).expect("owned string conversion"); + assert!(from_string.is_bytes()); + assert_eq!(::name(), &FieldName::Range); + } + + #[test] + fn defensive_owned_accessors_revalidate_private_wire_storage() { + let missing_separator = RangeOwned { + value: FieldValue::from_static("bytes"), + bytes: true, + }; + assert_eq!( + missing_separator.unit().expect_err("unit separator is missing").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert!(missing_separator.byte_ranges().is_none()); + + let invalid_utf8 = FieldValue::from_bytes(b"\xff=opaque").expect("obs-text field value"); + let invalid_utf8 = RangeOwned { + value: invalid_utf8, + bytes: false, + }; + assert_eq!( + invalid_utf8.unit().expect_err("unit is not UTF-8").kind(), + DecodeErrorKind::InvalidSyntax + ); + + assert_eq!( + ::decode_owned(FieldValue::from_static("bad range")) + .expect_err("invalid owned range") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn range_unit_projection_preserves_opaque_payloads_and_rejects_unicode_units() { + for (wire, mode, unit, payload) in [ + ( + b"items=first:last".as_slice(), + DecodeMode::Strict, + "items", + b"first:last".as_slice(), + ), + (b" Items = first:last ", DecodeMode::Relaxed, "Items", b"first:last"), + (b"Bytes=0-9", DecodeMode::Strict, "Bytes", b"0-9"), + (b" BYTES = 0 - 9 ", DecodeMode::Relaxed, "BYTES", b"0 - 9"), + ] { + let parsed = parse_range_with(wire, mode).unwrap(); + assert_eq!(parsed.unit, unit); + assert_eq!(parsed.payload, payload); + } + for wire in [ + b"it\xc3\xa9ms=opaque".as_slice(), + b"\xff=opaque", + b"Byt\x80s=0-9", + b"items=\xff\x80", + ] { + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert!(parse_range_with(wire, mode).is_err(), "{wire:?}"); + } + } + } + + #[test] + fn word_scanner_reads_ows_after_full_windows() { + for wire in [ + "100-109, 200-209", + "100-109,\t200-209", + "100-109,200-209", + "-123456, 200-209", + "123456-, 200-209", + ] { + assert!(scan_byte_range_set(wire.as_bytes(), 0), "{wire}"); + validate_byte_range_set_slow(wire.as_bytes()).unwrap(); + } + for separator in u8::MIN..=u8::MAX { + let mut wire = b"100-109,".to_vec(); + wire.push(separator); + wire.extend_from_slice(b"200-209"); + if scan_byte_range_set(&wire, 0) { + assert!(validate_byte_range_set_slow(&wire).is_ok(), "{wire:?}"); + } + } + for wire in ["100-109, \t200-209", "100-109, ", "100-109, 300-299", "100-109,\t"] { + assert!(!scan_byte_range_set(wire.as_bytes(), 0), "{wire}"); + } + } + + #[test] + fn word_scanner_helpers_cover_window_boundaries_and_rejections() { + let digits = nondigit_bits(word(b"12345678")); + assert_eq!(digit_run(digits), 8); + assert_eq!(digit_run(nondigit_bits(word(b"12-45678"))), 2); + + assert_eq!(word_at(b"12345678", 0), word(b"12345678")); + assert_eq!(word_at(b"123456789", 7) & 0xff, u64::from(b'8')); + assert_eq!(scan_byte_range_spec(word(b"-12,rest")), Some(3)); + assert_eq!(scan_byte_range_spec(word(b"12-,rest")), Some(3)); + assert_eq!(scan_byte_range_spec(word(b"12-34,re")), Some(5)); + for rejected in [ + word(b"x2-3,res"), + word(b"-x,rest!"), + word(b"-1234567"), + word(b"1234567-"), + word(b"12x34,re"), + word(b"12-034,r"), + word(b"123-12,r"), + word(b"34-12,re"), + ] { + assert_eq!(scan_byte_range_spec(rejected), None); + } + assert!(scan_byte_range_set(b"0-1, 2-3", 0)); + assert!(!scan_byte_range_set(b"short", 0)); + assert!(!scan_byte_range_set(b"0-1; 2-3", 0)); + assert!(valid_extension_payload(b"!opaque~")); + assert!(!valid_extension_payload(b"")); + assert!(!valid_extension_payload(b"has space")); + } +} diff --git a/crates/http_headers/src/headers/range/shared.rs b/crates/http_headers/src/headers/range/shared.rs new file mode 100644 index 000000000..ae76b7270 --- /dev/null +++ b/crates/http_headers/src/headers/range/shared.rs @@ -0,0 +1,126 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use crate::{DecodeError, DecodeErrorKind, FieldName, validate}; + +pub(super) fn validate_range_unit_for(bytes: &[u8], name: &'static FieldName) -> Result<(), DecodeError> { + if validate::token(bytes) { + Ok(()) + } else { + Err(DecodeError::new(name, DecodeErrorKind::InvalidToken)) + } +} + +#[inline] +pub(super) fn parse_number(bytes: &[u8], name: &'static FieldName) -> Result { + validate::decimal_u64(bytes).ok_or_else(|| DecodeError::new(name, DecodeErrorKind::InvalidNumber)) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::super::{accept_ranges, range}; + + #[test] + fn fast_scanners_agree_with_the_general_implementations() { + const ALPHABET: &[u8] = b"01-, a=\t9"; + const RANGE_ALPHABET: &[u8] = b"01-, "; + + let mut payload = Vec::new(); + // Coprime strides keep every alphabet symbol represented at every position under Miri. + for length in 0..=4_u32 { + let step = if cfg!(miri) && length > 2 { 11 } else { 1 }; + for encoded in (0..ALPHABET.len().pow(length)).step_by(step) { + payload.clear(); + let mut encoded = encoded; + for _position in 0..length { + payload.push(ALPHABET[encoded % ALPHABET.len()]); + encoded /= ALPHABET.len(); + } + + assert_eq!( + range::validate_byte_range_set(&payload, 0).map_err(|error| error.kind()), + range::validate_byte_range_set_slow(&payload).map_err(|error| error.kind()), + "byte range set {:?}", + String::from_utf8_lossy(&payload) + ); + + let fast = accept_ranges::scan_units(&payload).map(|(units, none)| { + accept_ranges::validate_none_cardinality(units, none) + .map(|()| none) + .map_err(|error| error.kind()) + }); + if let Some(fast) = fast { + assert_eq!( + fast, + accept_ranges::validate_units_slow(&payload).map_err(|error| error.kind()), + "unit list {:?}", + String::from_utf8_lossy(&payload) + ); + } + + let mut wire = b"bytes=".to_vec(); + wire.extend_from_slice(&payload); + for wire in [payload.as_slice(), wire.as_slice()] { + let fast = range::parse_range(wire) + .map(|value| (value.unit.as_bytes(), value.payload, value.bytes)) + .map_err(|error| error.kind()); + let slow = range::parse_range_slow(wire) + .map(|value| (value.unit.as_bytes(), value.payload, value.bytes)) + .map_err(|error| error.kind()); + assert_eq!(fast, slow, "range {:?}", String::from_utf8_lossy(wire)); + } + } + } + + // Longer payloads exercise the multi-item paths of the word scanner, + // which the short exhaustive sweep above never reaches. + for length in 5..=7_u32 { + let step = if cfg!(miri) { 61 } else { 1 }; + for encoded in (0..RANGE_ALPHABET.len().pow(length)).step_by(step) { + payload.clear(); + payload.extend_from_slice(b"bytes="); + let mut encoded = encoded; + for _position in 0..length { + payload.push(RANGE_ALPHABET[encoded % RANGE_ALPHABET.len()]); + encoded /= RANGE_ALPHABET.len(); + } + + let fast = range::parse_range(&payload) + .map(|value| (value.unit.as_bytes(), value.payload, value.bytes)) + .map_err(|error| error.kind()); + let slow = range::parse_range_slow(&payload) + .map(|value| (value.unit.as_bytes(), value.payload, value.bytes)) + .map_err(|error| error.kind()); + assert_eq!(fast, slow, "range {:?}", String::from_utf8_lossy(&payload)); + + // The accelerated scanner is only ever allowed to settle a set + // the general implementation also accepts. + if http_headers_simd::scan_byte_range_set(&payload, 6) { + assert!( + range::validate_byte_range_set_slow(&payload[6..]).is_ok(), + "accelerated scanner accepted {:?}", + String::from_utf8_lossy(&payload) + ); + } + } + } + + for payload in [ + "9999999999999999999-", + "18446744073709551615-", + "18446744073709551616-", + "0-18446744073709551615", + "0-18446744073709551616", + "00000000000000000000005-6", + "-18446744073709551616", + "0-1, 2-3 ,4-", + ] { + assert_eq!( + range::validate_byte_range_set(payload.as_bytes(), 0).map_err(|error| error.kind()), + range::validate_byte_range_set_slow(payload.as_bytes()).map_err(|error| error.kind()), + "byte range set {payload:?}" + ); + } + } +} diff --git a/crates/http_headers/src/headers/security/content_security_policy.rs b/crates/http_headers/src/headers/security/content_security_policy.rs new file mode 100644 index 000000000..18f4afa50 --- /dev/null +++ b/crates/http_headers/src/headers/security/content_security_policy.rs @@ -0,0 +1,480 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::str; + +use super::super::shared::FieldLinesIter; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef}; + +/// Defines the `Content-Security-Policy` header. +/// +/// # Specification +/// +/// Defined by [Content Security Policy Level 3 section 7.1](https://www.w3.org/TR/CSP3/#csp-header). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{ContentSecurityPolicy, ContentSecurityPolicyOwned}; +/// +/// let mut map = HeaderMap::new(); +/// ContentSecurityPolicy::insert( +/// &mut map, +/// ContentSecurityPolicyOwned::new("default-src 'self'")?, +/// )?; +/// assert!(ContentSecurityPolicy::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct ContentSecurityPolicy { + _private: (), +} + +/// Owned value for the `Content-Security-Policy` header. +/// +/// # Specification +/// +/// Defined by [Content Security Policy Level 3 section 7.1]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::ContentSecurityPolicyOwned::new("default-src 'self'")?; +/// assert_eq!(value.policies().count(), 1); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Content-Security-Policy: default-src 'self'` sets a default policy. +/// `Content-Security-Policy: default-src 'self'; script-src 'nonce-abc123'` +/// adds a script-specific source list. Multiple field lines are independent +/// policies and are preserved separately. +/// +/// [Content Security Policy Level 3 section 7.1]: https://www.w3.org/TR/CSP3/#csp-header +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct ContentSecurityPolicyOwned { + values: FieldLinesIter, +} + +/// Borrowed value for the `Content-Security-Policy` header. +/// # Examples +/// +/// ``` +/// use http_headers::headers::{ContentSecurityPolicy, ContentSecurityPolicyView}; +/// use http_headers::source::{FieldLines, FieldSource}; +/// use http_headers::{Field, FieldName}; +/// +/// struct Source; +/// impl FieldSource for Source { +/// fn lines(&self, name: &'static FieldName) -> Option> { +/// (name == &FieldName::ContentSecurityPolicy) +/// .then(|| FieldLines::single(name, b"default-src 'self'; img-src *")) +/// } +/// } +/// +/// let view: ContentSecurityPolicyView<'_> = +/// ContentSecurityPolicy::view(&Source)?.expect("header is present"); +/// let policies = view.policy_strs().collect::, _>>()?; +/// assert_eq!(policies, ["default-src 'self'; img-src *"]); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct ContentSecurityPolicyView<'a> { + values: FieldLines<'a>, +} + +super::super::shared::impl_value_count_debug!( + ContentSecurityPolicyOwned => "ContentSecurityPolicyOwned", + ContentSecurityPolicyView<'_> => "ContentSecurityPolicyView", +); + +impl ContentSecurityPolicyOwned { + /// Constructs one opaque policy field value. + /// + /// # Errors + /// + /// Returns an error when `policy` is not a safe HTTP field value. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentSecurityPolicyOwned; + /// + /// let value = ContentSecurityPolicyOwned::new("default-src 'self'")?; + /// let policies = value.policy_strs().collect::, _>>()?; + /// assert_eq!(policies, ["default-src 'self'"]); + /// + /// assert!(ContentSecurityPolicyOwned::new("default-src\nscript-src").is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn new(policy: impl AsRef) -> Result { + let policy = policy.as_ref(); + let value = FieldValue::from_str(policy).map_err(|_invalid| super::super::invalid_syntax(&FieldName::ContentSecurityPolicy))?; + Ok(Self::from_field_value(value)) + } + + /// Constructs one opaque policy from arbitrary safe field-value bytes. + /// + /// # Errors + /// + /// Returns an error for controls, DEL, CR, or LF. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentSecurityPolicyOwned; + /// + /// let value = ContentSecurityPolicyOwned::from_bytes(b"default-src 'self'; img-src *")?; + /// assert_eq!( + /// value.policies().next(), + /// Some(b"default-src 'self'; img-src *".as_slice()), + /// ); + /// + /// assert!(ContentSecurityPolicyOwned::from_bytes(b"default-src\x7f").is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_bytes(policy: impl AsRef<[u8]>) -> Result { + let policy = policy.as_ref(); + let value = FieldValue::from_bytes(policy).map_err(|_invalid| super::super::invalid_syntax(&FieldName::ContentSecurityPolicy))?; + Ok(Self::from_field_value(value)) + } + + /// Adds another independently enforced policy field line. + /// + /// # Errors + /// + /// Returns an error when `policy` is not a safe HTTP field value. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentSecurityPolicyOwned; + /// + /// let value = + /// ContentSecurityPolicyOwned::new("default-src 'self'")?.with_policy("script-src 'none'")?; + /// let policies = value.policy_strs().collect::, _>>()?; + /// assert_eq!(policies, ["default-src 'self'", "script-src 'none'"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn with_policy(mut self, policy: &str) -> Result { + let value = FieldValue::from_str(policy).map_err(|_invalid| super::super::invalid_syntax(&FieldName::ContentSecurityPolicy))?; + self.values.push(value); + Ok(self) + } + + /// Iterates raw policy bytes in field-line order. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentSecurityPolicyOwned; + /// + /// let value = ContentSecurityPolicyOwned::new("default-src 'self'")?.with_policy("img-src *")?; + /// let policies = value.policies().collect::>(); + /// assert_eq!( + /// policies, + /// [b"default-src 'self'".as_slice(), b"img-src *".as_slice()], + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn policies(&self) -> impl Iterator { + self.values.iter().map(FieldValue::as_bytes) + } + + /// Iterates policy text as UTF-8. + /// + /// # Errors + /// + /// An item contains [`DecodeErrorKind::InvalidUtf8`] when that field line + /// is not UTF-8. Iteration resumes with the following field line. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentSecurityPolicyOwned; + /// + /// let value = ContentSecurityPolicyOwned::from_bytes(b"default-src 'self'")?; + /// let policies = value.policy_strs().collect::, _>>()?; + /// assert_eq!(policies, ["default-src 'self'"]); + /// + /// let opaque = ContentSecurityPolicyOwned::from_bytes(b"\xff")?; + /// assert!(opaque.policy_strs().next().expect("one policy").is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn policy_strs(&self) -> impl Iterator> { + self.policies().map(policy_str) + } + + #[cfg(all(feature = "serde", feature = "headers-security"))] + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.iter().map(FieldValue::as_field_value_ref) + } + + fn from_field_value(value: FieldValue) -> Self { + Self { + values: FieldLinesIter::one(value), + } + } +} + +impl ContentSecurityPolicyView<'_> { + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.repeated() + } + + /// Iterates raw policy bytes in field-line order. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentSecurityPolicy; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source; + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::ContentSecurityPolicy) + /// .then(|| FieldLines::single(name, b"default-src 'self'")) + /// } + /// } + /// + /// let view = ContentSecurityPolicy::view(&Source)?.expect("header is present"); + /// assert_eq!( + /// view.policies().collect::>(), + /// [b"default-src 'self'".as_slice()], + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn policies(&self) -> impl Iterator { + self.values.repeated().map(FieldValueRef::as_bytes) + } + + /// Iterates policy text as UTF-8. + /// + /// # Errors + /// + /// An item contains [`DecodeErrorKind::InvalidUtf8`] when that field line + /// is not UTF-8. Iteration resumes with the following field line. + /// # Examples + /// + /// ``` + /// use http_headers::headers::ContentSecurityPolicy; + /// use http_headers::source::{FieldLines, FieldSource}; + /// use http_headers::{Field, FieldName}; + /// + /// struct Source; + /// impl FieldSource for Source { + /// fn lines(&self, name: &'static FieldName) -> Option> { + /// (name == &FieldName::ContentSecurityPolicy) + /// .then(|| FieldLines::single(name, b"default-src 'self'; img-src *")) + /// } + /// } + /// + /// let view = ContentSecurityPolicy::view(&Source)?.expect("header is present"); + /// let policies = view.policy_strs().collect::, _>>()?; + /// assert_eq!(policies, ["default-src 'self'; img-src *"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn policy_strs(&self) -> impl Iterator> { + self.policies().map(policy_str) + } +} + +impl Field for ContentSecurityPolicy { + type View<'a> = ContentSecurityPolicyView<'a>; + type Owned = ContentSecurityPolicyOwned; + + fn name() -> &'static FieldName { + &FieldName::ContentSecurityPolicy + } + + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(values) = source.lines(Self::name()) else { + return Ok(None); + }; + values.validate_custom_source()?; + Ok(Some(ContentSecurityPolicyView { values })) + } + + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + let mut copied = FieldLinesIter::empty(); + for (_, owned) in lines.repeated_owned()? { + copied.push(owned); + } + Ok(Some(ContentSecurityPolicyOwned { values: copied })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), value.values.into_encoded()) + } +} + +impl TryFrom<&str> for ContentSecurityPolicyOwned { + type Error = DecodeError; + + fn try_from(value: &str) -> Result { + Self::new(value) + } +} + +impl TryFrom for ContentSecurityPolicyOwned { + type Error = DecodeError; + + fn try_from(value: String) -> Result { + let value = FieldValue::try_from(value).map_err(|_invalid| super::super::invalid_syntax(&FieldName::ContentSecurityPolicy))?; + Self::try_from(value) + } +} + +impl TryFrom for ContentSecurityPolicyOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + Ok(Self::from_field_value(value)) + } +} + +fn policy_str(bytes: &[u8]) -> Result<&str, DecodeError> { + str::from_utf8(bytes).map_err(|_invalid| DecodeError::new(&FieldName::ContentSecurityPolicy, DecodeErrorKind::InvalidUtf8)) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{ContentSecurityPolicy, ContentSecurityPolicyOwned}; + use crate::sink::{EncodedValues, FieldSink}; + use crate::source::FieldSource; + use crate::{DecodeErrorKind, Field, FieldValue, TestSink}; + + #[test] + fn constructors_preserve_policy_lines_and_reject_unsafe_bytes() { + let policy = ContentSecurityPolicyOwned::new("default-src 'self'") + .expect("valid policy") + .with_policy("script-src 'none'") + .expect("valid second policy"); + assert_eq!( + policy.policies().collect::>(), + [b"default-src 'self'".as_slice(), b"script-src 'none'".as_slice()] + ); + assert_eq!( + policy.policy_strs().collect::, _>>(), + Ok(vec!["default-src 'self'", "script-src 'none'"]) + ); + assert_eq!(format!("{policy:?}"), "ContentSecurityPolicyOwned { value_count: 2 }"); + + let bytes = ContentSecurityPolicyOwned::from_bytes(b"sandbox").expect("safe bytes are accepted"); + assert_eq!(bytes.policies().next(), Some(b"sandbox".as_slice())); + assert_eq!( + ContentSecurityPolicyOwned::new("default-src\r\nscript-src") + .expect_err("line breaks are unsafe") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ContentSecurityPolicyOwned::from_bytes(b"default-src\x7f") + .expect_err("DEL is unsafe") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ContentSecurityPolicyOwned::new("default-src 'self'") + .expect("valid policy") + .with_policy("script-src\n") + .expect_err("unsafe appended policy") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn conversions_and_utf8_errors_are_reported_per_policy() { + let borrowed = ContentSecurityPolicyOwned::try_from("default-src *").expect("borrowed string converts"); + let owned = ContentSecurityPolicyOwned::try_from(String::from("sandbox")).expect("owned string converts"); + let field = FieldValue::from_static("upgrade-insecure-requests"); + let from_field = ContentSecurityPolicyOwned::try_from(field).expect("field converts"); + assert_eq!(borrowed.policy_strs().next(), Some(Ok("default-src *"))); + assert_eq!(owned.policy_strs().next(), Some(Ok("sandbox"))); + assert_eq!(from_field.policy_strs().next(), Some(Ok("upgrade-insecure-requests"))); + assert_eq!( + ContentSecurityPolicyOwned::try_from(String::from("sandbox\n")) + .expect_err("invalid field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let opaque = FieldValue::from_bytes([0xff]).expect("obs-text is a safe field value"); + let policy = ContentSecurityPolicyOwned::try_from(opaque).expect("opaque bytes are allowed"); + assert_eq!( + policy + .policy_strs() + .next() + .expect("one policy") + .expect_err("policy is not UTF-8") + .kind(), + DecodeErrorKind::InvalidUtf8 + ); + } + + #[test] + fn header_round_trip_covers_views_owned_values_and_absence() { + let mut table = TestSink::new(); + assert!(ContentSecurityPolicy::view(&table).expect("absent view succeeds").is_none()); + assert!( + ContentSecurityPolicy::owned(&table) + .expect("absent owned decode succeeds") + .is_none() + ); + + table + .set_values( + ContentSecurityPolicy::name(), + EncodedValues::from_vec(vec![ + FieldValue::from_static("default-src 'self'"), + FieldValue::from_static("script-src 'none'"), + ]), + ) + .expect("table accepts policies"); + let view = ContentSecurityPolicy::view(&table) + .expect("view decodes") + .expect("header is present"); + assert_eq!( + view.policies().collect::>(), + [b"default-src 'self'".as_slice(), b"script-src 'none'".as_slice()] + ); + assert_eq!( + view.policy_strs().collect::, _>>(), + Ok(vec!["default-src 'self'", "script-src 'none'",]) + ); + assert_eq!(view.field_values().count(), 2); + assert_eq!(format!("{view:?}"), "ContentSecurityPolicyView { value_count: 2 }"); + + let owned = ContentSecurityPolicy::owned(&table) + .expect("owned value decodes") + .expect("header is present"); + let mut output = TestSink::new(); + ContentSecurityPolicy::insert(&mut output, owned).expect("owned value inserts"); + assert_eq!( + output + .lines(ContentSecurityPolicy::name()) + .expect("inserted values") + .repeated() + .map(crate::FieldValueRef::as_bytes) + .collect::>(), + [b"default-src 'self'".as_slice(), b"script-src 'none'".as_slice()] + ); + } +} diff --git a/crates/http_headers/src/headers/security/mod.rs b/crates/http_headers/src/headers/security/mod.rs new file mode 100644 index 000000000..11942f7e1 --- /dev/null +++ b/crates/http_headers/src/headers/security/mod.rs @@ -0,0 +1,20 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Browser security policy and transport-security headers. + +mod content_security_policy; +mod referrer_policy; +mod strict_transport_security; +mod x_content_type_options; + +#[doc(inline)] +pub use content_security_policy::{ContentSecurityPolicy, ContentSecurityPolicyOwned, ContentSecurityPolicyView}; +#[doc(inline)] +pub use referrer_policy::{ReferrerPolicy, ReferrerPolicyOwned, ReferrerPolicyTokenView, ReferrerPolicyValue, ReferrerPolicyView}; +#[doc(inline)] +pub use strict_transport_security::{ + HstsDirectiveView, StrictTransportSecurity, StrictTransportSecurityBuilder, StrictTransportSecurityOwned, StrictTransportSecurityView, +}; +#[doc(inline)] +pub use x_content_type_options::{XContentTypeOptions, XContentTypeOptionsOwned, XContentTypeOptionsView}; diff --git a/crates/http_headers/src/headers/security/referrer_policy.rs b/crates/http_headers/src/headers/security/referrer_policy.rs new file mode 100644 index 000000000..0437ff7a0 --- /dev/null +++ b/crates/http_headers/src/headers/security/referrer_policy.rs @@ -0,0 +1,909 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::{fmt, str}; + +use super::super::shared::FieldLinesIter; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// A recognized Referrer Policy token. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +#[non_exhaustive] +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::ReferrerPolicyOwned::new( +/// http_headers::headers::ReferrerPolicyValue::NoReferrer, +/// ); +/// assert_eq!( +/// value.preferred()?, +/// http_headers::headers::ReferrerPolicyValue::NoReferrer +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub enum ReferrerPolicyValue { + /// Omits the `Referer` header. + NoReferrer, + /// Sends full referrers except on a secure-to-insecure downgrade. + NoReferrerWhenDowngrade, + /// Sends only the origin. + Origin, + /// Sends a full same-origin referrer and only the origin cross-origin. + OriginWhenCrossOrigin, + /// Sends referrers only for same-origin requests. + SameOrigin, + /// Sends the origin except on a secure-to-insecure downgrade. + StrictOrigin, + /// Sends a full same-origin referrer, an origin cross-origin, and nothing + /// on a secure-to-insecure downgrade. + StrictOriginWhenCrossOrigin, + /// Sends the full referrer for all requests. + UnsafeUrl, +} + +/// One Referrer Policy token, including unrecognized extensions. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{ +/// ReferrerPolicyOwned, ReferrerPolicyTokenView, ReferrerPolicyValue, +/// }; +/// +/// let value = ReferrerPolicyOwned::try_from("future-policy, strict-origin")?; +/// let mut tokens = value.tokens(); +/// let extension: ReferrerPolicyTokenView<'_> = tokens.next().expect("extension token")?; +/// assert_eq!(extension.as_str(), "future-policy"); +/// assert_eq!(extension.policy(), None); +/// let recognized = tokens.next().expect("recognized token")?; +/// assert_eq!(recognized.policy(), Some(ReferrerPolicyValue::StrictOrigin)); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct ReferrerPolicyTokenView<'a> { + token: &'a str, + policy: Option, +} + +/// Defines the `Referrer-Policy` header. +/// +/// # Specification +/// +/// Defined by [Referrer Policy section 8.1](https://www.w3.org/TR/referrer-policy/#referrer-policy-header). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{ReferrerPolicy, ReferrerPolicyOwned, ReferrerPolicyValue}; +/// +/// let mut map = HeaderMap::new(); +/// ReferrerPolicy::insert( +/// &mut map, +/// ReferrerPolicyOwned::new(ReferrerPolicyValue::NoReferrer), +/// )?; +/// assert!(ReferrerPolicy::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct ReferrerPolicy { + _private: (), +} + +/// Owned value for the `Referrer-Policy` header. +/// +/// # Specification +/// +/// Defined by [Referrer Policy section 8.1]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::ReferrerPolicyOwned::new( +/// http_headers::headers::ReferrerPolicyValue::NoReferrer, +/// ); +/// assert_eq!( +/// value.preferred()?, +/// http_headers::headers::ReferrerPolicyValue::NoReferrer +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Referrer-Policy: no-referrer` suppresses the referrer. +/// `Referrer-Policy: no-referrer, strict-origin-when-cross-origin` demonstrates +/// fallback-list processing, where the last recognized token is effective. +/// +/// [Referrer Policy section 8.1]: https://www.w3.org/TR/referrer-policy/#referrer-policy-header +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct ReferrerPolicyOwned { + values: FieldLinesIter, +} + +/// Borrowed value for the `Referrer-Policy` header. +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), http_headers::DecodeError> { +/// use http::{HeaderMap, HeaderValue}; +/// use http_headers::Field; +/// use http_headers::headers::{ReferrerPolicy, ReferrerPolicyValue, ReferrerPolicyView}; +/// +/// let mut map = HeaderMap::new(); +/// map.insert( +/// "referrer-policy", +/// HeaderValue::from_static("future-policy, strict-origin"), +/// ); +/// let view: ReferrerPolicyView<'_> = ReferrerPolicy::view(&map)?.expect("header present"); +/// assert_eq!(view.preferred()?, ReferrerPolicyValue::StrictOrigin); +/// # Ok::<(), http_headers::DecodeError>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +pub struct ReferrerPolicyView<'a> { + values: FieldLines<'a>, +} + +super::super::shared::impl_value_count_debug!( + ReferrerPolicyOwned => "ReferrerPolicyOwned", + ReferrerPolicyView<'_> => "ReferrerPolicyView", +); + +impl fmt::Display for ReferrerPolicyOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + super::super::shared::fmt_ascii_values(self.values.iter().map(FieldValue::as_field_value_ref), f) + } +} + +impl ReferrerPolicyValue { + /// Returns the serialized policy token. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ReferrerPolicyValue; + /// + /// let policy = ReferrerPolicyValue::StrictOriginWhenCrossOrigin; + /// assert_eq!(policy.as_str(), "strict-origin-when-cross-origin"); + /// ``` + pub const fn as_str(self) -> &'static str { + match self { + Self::NoReferrer => "no-referrer", + Self::NoReferrerWhenDowngrade => "no-referrer-when-downgrade", + Self::Origin => "origin", + Self::OriginWhenCrossOrigin => "origin-when-cross-origin", + Self::SameOrigin => "same-origin", + Self::StrictOrigin => "strict-origin", + Self::StrictOriginWhenCrossOrigin => "strict-origin-when-cross-origin", + Self::UnsafeUrl => "unsafe-url", + } + } +} + +impl<'a> ReferrerPolicyTokenView<'a> { + /// Returns the original policy token. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ReferrerPolicyOwned; + /// + /// let value = ReferrerPolicyOwned::try_from("future-policy, strict-origin")?; + /// let token = value.tokens().next().expect("token present")?; + /// assert_eq!(token.as_str(), "future-policy"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_str(self) -> &'a str { + self.token + } + + /// Returns the recognized policy, or `None` for a future extension token. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ReferrerPolicyOwned, ReferrerPolicyValue}; + /// + /// let value = ReferrerPolicyOwned::try_from("future-policy, strict-origin")?; + /// let policies = value + /// .tokens() + /// .map(|token| token.map(|token| token.policy())) + /// .collect::, _>>()?; + /// assert_eq!(policies, [None, Some(ReferrerPolicyValue::StrictOrigin)]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn policy(self) -> Option { + self.policy + } +} + +impl ReferrerPolicyOwned { + /// Constructs a list containing one policy. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ReferrerPolicyOwned::new( + /// http_headers::headers::ReferrerPolicyValue::NoReferrer, + /// ); + /// assert_eq!( + /// value.preferred()?, + /// http_headers::headers::ReferrerPolicyValue::NoReferrer + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn new(policy: ReferrerPolicyValue) -> Self { + Self { + values: FieldLinesIter::one(FieldValue::from_static(policy.as_str())), + } + } + + /// Adds a fallback policy as another field line. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ReferrerPolicyOwned, ReferrerPolicyValue}; + /// + /// let value = ReferrerPolicyOwned::new(ReferrerPolicyValue::NoReferrer) + /// .with_fallback(ReferrerPolicyValue::StrictOriginWhenCrossOrigin); + /// assert_eq!( + /// value.preferred()?, + /// ReferrerPolicyValue::StrictOriginWhenCrossOrigin, + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn with_fallback(mut self, policy: ReferrerPolicyValue) -> Self { + self.values.push(FieldValue::from_static(policy.as_str())); + self + } + + /// Iterates all tokens, including unrecognized extensions, in wire order. + /// + /// # Errors + /// + /// An item contains [`DecodeErrorKind::InvalidToken`] when malformed + /// stored data is encountered. Iteration resumes with the following token. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::ReferrerPolicyOwned; + /// + /// let value = ReferrerPolicyOwned::try_from("no-referrer, strict-origin")?; + /// let tokens = value + /// .tokens() + /// .map(|token| token.map(|token| token.as_str())) + /// .collect::, _>>()?; + /// assert_eq!(tokens, ["no-referrer", "strict-origin"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn tokens(&self) -> impl Iterator, DecodeError>> { + self.values + .iter() + .flat_map(|value| CommaItems::new(value.as_bytes())) + .map(parse_referrer_policy_token) + } + + /// Iterates recognized policies in wire order. + /// + /// # Errors + /// + /// Items propagate the token errors documented by [`Self::tokens`]. + /// Iteration resumes with the following token. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ReferrerPolicyOwned, ReferrerPolicyValue}; + /// + /// let value = ReferrerPolicyOwned::try_from( + /// "future-policy, no-referrer, strict-origin-when-cross-origin", + /// )?; + /// let policies = value.policies().collect::, _>>()?; + /// assert_eq!( + /// policies, + /// [ + /// ReferrerPolicyValue::NoReferrer, + /// ReferrerPolicyValue::StrictOriginWhenCrossOrigin, + /// ], + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn policies(&self) -> impl Iterator> + '_ { + self.tokens().filter_map(recognized_policy) + } + + /// Returns the preferred policy from a fallback list. + /// + /// # Errors + /// + /// Returns an error if stored wire data is invalid or no recognized + /// policy is present. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ReferrerPolicyOwned::new( + /// http_headers::headers::ReferrerPolicyValue::NoReferrer, + /// ); + /// assert_eq!( + /// value.preferred()?, + /// http_headers::headers::ReferrerPolicyValue::NoReferrer + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn preferred(&self) -> Result { + self.policies() + .last() + .transpose()? + .ok_or_else(|| super::super::invalid_syntax(&FieldName::ReferrerPolicy)) + } + + #[cfg(all(feature = "serde", feature = "headers-security"))] + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.iter().map(FieldValue::as_field_value_ref) + } +} + +impl ReferrerPolicyView<'_> { + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.repeated() + } + + /// Iterates all tokens, including unrecognized extensions, in wire order. + /// + /// # Errors + /// + /// An item contains [`DecodeErrorKind::InvalidToken`] when malformed + /// stored data is encountered. Iteration resumes with the following token. + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::ReferrerPolicy; + /// + /// let mut map = HeaderMap::new(); + /// map.insert( + /// "referrer-policy", + /// HeaderValue::from_static("future-policy, no-referrer"), + /// ); + /// let view = ReferrerPolicy::view(&map)?.expect("header present"); + /// let tokens = view + /// .tokens() + /// .map(|token| token.map(|token| token.as_str())) + /// .collect::, _>>()?; + /// assert_eq!(tokens, ["future-policy", "no-referrer"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn tokens(&self) -> impl Iterator, DecodeError>> + '_ { + self.values.comma_items().map(|item| item.and_then(parse_referrer_policy_token)) + } + + /// Iterates recognized policies in wire order. + /// + /// # Errors + /// + /// Items propagate the token errors documented by [`Self::tokens`]. + /// Iteration resumes with the following token. + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::{HeaderMap, HeaderValue}; + /// use http_headers::Field; + /// use http_headers::headers::{ReferrerPolicy, ReferrerPolicyValue}; + /// + /// let mut map = HeaderMap::new(); + /// map.append( + /// "referrer-policy", + /// HeaderValue::from_static("future-policy, origin"), + /// ); + /// map.append("referrer-policy", HeaderValue::from_static("strict-origin")); + /// let view = ReferrerPolicy::view(&map)?.expect("header present"); + /// let policies = view.policies().collect::, _>>()?; + /// assert_eq!( + /// policies, + /// [ + /// ReferrerPolicyValue::Origin, + /// ReferrerPolicyValue::StrictOrigin, + /// ], + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn policies(&self) -> impl Iterator> + '_ { + self.tokens().filter_map(recognized_policy) + } + + /// Returns the preferred policy from a fallback list. + /// + /// # Errors + /// + /// Returns an error if stored wire data is invalid or no recognized + /// policy is present. + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::ReferrerPolicyOwned::new( + /// http_headers::headers::ReferrerPolicyValue::NoReferrer, + /// ); + /// assert_eq!( + /// value.preferred()?, + /// http_headers::headers::ReferrerPolicyValue::NoReferrer + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn preferred(&self) -> Result { + self.policies() + .last() + .transpose()? + .ok_or_else(|| super::super::invalid_syntax(&FieldName::ReferrerPolicy)) + } +} + +impl Field for ReferrerPolicy { + type View<'a> = ReferrerPolicyView<'a>; + type Owned = ReferrerPolicyOwned; + + fn name() -> &'static FieldName { + &FieldName::ReferrerPolicy + } + + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_custom_source()?; + let mut all_recognized = false; + for value in lines.repeated() { + if recognize_referrer_policy(super::super::trim_ows(value.as_bytes())).is_none() { + all_recognized = false; + break; + } + all_recognized = true; + } + if all_recognized { + return Ok(Some(ReferrerPolicyView { values: lines })); + } + + let mut count = 0_usize; + for item in lines.comma_items() { + parse_referrer_policy_token(item?)?; + count = increment_item_count(count)?; + } + if count == 0 { + return Err(DecodeError::new(&FieldName::ReferrerPolicy, DecodeErrorKind::MissingValue)); + } + Ok(Some(ReferrerPolicyView { values: lines })) + } + + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_custom_source()?; + let mut all_recognized = false; + for value in lines.repeated() { + if recognize_referrer_policy(super::super::trim_ows(value.as_bytes())).is_none() { + all_recognized = false; + break; + } + all_recognized = true; + } + let count = if all_recognized { + lines.len() + } else { + let mut count = 0_usize; + for item in lines.comma_items() { + parse_referrer_policy_token(item?)?; + count = increment_item_count(count)?; + } + count + }; + if count == 0 { + return Err(DecodeError::new(&FieldName::ReferrerPolicy, DecodeErrorKind::MissingValue)); + } + let mut copied = FieldLinesIter::empty(); + for (_value, owned) in lines.repeated_owned()? { + copied.push(owned); + } + Ok(Some(ReferrerPolicyOwned { values: copied })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), value.values.into_encoded()) + } +} + +super::super::shared::impl_string_conversions!(ReferrerPolicyOwned, &FieldName::ReferrerPolicy, super::super::invalid_syntax, value); + +impl TryFrom for ReferrerPolicyOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + if recognize_referrer_policy(super::super::trim_ows(value.as_bytes())).is_some() { + return Ok(Self { + values: FieldLinesIter::one(value), + }); + } + + let mut count = 0_usize; + for item in CommaItems::new(value.as_bytes()) { + parse_referrer_policy_token(item)?; + count = increment_item_count(count)?; + } + if count == 0 { + return Err(super::super::invalid_syntax(&FieldName::ReferrerPolicy)); + } + Ok(Self { + values: FieldLinesIter::one(value), + }) + } +} + +fn parse_referrer_policy_token(bytes: &[u8]) -> Result, DecodeError> { + if let Some(policy) = recognize_referrer_policy(bytes) { + return Ok(ReferrerPolicyTokenView { + token: policy.as_str(), + policy: Some(policy), + }); + } + if !validate::token(bytes) { + return Err(DecodeError::new(&FieldName::ReferrerPolicy, DecodeErrorKind::InvalidToken)); + } + let token = str::from_utf8(bytes).expect("HTTP token validation guarantees ASCII"); + Ok(ReferrerPolicyTokenView { token, policy: None }) +} + +fn increment_item_count(count: usize) -> Result { + count + .checked_add(1) + .ok_or_else(|| DecodeError::new(&FieldName::ReferrerPolicy, DecodeErrorKind::InvalidNumber)) +} + +fn recognize_referrer_policy(bytes: &[u8]) -> Option { + if bytes == b"strict-origin-when-cross-origin" { + return Some(ReferrerPolicyValue::StrictOriginWhenCrossOrigin); + } + match bytes.len() { + 11 => match bytes[0] { + b'n' if &bytes[1..3] == b"o-" && &bytes[3..] == b"referrer" => Some(ReferrerPolicyValue::NoReferrer), + b's' if bytes == b"same-origin" => Some(ReferrerPolicyValue::SameOrigin), + _ => None, + }, + 26 if bytes.starts_with(b"no-") && bytes == b"no-referrer-when-downgrade" => Some(ReferrerPolicyValue::NoReferrerWhenDowngrade), + 6 if bytes[0] == b'o' && bytes == b"origin" => Some(ReferrerPolicyValue::Origin), + 24 if bytes[0] == b'o' && bytes == b"origin-when-cross-origin" => Some(ReferrerPolicyValue::OriginWhenCrossOrigin), + 13 if bytes[0] == b's' && bytes == b"strict-origin" => Some(ReferrerPolicyValue::StrictOrigin), + 10 if bytes[0] == b'u' && bytes == b"unsafe-url" => Some(ReferrerPolicyValue::UnsafeUrl), + _ => None, + } +} + +fn recognized_policy(token: Result, DecodeError>) -> Option> { + match token { + Ok(token) => token.policy().map(Ok), + Err(error) => Some(Err(error)), + } +} + +struct CommaItems<'a> { + bytes: &'a [u8], + start: usize, + position: usize, + finished: bool, +} + +impl<'a> CommaItems<'a> { + const fn new(bytes: &'a [u8]) -> Self { + Self { + bytes, + start: 0, + position: 0, + finished: false, + } + } +} + +impl<'a> Iterator for CommaItems<'a> { + type Item = &'a [u8]; + + fn next(&mut self) -> Option { + while !self.finished { + if let Some(relative) = self.bytes[self.position..].iter().position(|byte| *byte == b',') { + let end = self.position + relative; + let item = super::super::trim_ows(&self.bytes[self.start..end]); + self.position = end + 1; + self.start = self.position; + if item.is_empty() { + continue; + } + return Some(item); + } + self.finished = true; + let item = super::super::trim_ows(&self.bytes[self.start..]); + if !item.is_empty() { + return Some(item); + } + } + None + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{ + CommaItems, ReferrerPolicy, ReferrerPolicyOwned, ReferrerPolicyTokenView, ReferrerPolicyValue, increment_item_count, + parse_referrer_policy_token, recognize_referrer_policy, + }; + use crate::sink::{EncodedValues, FieldSink}; + use crate::source::FieldSource; + use crate::{DecodeErrorKind, Field, FieldValue, TestSink}; + + #[test] + fn recognized_values_and_token_accessors_cover_every_policy() { + assert_eq!( + increment_item_count(usize::MAX - 1).expect("last count is representable"), + usize::MAX + ); + assert_eq!( + increment_item_count(usize::MAX).expect_err("count overflow is rejected").kind(), + DecodeErrorKind::InvalidNumber + ); + + let policies = [ + (ReferrerPolicyValue::NoReferrer, "no-referrer"), + (ReferrerPolicyValue::NoReferrerWhenDowngrade, "no-referrer-when-downgrade"), + (ReferrerPolicyValue::Origin, "origin"), + (ReferrerPolicyValue::OriginWhenCrossOrigin, "origin-when-cross-origin"), + (ReferrerPolicyValue::SameOrigin, "same-origin"), + (ReferrerPolicyValue::StrictOrigin, "strict-origin"), + (ReferrerPolicyValue::StrictOriginWhenCrossOrigin, "strict-origin-when-cross-origin"), + (ReferrerPolicyValue::UnsafeUrl, "unsafe-url"), + ]; + for (policy, wire) in policies { + assert_eq!(policy.as_str(), wire); + let token = parse_referrer_policy_token(wire.as_bytes()).expect("recognized policy parses"); + assert_eq!(token.as_str(), wire); + assert_eq!(token.policy(), Some(policy)); + let mut neighbor = wire.as_bytes().to_vec(); + for index in 0..neighbor.len() { + for byte in crate::test_support::substitution_bytes(wire.as_bytes()[index], index, wire.len()) { + neighbor[index] = byte; + assert_eq!( + recognize_referrer_policy(&neighbor), + (byte == wire.as_bytes()[index]).then_some(policy), + "{neighbor:?}" + ); + } + neighbor[index] = wire.as_bytes()[index]; + } + } + + let extension = parse_referrer_policy_token(b"future-policy").expect("extension token parses"); + assert_eq!(extension.as_str(), "future-policy"); + assert_eq!(extension.policy(), None); + assert_eq!( + parse_referrer_policy_token(b"not a token") + .expect_err("spaces are not token bytes") + .kind(), + DecodeErrorKind::InvalidToken + ); + } + + #[test] + fn constructors_fallbacks_and_conversions_preserve_wire_order() { + let policy = ReferrerPolicyOwned::new(ReferrerPolicyValue::NoReferrer) + .with_fallback(ReferrerPolicyValue::StrictOrigin) + .with_fallback(ReferrerPolicyValue::UnsafeUrl); + assert_eq!( + policy.policies().collect::, _>>(), + Ok(vec![ + ReferrerPolicyValue::NoReferrer, + ReferrerPolicyValue::StrictOrigin, + ReferrerPolicyValue::UnsafeUrl, + ]) + ); + assert_eq!(policy.preferred(), Ok(ReferrerPolicyValue::UnsafeUrl)); + assert_eq!(format!("{policy:?}"), "ReferrerPolicyOwned { value_count: 3 }"); + + for converted in [ + ReferrerPolicyOwned::try_from("future, origin"), + ReferrerPolicyOwned::try_from(String::from("future, origin")), + ReferrerPolicyOwned::try_from(FieldValue::from_static("future, origin")), + ] { + let converted = converted.expect("valid fallback list"); + let tokens = converted + .tokens() + .map(|token| token.map(ReferrerPolicyTokenView::as_str)) + .collect::, _>>(); + assert_eq!(tokens, Ok(vec!["future", "origin"])); + assert_eq!(converted.preferred(), Ok(ReferrerPolicyValue::Origin)); + } + assert_eq!( + ReferrerPolicyOwned::try_from(FieldValue::from_static("origin")) + .expect("recognized field value") + .preferred(), + Ok(ReferrerPolicyValue::Origin) + ); + + assert_eq!( + ReferrerPolicyOwned::try_from(" , \t, ").expect_err("empty list").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ReferrerPolicyOwned::try_from(String::from("origin\n")) + .expect_err("invalid field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ReferrerPolicyOwned::try_from("origin\n") + .expect_err("invalid borrowed field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ReferrerPolicyOwned::try_from("origin, not valid") + .expect_err("invalid token") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + ReferrerPolicyOwned::try_from("future") + .expect("extension-only list is syntactically valid") + .preferred() + .expect_err("no recognized policy") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn borrowed_and_owned_decoders_share_comma_list_classification() { + let mut table = TestSink::new(); + assert!(ReferrerPolicy::view(&table).expect("absent view succeeds").is_none()); + assert!(ReferrerPolicy::owned(&table).expect("absent owned decode succeeds").is_none()); + + table + .set_values( + ReferrerPolicy::name(), + EncodedValues::from_vec(vec![ + FieldValue::from_static("no-referrer"), + FieldValue::from_static("future, origin"), + ]), + ) + .expect("table accepts policies"); + let view = ReferrerPolicy::view(&table).expect("view decodes").expect("header is present"); + assert_eq!(view.field_values().count(), 2); + assert_eq!( + view.tokens() + .map(|token| token.map(ReferrerPolicyTokenView::as_str)) + .collect::, _>>(), + Ok(vec!["no-referrer", "future", "origin"]) + ); + assert_eq!(view.preferred(), Ok(ReferrerPolicyValue::Origin)); + assert_eq!(format!("{view:?}"), "ReferrerPolicyView { value_count: 2 }"); + + let owned = ReferrerPolicy::owned(&table) + .expect("owned value decodes") + .expect("header is present"); + let mut output = TestSink::new(); + ReferrerPolicy::insert(&mut output, owned).expect("owned value inserts"); + assert_eq!(output.lines(ReferrerPolicy::name()).expect("inserted values").len(), 2); + + for raw in [" , ", "origin, not valid", "\"origin"] { + table + .set_values( + ReferrerPolicy::name(), + EncodedValues::single(FieldValue::from_str(raw).expect("safe field value")), + ) + .expect("table accepts raw value"); + let view_error = ReferrerPolicy::view(&table).expect_err("view rejects invalid list"); + let owned_error = ReferrerPolicy::owned(&table).expect_err("owned rejects invalid list"); + assert_eq!(view_error.kind(), owned_error.kind()); + } + + table + .set_values( + ReferrerPolicy::name(), + EncodedValues::single(FieldValue::from_static("future-policy")), + ) + .expect("table accepts extension policy"); + assert_eq!( + ReferrerPolicy::view(&table) + .expect("extension view decodes") + .expect("header is present") + .preferred() + .expect_err("no recognized policy") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + table + .set_values( + ReferrerPolicy::name(), + EncodedValues::from_vec(vec![ + FieldValue::from_static(" no-referrer "), + FieldValue::from_static("strict-origin"), + ]), + ) + .expect("table accepts recognized policies"); + assert_eq!( + ReferrerPolicy::view(&table) + .expect("recognized fast-path view") + .expect("header is present") + .preferred(), + Ok(ReferrerPolicyValue::StrictOrigin) + ); + assert_eq!( + ReferrerPolicy::owned(&table) + .expect("recognized fast-path owned decode") + .expect("header is present") + .preferred(), + Ok(ReferrerPolicyValue::StrictOrigin) + ); + } + + #[test] + fn comma_items_skip_empty_members_and_trim_ows() { + assert_eq!( + CommaItems::new(b" , no-referrer,\t, origin , ").collect::>(), + [b"no-referrer".as_slice(), b"origin".as_slice()] + ); + assert_eq!(CommaItems::new(b",,,").next(), None); + } + + #[test] + fn malformed_private_storage_surfaces_iteration_errors() { + let malformed = ReferrerPolicyOwned { + values: super::FieldLinesIter::one(FieldValue::from_static("origin, bad token")), + }; + let mut tokens = malformed.tokens(); + assert_eq!( + tokens.next().expect("first token").expect("recognized token").policy(), + Some(ReferrerPolicyValue::Origin) + ); + assert_eq!( + tokens.next().expect("second token").expect_err("invalid token").kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + malformed.preferred().expect_err("iteration error propagates").kind(), + DecodeErrorKind::InvalidToken + ); + + let malformed_view = super::ReferrerPolicyView { + values: crate::source::FieldLines::single(&crate::FieldName::ReferrerPolicy, b"bad token"), + }; + assert_eq!( + malformed_view.preferred().expect_err("borrowed iteration error propagates").kind(), + DecodeErrorKind::InvalidToken + ); + } +} diff --git a/crates/http_headers/src/headers/security/strict_transport_security.rs b/crates/http_headers/src/headers/security/strict_transport_security.rs new file mode 100644 index 000000000..0598fce0f --- /dev/null +++ b/crates/http_headers/src/headers/security/strict_transport_security.rs @@ -0,0 +1,1301 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt::Write as _; +use std::str; +use std::time::Duration; + +use compact_str::CompactString; + +use super::super::ExtensionValue; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, SingleValueField, validate}; + +/// Defines the `Strict-Transport-Security` header. +/// +/// # Specification +/// +/// Defined by [RFC 6797 section 6.1](https://www.rfc-editor.org/rfc/rfc6797#section-6.1). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{StrictTransportSecurity, StrictTransportSecurityOwned}; +/// +/// let mut map = HeaderMap::new(); +/// StrictTransportSecurity::insert( +/// &mut map, +/// StrictTransportSecurityOwned::new(std::time::Duration::from_secs(60))?, +/// )?; +/// assert!(StrictTransportSecurity::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct StrictTransportSecurity { + _private: (), +} + +/// Owned value for the `Strict-Transport-Security` header. +/// +/// # Specification +/// +/// Defined by [RFC 6797 section 6.1]. The optional `preload` directive is a +/// nonstandard convention documented by the [HSTS preload service]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::StrictTransportSecurityOwned::builder( +/// std::time::Duration::from_secs(31_536_000), +/// ) +/// .include_subdomains() +/// .build()?; +/// assert_eq!(value.max_age(), std::time::Duration::from_secs(31_536_000)); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Strict-Transport-Security: max-age=31536000` sets the lifetime. +/// `Strict-Transport-Security: max-age=31536000; includeSubDomains; preload` +/// also covers subdomains and requests preload-list inclusion. +/// +/// [RFC 6797 section 6.1]: https://www.rfc-editor.org/rfc/rfc6797#section-6.1 +/// [HSTS preload service]: https://hstspreload.org/ +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub struct StrictTransportSecurityOwned { + value: FieldValue, + summary: HstsSummary, +} + +/// Borrowed value for the `Strict-Transport-Security` header. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ``` +/// use http_headers::headers::{StrictTransportSecurity, StrictTransportSecurityView}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let view: StrictTransportSecurityView<'_> = StrictTransportSecurity::decode_view( +/// FieldValueRef::new(b"max-age=31536000; includeSubDomains"), +/// )?; +/// assert_eq!(view.max_age().as_secs(), 31_536_000); +/// assert!(view.include_subdomains()); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct StrictTransportSecurityView<'a> { + value: FieldValueRef<'a>, + summary: HstsSummary, +} + +/// Builder for a canonical `Strict-Transport-Security` field value. +#[derive(Clone, Debug)] +/// # Examples +/// +/// ``` +/// use std::time::Duration; +/// +/// use http_headers::headers::{StrictTransportSecurityBuilder, StrictTransportSecurityOwned}; +/// +/// let builder: StrictTransportSecurityBuilder = +/// StrictTransportSecurityOwned::builder(Duration::from_secs(31_536_000)) +/// .include_subdomains() +/// .preload(); +/// let value = builder.build()?; +/// assert_eq!( +/// value.as_field_value().as_bytes(), +/// b"max-age=31536000; includeSubDomains; preload", +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct StrictTransportSecurityBuilder { + max_age: Duration, + include_subdomains: bool, + preload: bool, + extensions: Vec, +} + +#[derive(Clone, Debug)] +struct BuilderExtension { + name: CompactString, + value: Option, +} + +/// One borrowed HSTS directive. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ``` +/// use std::time::Duration; +/// +/// use http_headers::headers::{HstsDirectiveView, StrictTransportSecurityOwned}; +/// +/// let value = StrictTransportSecurityOwned::builder(Duration::from_secs(60)) +/// .include_subdomains() +/// .build()?; +/// let mut directives = value.directives(); +/// let first: HstsDirectiveView<'_> = directives.next().transpose()?.expect("max-age is present"); +/// assert_eq!(first.name(), "max-age"); +/// assert_eq!(first.value(), Some(&b"60"[..])); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct HstsDirectiveView<'a> { + raw: &'a [u8], + name: &'a str, + value: Option<&'a [u8]>, +} + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +struct HstsSummary { + max_age: Duration, + include_subdomains: bool, + preload: bool, +} + +impl StrictTransportSecurityOwned { + /// Constructs a header containing only `max-age`. + /// + /// # Errors + /// + /// Returns [`DecodeErrorKind::InvalidNumber`] if `max_age` contains + /// fractional seconds, or an error if field-value construction fails. + /// # Examples + /// + /// ``` + /// use std::time::Duration; + /// + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::new(Duration::from_secs(31_536_000))?; + /// assert_eq!(value.max_age(), Duration::from_secs(31_536_000)); + /// assert!(!value.include_subdomains()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn new(max_age: Duration) -> Result { + Self::builder(max_age).build() + } + + /// Creates an HSTS builder with the required `max-age`. + /// + /// Fractional seconds are rejected by [`StrictTransportSecurityBuilder::build`]. + #[must_use] + /// # Examples + /// + /// ``` + /// use std::time::Duration; + /// + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::builder(Duration::from_secs(600)) + /// .preload() + /// .build()?; + /// assert!(value.preload()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn builder(max_age: Duration) -> StrictTransportSecurityBuilder { + StrictTransportSecurityBuilder { + max_age, + include_subdomains: false, + preload: false, + extensions: Vec::new(), + } + } + + /// Returns the `max-age` duration. + #[must_use] + /// # Examples + /// + /// ``` + /// use std::time::Duration; + /// + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=31536000")?; + /// assert_eq!(value.max_age(), Duration::from_secs(31_536_000)); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn max_age(&self) -> Duration { + self.summary.max_age + } + + /// Returns whether `includeSubDomains` is present. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=60; includeSubDomains")?; + /// assert!(value.include_subdomains()); + /// + /// let narrow = StrictTransportSecurityOwned::try_from("max-age=60")?; + /// assert!(!narrow.include_subdomains()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn include_subdomains(&self) -> bool { + self.summary.include_subdomains + } + + /// Returns whether the nonstandard `preload` directive is present. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=60; preload")?; + /// assert!(value.preload()); + /// + /// let plain = StrictTransportSecurityOwned::try_from("max-age=60")?; + /// assert!(!plain.preload()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn preload(&self) -> bool { + self.summary.preload + } + + /// Iterates all directives, including extensions, in wire order. + /// + /// # Errors + /// + /// An item is [`Err`] when a stored directive has invalid syntax or a + /// non-UTF-8 name. Iteration resumes with the following directive. + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=60; includeSubDomains")?; + /// let names = value + /// .directives() + /// .map(|directive| Ok(directive?.name())) + /// .collect::, http_headers::DecodeError>>()?; + /// assert_eq!(names, ["max-age", "includeSubDomains"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn directives(&self) -> impl Iterator, DecodeError>> { + HstsItems::new(self.value.as_bytes()).map(parse_hsts_directive) + } + + /// Returns the stored field value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=60; preload")?; + /// assert_eq!(value.as_field_value().as_bytes(), b"max-age=60; preload"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=60")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"max-age=60"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(StrictTransportSecurityOwned, |value| value.value); + +impl<'a> StrictTransportSecurityView<'a> { + /// Returns the `max-age` duration. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurity; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = StrictTransportSecurity::decode_view(FieldValueRef::new(b"max-age=31536000"))?; + /// assert_eq!(view.max_age().as_secs(), 31_536_000); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn max_age(self) -> Duration { + self.summary.max_age + } + + /// Returns whether `includeSubDomains` is present. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurity; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = + /// StrictTransportSecurity::decode_view(FieldValueRef::new(b"max-age=60; includeSubDomains"))?; + /// assert!(view.include_subdomains()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn include_subdomains(self) -> bool { + self.summary.include_subdomains + } + + /// Returns whether the nonstandard `preload` directive is present. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurity; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = StrictTransportSecurity::decode_view(FieldValueRef::new(b"max-age=60; preload"))?; + /// assert!(view.preload()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn preload(self) -> bool { + self.summary.preload + } + + /// Iterates all directives, including extensions, in wire order. + /// + /// # Errors + /// + /// An item is [`Err`] when a stored directive has invalid syntax or a + /// non-UTF-8 name. Iteration resumes with the following directive. + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurity; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = StrictTransportSecurity::decode_view(FieldValueRef::new(b"max-age=60; preload"))?; + /// let names = view + /// .directives() + /// .map(|directive| Ok(directive?.name())) + /// .collect::, http_headers::DecodeError>>()?; + /// assert_eq!(names, ["max-age", "preload"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn directives(self) -> impl Iterator, DecodeError>> { + HstsItems::new(self.value.as_bytes()).map(parse_hsts_directive) + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurity; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = StrictTransportSecurity::decode_view(FieldValueRef::new(b"max-age=60"))?; + /// assert_eq!(view.as_field_value().as_bytes(), b"max-age=60"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } +} + +impl StrictTransportSecurityBuilder { + /// Adds `includeSubDomains`. + #[must_use] + /// # Examples + /// + /// ``` + /// use std::time::Duration; + /// + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::builder(Duration::from_secs(60)) + /// .include_subdomains() + /// .build()?; + /// assert_eq!( + /// value.as_field_value().as_bytes(), + /// b"max-age=60; includeSubDomains" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn include_subdomains(mut self) -> Self { + self.include_subdomains = true; + self + } + + /// Adds the nonstandard `preload` directive. + #[must_use] + /// # Examples + /// + /// ``` + /// use std::time::Duration; + /// + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::builder(Duration::from_secs(60)) + /// .preload() + /// .build()?; + /// assert_eq!(value.as_field_value().as_bytes(), b"max-age=60; preload"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn preload(mut self) -> Self { + self.preload = true; + self + } + + /// Adds an extension directive for validation by [`Self::build`]. + /// # Examples + /// + /// ``` + /// use std::time::Duration; + /// + /// use http_headers::headers::{ExtensionValue, StrictTransportSecurityOwned}; + /// + /// let value = StrictTransportSecurityOwned::builder(Duration::from_secs(60)) + /// .extension("report", ExtensionValue::Value("audit")) + /// .build()?; + /// assert_eq!( + /// value.as_field_value().as_bytes(), + /// b"max-age=60; report=audit" + /// ); + /// + /// // Reserved directive names are rejected. + /// assert!( + /// StrictTransportSecurityOwned::builder(Duration::from_secs(60)) + /// .extension("preload", ExtensionValue::Flag) + /// .build() + /// .is_err(), + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[must_use] + pub fn extension(mut self, name: impl AsRef, value: ExtensionValue<'_>) -> Self { + let name = name.as_ref(); + let value = match value { + ExtensionValue::Flag => None, + ExtensionValue::Value(value) => Some(CompactString::from(value)), + }; + self.extensions.push(BuilderExtension { + name: CompactString::from(name), + value, + }); + self + } + + /// Adds a flag-style extension directive. + #[must_use] + pub fn extension_flag(self, name: impl AsRef) -> Self { + self.extension(name, ExtensionValue::Flag) + } + + /// Adds an extension directive carrying a token or quoted-string value. + #[must_use] + pub fn extension_value(self, name: impl AsRef, value: impl AsRef) -> Self { + self.extension(name, ExtensionValue::Value(value.as_ref())) + } + + /// Builds the header. + /// + /// # Errors + /// + /// Returns an error if `max-age` contains fractional seconds, an extension + /// is malformed or uses a reserved name, or field-value construction fails. + /// # Examples + /// + /// ``` + /// use std::time::Duration; + /// + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::builder(Duration::from_secs(31_536_000)) + /// .include_subdomains() + /// .build()?; + /// assert_eq!(value.max_age(), Duration::from_secs(31_536_000)); + /// assert!(value.include_subdomains()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn build(self) -> Result { + if self.max_age.subsec_nanos() != 0 { + return Err(DecodeError::new( + &FieldName::StrictTransportSecurity, + DecodeErrorKind::InvalidNumber, + )); + } + if self.extensions.iter().any(|extension| { + !validate::token(extension.name.as_bytes()) + || known_hsts_name(extension.name.as_bytes()) + || extension + .value + .as_ref() + .is_some_and(|value| !valid_token_or_quoted(value.as_bytes())) + }) { + return Err(super::super::invalid_syntax(&FieldName::StrictTransportSecurity)); + } + let capacity = "max-age=".len() + + decimal_len(self.max_age.as_secs()) + + usize::from(self.include_subdomains) * "; includeSubDomains".len() + + usize::from(self.preload) * "; preload".len() + + self + .extensions + .iter() + .map(|extension| 2 + extension.name.len() + extension.value.as_ref().map_or(0, |value| 1 + value.len())) + .sum::(); + let mut wire = String::with_capacity(capacity); + write!(wire, "max-age={}", self.max_age.as_secs()).expect("writing to a String is infallible"); + if self.include_subdomains { + wire.push_str("; includeSubDomains"); + } + if self.preload { + wire.push_str("; preload"); + } + for extension in self.extensions { + wire.push_str("; "); + wire.push_str(&extension.name); + if let Some(value) = extension.value { + wire.push('='); + wire.push_str(&value); + } + } + StrictTransportSecurityOwned::try_from(wire) + } +} + +const fn decimal_len(mut value: u64) -> usize { + let mut length = 1; + while value >= 10 { + value /= 10; + length += 1; + } + length +} + +impl<'a> HstsDirectiveView<'a> { + /// Returns the directive name. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=60; includeSubDomains")?; + /// let directive = value + /// .directives() + /// .next() + /// .transpose()? + /// .expect("max-age is present"); + /// assert_eq!(directive.name(), "max-age"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn name(self) -> &'a str { + self.name + } + + /// Returns the raw token or quoted-string value. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=60; preload")?; + /// let mut directives = value.directives(); + /// let max_age = directives.next().transpose()?.expect("max-age is present"); + /// assert_eq!(max_age.value(), Some(&b"60"[..])); + /// + /// // Valueless directives report `None`. + /// let preload = directives.next().transpose()?.expect("preload is present"); + /// assert_eq!(preload.value(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn value(self) -> Option<&'a [u8]> { + self.value + } + + /// Returns the optional value as UTF-8. + /// + /// # Errors + /// + /// Returns an error when a quoted value contains non-UTF-8 `obs-text`. + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=60")?; + /// let directive = value + /// .directives() + /// .next() + /// .transpose()? + /// .expect("max-age is present"); + /// assert_eq!(directive.value_str()?, Some("60")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn value_str(self) -> Result, DecodeError> { + self.value + .map(str::from_utf8) + .transpose() + .map_err(|_invalid| DecodeError::new(&FieldName::StrictTransportSecurity, DecodeErrorKind::InvalidUtf8)) + } + + /// Returns the complete directive bytes after surrounding OWS trimming. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::StrictTransportSecurityOwned; + /// + /// let value = StrictTransportSecurityOwned::try_from("max-age=60; includeSubDomains")?; + /// let directive = value + /// .directives() + /// .nth(1) + /// .transpose()? + /// .expect("second directive"); + /// assert_eq!(directive.as_bytes(), b"includeSubDomains"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_bytes(self) -> &'a [u8] { + self.raw + } +} + +impl SingleValueField for StrictTransportSecurity { + type View<'a> = StrictTransportSecurityView<'a>; + type Owned = StrictTransportSecurityOwned; + + fn name() -> &'static FieldName { + &FieldName::StrictTransportSecurity + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + let summary = parse_hsts(value.as_bytes())?; + Ok(StrictTransportSecurityView { value, summary }) + } + + fn decode_owned(value: FieldValue) -> Result { + StrictTransportSecurityOwned::try_from(value) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +super::super::shared::impl_string_conversions!( + StrictTransportSecurityOwned, + &FieldName::StrictTransportSecurity, + super::super::invalid_syntax, + value +); + +impl TryFrom for StrictTransportSecurityOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + let summary = parse_hsts(value.as_bytes())?; + Ok(Self { value, summary }) + } +} + +fn parse_hsts(bytes: &[u8]) -> Result { + if bytes == b"max-age=31536000; includeSubDomains" { + return Ok(HstsSummary { + max_age: Duration::from_hours(8_760), + include_subdomains: true, + preload: false, + }); + } + if bytes == b"max-age=63072000; includeSubDomains; preload" { + return Ok(HstsSummary { + max_age: Duration::from_hours(17_520), + include_subdomains: true, + preload: true, + }); + } + if let Some(summary) = parse_canonical_hsts(bytes) { + return Ok(summary); + } + parse_hsts_directives(bytes) +} + +/// Parses the canonical serialization, returning `None` for anything else so +/// that [`parse_hsts_directives`] can apply the full grammar. +fn parse_canonical_hsts(bytes: &[u8]) -> Option { + let digits = bytes.strip_prefix(b"max-age=")?; + let mut max_age = 0_u64; + let mut digit_count = 0_usize; + let mut rest = digits; + while let Some((&byte, tail)) = rest.split_first() { + let digit = byte.wrapping_sub(b'0'); + if digit > 9 { + break; + } + max_age = max_age.checked_mul(10)?.checked_add(u64::from(digit))?; + digit_count += 1; + rest = tail; + } + if digit_count == 0 { + return None; + } + + let mut include_subdomains = false; + let mut preload = false; + loop { + rest = trim_start_ows(rest); + if rest.is_empty() { + break; + } + rest = trim_start_ows(rest.strip_prefix(b";")?); + if let Some(tail) = rest.strip_prefix(b"includeSubDomains") { + if include_subdomains { + return None; + } + include_subdomains = true; + rest = tail; + } else { + let tail = rest.strip_prefix(b"preload")?; + if preload { + return None; + } + preload = true; + rest = tail; + } + } + + Some(HstsSummary { + max_age: Duration::from_secs(max_age), + include_subdomains, + preload, + }) +} + +fn trim_start_ows(mut bytes: &[u8]) -> &[u8] { + while let Some((first, rest)) = bytes.split_first() { + if !matches!(first, b' ' | b'\t') { + break; + } + bytes = rest; + } + bytes +} + +fn parse_hsts_directives(bytes: &[u8]) -> Result { + let mut max_age = None; + let mut include_subdomains = false; + let mut preload = false; + for item in HstsItems::new(bytes) { + if let Some(seconds) = parse_unquoted_max_age(item) { + if max_age.is_some() { + return Err(super::super::invalid_syntax(&FieldName::StrictTransportSecurity)); + } + max_age = Some(seconds); + continue; + } + let (name, value) = split_hsts_directive_trimmed(item)?; + if validate::eq_ignore_ascii_case(name, b"max-age") { + if max_age.is_some() { + return Err(super::super::invalid_syntax(&FieldName::StrictTransportSecurity)); + } + let value = value.ok_or_else(|| super::super::invalid_syntax(&FieldName::StrictTransportSecurity))?; + max_age = Some(parse_delta_seconds(value)?); + } else if validate::eq_ignore_ascii_case(name, b"includesubdomains") { + if include_subdomains || value.is_some() { + return Err(super::super::invalid_syntax(&FieldName::StrictTransportSecurity)); + } + include_subdomains = true; + } else if validate::eq_ignore_ascii_case(name, b"preload") { + if preload || value.is_some() { + return Err(super::super::invalid_syntax(&FieldName::StrictTransportSecurity)); + } + preload = true; + } + } + Ok(HstsSummary { + max_age: Duration::from_secs(max_age.ok_or_else(|| super::super::invalid_syntax(&FieldName::StrictTransportSecurity))?), + include_subdomains, + preload, + }) +} + +fn parse_unquoted_max_age(bytes: &[u8]) -> Option { + if bytes.get(7) != Some(&b'=') || !validate::eq_ignore_ascii_case(&bytes[..7], b"max-age") { + return None; + } + validate::decimal_u64(&bytes[8..]) +} + +#[cfg(test)] +fn split_hsts_directive(bytes: &[u8]) -> Result<(&[u8], Option<&[u8]>), DecodeError> { + let bytes = super::super::trim_ows(bytes); + split_hsts_directive_trimmed(bytes) +} + +fn split_hsts_directive_trimmed(bytes: &[u8]) -> Result<(&[u8], Option<&[u8]>), DecodeError> { + let equals = bytes.iter().position(|byte| *byte == b'='); + let (name, value) = equals.map_or((bytes, None), |equals| (&bytes[..equals], Some(&bytes[equals + 1..]))); + if !validate::token(name) || value.is_some_and(|value| !valid_token_or_quoted(value)) { + return Err(super::super::invalid_syntax(&FieldName::StrictTransportSecurity)); + } + Ok((name, value)) +} + +fn parse_hsts_directive(bytes: &[u8]) -> Result, DecodeError> { + let bytes = super::super::trim_ows(bytes); + let (name, value) = split_hsts_directive_trimmed(bytes)?; + let name = str::from_utf8(name).expect("HTTP token validation guarantees ASCII"); + Ok(HstsDirectiveView { raw: bytes, name, value }) +} + +fn parse_delta_seconds(bytes: &[u8]) -> Result { + if bytes.is_empty() { + return Err(DecodeError::new( + &FieldName::StrictTransportSecurity, + DecodeErrorKind::InvalidNumber, + )); + } + if bytes.first() != Some(&b'"') { + return validate::decimal_u64(bytes) + .ok_or_else(|| DecodeError::new(&FieldName::StrictTransportSecurity, DecodeErrorKind::InvalidNumber)); + } + let inner = bytes + .strip_prefix(b"\"") + .and_then(|value| value.strip_suffix(b"\"")) + .ok_or_else(|| DecodeError::new(&FieldName::StrictTransportSecurity, DecodeErrorKind::InvalidNumber))?; + let mut value = 0_u64; + let mut digits = 0_usize; + let mut index = 0; + while index < inner.len() { + let byte = if inner[index] == b'\\' { + index += 1; + *inner + .get(index) + .ok_or_else(|| DecodeError::new(&FieldName::StrictTransportSecurity, DecodeErrorKind::InvalidNumber))? + } else { + inner[index] + }; + if !byte.is_ascii_digit() { + return Err(DecodeError::new( + &FieldName::StrictTransportSecurity, + DecodeErrorKind::InvalidNumber, + )); + } + value = value + .checked_mul(10) + .and_then(|current| current.checked_add(u64::from(byte - b'0'))) + .ok_or_else(|| DecodeError::new(&FieldName::StrictTransportSecurity, DecodeErrorKind::InvalidNumber))?; + digits += 1; + index += 1; + } + if digits == 0 { + return Err(DecodeError::new( + &FieldName::StrictTransportSecurity, + DecodeErrorKind::InvalidNumber, + )); + } + Ok(value) +} + +fn known_hsts_name(name: &[u8]) -> bool { + validate::eq_ignore_ascii_case(name, b"max-age") + || validate::eq_ignore_ascii_case(name, b"includesubdomains") + || validate::eq_ignore_ascii_case(name, b"preload") +} + +fn valid_token_or_quoted(bytes: &[u8]) -> bool { + validate::token(bytes) || valid_quoted_string(bytes) +} + +fn valid_quoted_string(bytes: &[u8]) -> bool { + if bytes.len() < 2 || bytes.first() != Some(&b'"') || bytes.last() != Some(&b'"') { + return false; + } + let mut escaped = false; + for byte in &bytes[1..bytes.len() - 1] { + if escaped { + if !matches!(byte, b'\t' | b' '..=b'~' | 0x80..=0xff) { + return false; + } + escaped = false; + } else if *byte == b'\\' { + escaped = true; + } else if !matches!(byte, b'\t' | b' ' | b'!' | b'#'..=b'[' | b']'..=b'~' | 0x80..=0xff) { + return false; + } + } + !escaped +} + +struct HstsItems<'a> { + bytes: &'a [u8], + start: usize, + position: usize, + finished: bool, +} + +impl<'a> HstsItems<'a> { + const fn new(bytes: &'a [u8]) -> Self { + Self { + bytes, + start: 0, + position: 0, + finished: false, + } + } +} + +impl<'a> Iterator for HstsItems<'a> { + type Item = &'a [u8]; + + fn next(&mut self) -> Option { + while !self.finished { + let mut quoted = false; + let mut escaped = false; + while let Some(byte) = self.bytes.get(self.position).copied() { + if escaped { + escaped = false; + } else if quoted && byte == b'\\' { + escaped = true; + } else if byte == b'"' { + quoted = !quoted; + } else if !quoted && byte == b';' { + let item = super::super::trim_ows(&self.bytes[self.start..self.position]); + self.position += 1; + self.start = self.position; + if item.is_empty() { + continue; + } + return Some(item); + } + self.position += 1; + } + self.finished = true; + let item = super::super::trim_ows(&self.bytes[self.start..]); + if !item.is_empty() { + return Some(item); + } + } + None + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::time::Duration; + + use super::{ + HstsItems, StrictTransportSecurity, StrictTransportSecurityOwned, decimal_len, parse_delta_seconds, parse_hsts, + parse_hsts_directive, parse_unquoted_max_age, split_hsts_directive, valid_quoted_string, + }; + use crate::sink::{EncodedValues, FieldSink}; + use crate::{DecodeErrorKind, Field, FieldValue, FieldValueRef, TestSink}; + + #[test] + fn unquoted_max_age_preserves_fallback_boundaries() { + for (wire, expected) in [ + (b"max-age=0".as_slice(), Some(0)), + (b"MAX-AGE=60", Some(60)), + (b"max-age=000000000000000000000000000060", Some(60)), + (b"max-age=18446744073709551615", Some(u64::MAX)), + (b"max-age=18446744073709551616", None), + (b"", None), + (b"max-age", None), + (b"max-age=", None), + (b"max-age=\"60\"", None), + (b"max-age =60", None), + (b"max-age= 60", None), + (b"max-age=60; x=y", None), + (b"x-extension=max-age=60", None), + ] { + assert_eq!(parse_unquoted_max_age(wire), expected, "{wire:?}"); + } + } + + #[test] + fn hsts_builder_keeps_short_extensions_inline_and_spills_long_ones() { + let short = StrictTransportSecurityOwned::builder(Duration::from_mins(1)).extension_value("x-mode", "stable"); + let long = StrictTransportSecurityOwned::builder(Duration::from_mins(1)) + .extension_value("extension-directive-with-a-value-that-exceeds-inline-storage", "enabled"); + assert!(!short.extensions[0].name.is_heap_allocated()); + assert!(long.extensions[0].name.is_heap_allocated()); + } + + #[test] + fn hsts_canonical_fast_path_agrees_with_the_general_parser() { + let inputs: &[&[u8]] = &[ + b"max-age=0", + b"max-age=31536000", + b"max-age=31536000; includeSubDomains", + b"max-age=31536000; includeSubDomains; preload", + b"max-age=60;preload", + b"max-age=60 ; includeSubDomains", + b"max-age=60\t;\tpreload ", + b"max-age=18446744073709551615", + b"max-age=18446744073709551616", + b"max-age=184467440737095516150", + b"max-age=", + b"max-age=60; includeSubDomains; includeSubDomains", + b"max-age=60; preload; preload", + b"max-age=60; preload=yes", + b"max-age=60; unknown", + b"max-age=60;;preload", + b"max-age=60; includeSubDomainsX", + b"MAX-AGE=60; includeSubDomains", + b"max-age=\"60\"", + b"includeSubDomains; max-age=60", + b"", + ]; + for input in inputs { + if let Some(summary) = super::parse_canonical_hsts(input) { + assert_eq!( + super::parse_hsts_directives(input), + Ok(summary), + "fast path disagreed for {input:?}" + ); + } + assert_eq!( + super::parse_hsts(input).ok(), + super::parse_hsts_directives(input).ok(), + "dispatch disagreed for {input:?}" + ); + } + } + + #[test] + fn builder_accessors_directives_and_header_round_trip() { + let value = StrictTransportSecurityOwned::builder(Duration::from_hours(8_760)) + .include_subdomains() + .preload() + .extension_value("report-to", "\"hsts\"") + .extension_flag("flag") + .build() + .expect("header builds"); + assert_eq!(value.max_age(), Duration::from_hours(8_760)); + assert!(value.include_subdomains()); + assert!(value.preload()); + assert_eq!( + value.as_field_value().as_bytes(), + b"max-age=31536000; includeSubDomains; preload; report-to=\"hsts\"; flag" + ); + + let directives = value.directives().collect::, _>>().expect("built directives parse"); + assert_eq!(directives[0].name(), "max-age"); + assert_eq!(directives[0].value(), Some(b"31536000".as_slice())); + assert_eq!(directives[0].value_str(), Ok(Some("31536000"))); + assert_eq!(directives[0].as_bytes(), b"max-age=31536000"); + assert_eq!(directives[4].name(), "flag"); + assert_eq!(directives[4].value(), None); + assert_eq!(directives[4].value_str(), Ok(None)); + assert_eq!(value.clone().into_field_value().as_bytes(), value.as_field_value().as_bytes()); + + let mut table = TestSink::new(); + assert!(StrictTransportSecurity::view(&table).expect("absent header succeeds").is_none()); + StrictTransportSecurity::insert(&mut table, value).expect("header inserts"); + let view = StrictTransportSecurity::view(&table) + .expect("view decodes") + .expect("header is present"); + assert_eq!(view.max_age(), Duration::from_hours(8_760)); + assert!(view.include_subdomains()); + assert!(view.preload()); + assert_eq!( + view.as_field_value().as_bytes(), + b"max-age=31536000; includeSubDomains; preload; report-to=\"hsts\"; flag" + ); + assert_eq!(view.directives().count(), 5); + let owned = StrictTransportSecurity::owned(&table) + .expect("owned value decodes") + .expect("header is present"); + assert_eq!(owned.max_age(), Duration::from_hours(8_760)); + } + + #[test] + fn constructors_and_conversions_accept_general_grammar() { + let simple = StrictTransportSecurityOwned::new(Duration::from_mins(1)).expect("header builds"); + assert_eq!(simple.as_field_value().as_bytes(), b"max-age=60"); + assert_eq!( + ::as_field_value(&simple).as_bytes(), + b"max-age=60" + ); + assert_eq!( + ::into_field_value(simple).as_bytes(), + b"max-age=60" + ); + assert_eq!( + ::decode_owned(FieldValue::from_static("max-age=60"),) + .expect("direct owned decode") + .max_age(), + Duration::from_mins(1) + ); + assert_eq!( + ::decode_view(FieldValueRef::new(b"max-age=x"),) + .expect_err("invalid borrowed value") + .kind(), + DecodeErrorKind::InvalidNumber + ); + + for value in [ + StrictTransportSecurityOwned::try_from("MAX-AGE=\"6\\0\"; includeSubDomains; preload; future=\"a,b\""), + StrictTransportSecurityOwned::try_from(String::from("MAX-AGE=\"6\\0\"; includeSubDomains; preload; future=\"a,b\"")), + StrictTransportSecurityOwned::try_from( + FieldValue::from_str("MAX-AGE=\"6\\0\"; includeSubDomains; preload; future=\"a,b\"").expect("safe field value"), + ), + ] { + let value = value.expect("general HSTS grammar is valid"); + assert_eq!(value.max_age(), Duration::from_mins(1)); + assert!(value.include_subdomains()); + assert!(value.preload()); + } + + assert_eq!( + StrictTransportSecurityOwned::try_from(String::from("max-age=1\n")) + .expect_err("invalid field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + StrictTransportSecurityOwned::try_from("max-age=1\n") + .expect_err("invalid borrowed field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + StrictTransportSecurityOwned::try_from(FieldValue::from_static("max-age=x")) + .expect_err("invalid stored value") + .kind(), + DecodeErrorKind::InvalidNumber + ); + assert_eq!( + parse_hsts(b"max-age=63072000; includeSubDomains; preload") + .expect("common preload value") + .max_age, + Duration::from_hours(17_520) + ); + } + + #[test] + fn builder_rejects_reserved_and_invalid_extensions() { + for (name, value) in [ + ("max-age", None), + ("IncludeSubDomains", None), + ("preload", None), + ("bad name", None), + ("future", Some("bad value")), + ("future", Some("\"unterminated")), + ] { + let builder = StrictTransportSecurityOwned::builder(Duration::from_secs(1)); + let builder = match value { + Some(value) => builder.extension_value(name, value), + None => builder.extension_flag(name), + }; + assert_eq!( + builder.build().expect_err("invalid extension").kind(), + DecodeErrorKind::InvalidSyntax + ); + } + } + + #[test] + fn parser_reports_duplicate_missing_and_malformed_directives() { + let invalid = [ + b"".as_slice(), + b"includeSubDomains", + b"max-age", + b"max-age=", + b"max-age=x", + b"max-age=1; max-age=2", + b"max-age=\"1\"; max-age=\"2\"", + b"max-age=1; includeSubDomains; includeSubDomains", + b"max-age=1; includeSubDomains=yes", + b"max-age=1; preload; preload", + b"max-age=1; preload=yes", + b"max-age=1; =value", + b"max-age=1; future=\"unterminated", + b"max-age=18446744073709551616", + b"max-age=\"\"", + b"max-age=\"x\"", + b"max-age=\"1\\\"", + b"max-age=\"18446744073709551616\"", + ]; + for wire in invalid { + let _error = parse_hsts(wire).expect_err("malformed HSTS must fail"); + } + assert_eq!( + parse_delta_seconds(b"").expect_err("empty number").kind(), + DecodeErrorKind::InvalidNumber + ); + assert_eq!( + parse_delta_seconds(b"\"12").expect_err("unterminated number").kind(), + DecodeErrorKind::InvalidNumber + ); + assert_eq!( + parse_delta_seconds(b"\"1\\\"").expect_err("trailing quoted escape").kind(), + DecodeErrorKind::InvalidNumber + ); + } + + #[test] + fn directive_helpers_handle_quoted_delimiters_obs_text_and_invalid_utf8_values() { + assert_eq!( + HstsItems::new(b"max-age=1; future=\"a;b\";; preload").collect::>(), + [b"max-age=1".as_slice(), b"future=\"a;b\"".as_slice(), b"preload".as_slice(),] + ); + assert_eq!( + HstsItems::new(b"max-age=1; future=\"a\\\";b\"; preload").collect::>(), + [b"max-age=1".as_slice(), b"future=\"a\\\";b\"".as_slice(), b"preload".as_slice(),] + ); + assert_eq!( + split_hsts_directive(b" future=token ").expect("directive parses"), + (b"future".as_slice(), Some(b"token".as_slice())) + ); + assert!(valid_quoted_string(b"\"quoted\\\"value\"")); + assert!(valid_quoted_string(b"\"\xff\"")); + assert!(valid_quoted_string(b"\"\\\xff\"")); + assert!(!valid_quoted_string(b"\"trailing\\\"")); + assert!(!valid_quoted_string(b"\"line\nbreak\"")); + assert!(!valid_quoted_string(b"\"escaped\\\ncontrol\"")); + assert_eq!( + parse_hsts_directive(b"bad name").expect_err("invalid directive").kind(), + DecodeErrorKind::InvalidSyntax + ); + + let raw = b"future=\"\xff\""; + let directive = parse_hsts_directive(raw).expect("obs-text value parses"); + assert_eq!(directive.name(), "future"); + assert_eq!(directive.value(), Some(b"\"\xff\"".as_slice())); + assert_eq!( + directive.value_str().expect_err("obs-text is not UTF-8").kind(), + DecodeErrorKind::InvalidUtf8 + ); + } + + #[test] + fn singleton_decoding_rejects_multiple_values_and_decimal_len_counts_digits() { + let mut table = TestSink::new(); + table + .set_values( + StrictTransportSecurity::name(), + EncodedValues::from_vec(vec![FieldValue::from_static("max-age=1"), FieldValue::from_static("max-age=2")]), + ) + .expect("table accepts raw values"); + assert_eq!( + StrictTransportSecurity::view(&table) + .expect_err("singleton view rejects duplicates") + .kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!(decimal_len(0), 1); + assert_eq!(decimal_len(9), 1); + assert_eq!(decimal_len(10), 2); + assert_eq!(decimal_len(u64::MAX), 20); + } +} diff --git a/crates/http_headers/src/headers/security/x_content_type_options.rs b/crates/http_headers/src/headers/security/x_content_type_options.rs new file mode 100644 index 000000000..645bcd069 --- /dev/null +++ b/crates/http_headers/src/headers/security/x_content_type_options.rs @@ -0,0 +1,270 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +/// Defines the `X-Content-Type-Options` header. +/// +/// # Specification +/// +/// Defined by the Fetch standard's +/// [X-Content-Type-Options section](https://fetch.spec.whatwg.org/#x-content-type-options-header). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{XContentTypeOptions, XContentTypeOptionsOwned}; +/// +/// let mut map = HeaderMap::new(); +/// XContentTypeOptions::insert(&mut map, XContentTypeOptionsOwned::nosniff())?; +/// assert!(XContentTypeOptions::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct XContentTypeOptions { + _private: (), +} + +/// Owned value for the `X-Content-Type-Options` header. +/// +/// # Specification +/// +/// Defined by the Fetch standard's [X-Content-Type-Options section]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::XContentTypeOptionsOwned::nosniff(); +/// assert_eq!(value.as_field_value(), "nosniff"); +/// ``` +/// +/// `X-Content-Type-Options: nosniff` is the only supported value. +/// +/// [X-Content-Type-Options section]: https://fetch.spec.whatwg.org/#x-content-type-options-header +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub struct XContentTypeOptionsOwned { + value: FieldValue, +} + +/// Borrowed value for the `X-Content-Type-Options` header. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{XContentTypeOptions, XContentTypeOptionsView}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let view: XContentTypeOptionsView<'_> = +/// XContentTypeOptions::decode_view(FieldValueRef::new(b"nosniff"))?; +/// assert_eq!(view.as_field_value().as_bytes(), b"nosniff"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct XContentTypeOptionsView<'a> { + value: FieldValueRef<'a>, +} + +impl XContentTypeOptionsOwned { + /// Constructs the only valid value, `nosniff`. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::XContentTypeOptionsOwned; + /// + /// let value = XContentTypeOptionsOwned::nosniff(); + /// assert_eq!(value.as_field_value().as_bytes(), b"nosniff"); + /// ``` + pub fn nosniff() -> Self { + Self { + value: FieldValue::from_static("nosniff"), + } + } + + /// Returns the stored field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::XContentTypeOptionsOwned; + /// + /// let value = XContentTypeOptionsOwned::nosniff(); + /// assert_eq!(value.as_field_value().as_bytes(), b"nosniff"); + /// assert!(XContentTypeOptionsOwned::try_from("NoSniff").is_err()); + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::XContentTypeOptionsOwned; + /// + /// let value = XContentTypeOptionsOwned::nosniff().into_field_value(); + /// assert_eq!(value.as_bytes(), b"nosniff"); + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(XContentTypeOptionsOwned, |value| value.value); + +impl<'a> XContentTypeOptionsView<'a> { + /// Returns the original `nosniff` field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::XContentTypeOptions; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = XContentTypeOptions::decode_view(FieldValueRef::new(b"nosniff"))?; + /// assert_eq!(view.as_field_value().as_bytes(), b"nosniff"); + /// assert!(XContentTypeOptions::decode_view(FieldValueRef::new(b"NoSniff")).is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } +} + +impl Default for XContentTypeOptionsOwned { + fn default() -> Self { + Self::nosniff() + } +} + +impl SingleValueField for XContentTypeOptions { + type View<'a> = XContentTypeOptionsView<'a>; + type Owned = XContentTypeOptionsOwned; + + fn name() -> &'static FieldName { + &FieldName::XContentTypeOptions + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + validate_nosniff(value)?; + Ok(XContentTypeOptionsView { value }) + } + + fn decode_owned(value: FieldValue) -> Result { + XContentTypeOptionsOwned::try_from(value) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +super::super::shared::impl_string_conversions!( + XContentTypeOptionsOwned, + &FieldName::XContentTypeOptions, + super::super::invalid_syntax, + value +); + +impl TryFrom for XContentTypeOptionsOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + validate_nosniff(value.as_field_value_ref())?; + Ok(Self { value }) + } +} + +fn validate_nosniff(value: FieldValueRef<'_>) -> Result<(), DecodeError> { + if value.as_bytes() == b"nosniff" { + Ok(()) + } else { + Err(super::super::invalid_syntax(&FieldName::XContentTypeOptions)) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{XContentTypeOptions, XContentTypeOptionsOwned}; + use crate::{DecodeErrorKind, FieldValue, FieldValueRef, TestSink}; + + #[test] + fn constructors_accessors_and_header_round_trip_use_exact_nosniff_value() { + let value = XContentTypeOptionsOwned::nosniff(); + assert_eq!(value.as_field_value().as_bytes(), b"nosniff"); + assert_eq!(value.clone().into_field_value().as_bytes(), b"nosniff"); + assert_eq!(value, XContentTypeOptionsOwned::default()); + assert_eq!( + ::as_field_value(&value).as_bytes(), + b"nosniff" + ); + assert_eq!( + ::into_field_value(value.clone()).as_bytes(), + b"nosniff" + ); + assert_eq!( + ::decode_owned(FieldValue::from_static("nosniff"),) + .expect("direct owned decode"), + value + ); + assert_eq!( + ::decode_view(FieldValueRef::new(b"NoSniff",)) + .expect_err("borrowed value is case-sensitive") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let mut table = TestSink::new(); + assert!(XContentTypeOptions::view(&table).expect("absent header succeeds").is_none()); + XContentTypeOptions::insert(&mut table, value).expect("header inserts"); + let view = XContentTypeOptions::view(&table).expect("view decodes").expect("header is present"); + assert_eq!(view.as_field_value().as_bytes(), b"nosniff"); + let owned = XContentTypeOptions::owned(&table) + .expect("owned value decodes") + .expect("header is present"); + assert_eq!(owned.as_field_value().as_bytes(), b"nosniff"); + } + + #[test] + fn conversions_reject_case_changes_invalid_syntax_and_bad_field_strings() { + for value in [ + XContentTypeOptionsOwned::try_from("nosniff"), + XContentTypeOptionsOwned::try_from(String::from("nosniff")), + XContentTypeOptionsOwned::try_from(FieldValue::from_static("nosniff")), + ] { + assert_eq!(value.expect("exact value is valid").as_field_value().as_bytes(), b"nosniff"); + } + for raw in ["NoSniff", "nosniff ", ""] { + assert_eq!( + XContentTypeOptionsOwned::try_from(raw) + .expect_err("only exact nosniff is valid") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + assert_eq!( + XContentTypeOptionsOwned::try_from(String::from("nosniff\n")) + .expect_err("line break is not a field value") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + XContentTypeOptionsOwned::try_from("nosniff\n") + .expect_err("invalid borrowed field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } +} diff --git a/crates/http_headers/src/headers/set_cookie.rs b/crates/http_headers/src/headers/set_cookie.rs new file mode 100644 index 000000000..ac88e2515 --- /dev/null +++ b/crates/http_headers/src/headers/set_cookie.rs @@ -0,0 +1,576 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Opaque repeated `Set-Cookie` field-line storage. + +use std::str::FromStr; +use std::{fmt, slice, vec}; + +use super::shared::FieldLinesIter; +use crate::sink::{FieldSink, InsertError, InsertErrorKind}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, Field, FieldName, FieldValue, FieldValueRef}; + +/// Defines the `Set-Cookie` header. +/// +/// # Specification +/// +/// Defined by [RFC 6265 section 4.1](https://www.rfc-editor.org/rfc/rfc6265#section-4.1). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{SetCookie, SetCookieOwned}; +/// +/// let mut map = HeaderMap::new(); +/// let mut cookies = SetCookieOwned::new(); +/// cookies.push_str("session=abc123; Path=/; HttpOnly; Secure")?; +/// SetCookie::insert(&mut map, cookies)?; +/// assert!(SetCookie::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct SetCookie { + _private: (), +} + +/// Owned value for the `Set-Cookie` header. +/// +/// # Specification +/// +/// Defined by [RFC 6265 section 4.1]. +/// +/// # Examples +/// +/// ```rust +/// let mut cookies = http_headers::headers::SetCookieOwned::new(); +/// cookies.push_str("session=abc123; Path=/; HttpOnly; Secure")?; +/// assert_eq!(cookies.len(), 1); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// [RFC 6265 section 4.1]: https://www.rfc-editor.org/rfc/rfc6265#section-4.1 +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct SetCookieOwned { + values: FieldLinesIter, +} + +/// Borrowed value for the `Set-Cookie` header. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), http_headers::DecodeError> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{SetCookie, SetCookieView}; +/// +/// let mut map = HeaderMap::new(); +/// map.append( +/// http::header::SET_COOKIE, +/// http::HeaderValue::from_static("id=a3fWa; Max-Age=2592000; Secure; HttpOnly"), +/// ); +/// let cookies: SetCookieView<'_> = SetCookie::view(&map)?.expect("present"); +/// assert_eq!(cookies.len(), 1); +/// # Ok::<(), http_headers::DecodeError>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +pub struct SetCookieView<'a> { + values: FieldLines<'a>, +} + +impl fmt::Debug for SetCookieOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("SetCookieOwned").field("value_count", &self.values.len()).finish() + } +} + +impl fmt::Debug for SetCookieView<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("SetCookieView").field("value_count", &self.values.len()).finish() + } +} + +impl SetCookieOwned { + /// Creates an empty collection. + /// + /// # Examples + /// + /// ```rust + /// assert!(http_headers::headers::SetCookieOwned::new().is_empty()); + /// ``` + #[must_use] + pub fn new() -> Self { + Self { + values: FieldLinesIter::empty(), + } + } + + /// Adds one nonempty HTTP field value. + /// + /// This method does not parse RFC 6265 cookie grammar. Callers that include + /// untrusted data must encode it before constructing the field value; + /// delimiters such as `;` are accepted as opaque cookie attributes. + /// + /// # Errors + /// + /// Returns an error when the value is empty. + /// + /// # Examples + /// + /// ```rust + /// let mut cookies = http_headers::headers::SetCookieOwned::new(); + /// cookies.push(http_headers::FieldValue::from_static("a=1"))?; + /// assert_eq!(cookies.len(), 1); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn push(&mut self, mut value: FieldValue) -> Result<(), DecodeError> { + validate(value.as_field_value_ref())?; + value.set_sensitive(true); + self.values.push(value); + Ok(()) + } + + /// Adds one nonempty cookie string as an opaque field value. + /// + /// This method does not parse RFC 6265 cookie grammar. Callers that include + /// untrusted data must encode it before interpolation; `;` and `=` are + /// accepted because they are meaningful cookie delimiters. + /// + /// # Errors + /// + /// Returns an error when the string is not a valid nonempty field value. + /// + /// # Examples + /// + /// ```rust + /// let mut cookies = http_headers::headers::SetCookieOwned::new(); + /// cookies.push_str("a=1; HttpOnly")?; + /// assert_eq!(cookies.len(), 1); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn push_str(&mut self, value: &str) -> Result<(), DecodeError> { + let value = match FieldValue::from_str(value) { + Ok(value) => value, + Err(_invalid) => return Err(super::invalid_syntax(&FieldName::SetCookie)), + }; + self.push(value) + } + + /// Iterates the stored field values. + /// + /// # Examples + /// + /// ```rust + /// let mut cookies = http_headers::headers::SetCookieOwned::new(); + /// cookies.push_str("a=1")?; + /// assert!(cookies.iter().any(|cookie| cookie == "a=1")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn iter(&self) -> slice::Iter<'_, FieldValue> { + self.values.iter() + } + + /// Mutably iterates the stored field values. + /// + /// Each value must remain nonempty for later insertion to succeed. + /// Insertion restores the sensitive marker on every value. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldSensitivity; + /// + /// let mut cookies = http_headers::headers::SetCookieOwned::new(); + /// cookies.push_str("a=1")?; + /// cookies + /// .iter_mut() + /// .for_each(|cookie| cookie.set_sensitivity(FieldSensitivity::Sensitive)); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn iter_mut(&mut self) -> slice::IterMut<'_, FieldValue> { + self.values.iter_mut() + } + + /// Returns the number of cookie field values. + /// + /// # Examples + /// + /// ```rust + /// let mut cookies = http_headers::headers::SetCookieOwned::new(); + /// cookies.push_str("a=1")?; + /// assert_eq!(cookies.len(), 1); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[must_use] + pub fn len(&self) -> usize { + self.values.len() + } + + /// Returns whether no cookie field values are stored. + /// + /// # Examples + /// + /// ```rust + /// assert!(http_headers::headers::SetCookieOwned::new().is_empty()); + /// ``` + #[must_use] + pub fn is_empty(&self) -> bool { + self.values.is_empty() + } + + pub(crate) fn into_sensitive_encoded_values(self) -> Result { + fn restore_sensitivity(mut value: FieldValue) -> FieldValue { + value.set_sensitive(true); + value + } + + if self.values.iter().any(FieldValue::is_empty) { + return Err(InsertError::new(InsertErrorKind::InvalidValue)); + } + + Ok(match self.values { + FieldLinesIter::Empty => crate::sink::EncodedValues::new(), + FieldLinesIter::One(value) => crate::sink::EncodedValues::single(restore_sensitivity(value)), + FieldLinesIter::Many(values) => crate::sink::EncodedValues::from_vec(values.into_iter().map(restore_sensitivity).collect()), + }) + } +} + +impl FromStr for SetCookieOwned { + type Err = DecodeError; + + fn from_str(value: &str) -> Result { + let mut cookies = Self::new(); + cookies.push_str(value)?; + Ok(cookies) + } +} + +impl Default for SetCookieOwned { + fn default() -> Self { + Self::new() + } +} + +impl<'a> SetCookieView<'a> { + /// Iterates the borrowed field values. + /// + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::SetCookie; + /// let mut map = HeaderMap::new(); + /// map.append( + /// http::header::SET_COOKIE, + /// http::HeaderValue::from_static("a=1"), + /// ); + /// assert_eq!( + /// SetCookie::view(&map)?.map(|cookies| cookies.iter().count()), + /// Some(1) + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn iter(&self) -> impl Iterator> + '_ { + self.values.repeated().map(|value| value.with_sensitive(true)) + } + + /// Returns the number of cookie field values. + /// + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::SetCookie; + /// let mut map = HeaderMap::new(); + /// map.append( + /// http::header::SET_COOKIE, + /// http::HeaderValue::from_static("a=1"), + /// ); + /// assert_eq!(SetCookie::view(&map)?.map(|cookies| cookies.len()), Some(1)); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + #[must_use] + pub fn len(&self) -> usize { + self.values.len() + } + + /// Returns whether no cookie field values are present. + /// + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), http_headers::DecodeError> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::SetCookie; + /// let mut map = HeaderMap::new(); + /// map.append( + /// http::header::SET_COOKIE, + /// http::HeaderValue::from_static("a=1"), + /// ); + /// assert_eq!( + /// SetCookie::view(&map)?.map(|cookies| cookies.is_empty()), + /// Some(false) + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + #[must_use] + #[expect( + clippy::unused_self, + reason = "a SetCookie view is nonempty by construction and this query keeps collection semantics" + )] + pub fn is_empty(&self) -> bool { + false + } +} + +impl<'a> IntoIterator for &'a SetCookieOwned { + type Item = &'a FieldValue; + type IntoIter = slice::Iter<'a, FieldValue>; + + fn into_iter(self) -> Self::IntoIter { + self.values.iter() + } +} + +impl<'a> IntoIterator for &'a mut SetCookieOwned { + type Item = &'a mut FieldValue; + type IntoIter = slice::IterMut<'a, FieldValue>; + + fn into_iter(self) -> Self::IntoIter { + self.values.iter_mut() + } +} + +impl IntoIterator for SetCookieOwned { + type Item = FieldValue; + type IntoIter = vec::IntoIter; + + fn into_iter(self) -> Self::IntoIter { + self.values.into_vec().into_iter() + } +} + +impl Field for SetCookie { + type View<'a> = SetCookieView<'a>; + type Owned = SetCookieOwned; + + fn name() -> &'static FieldName { + &FieldName::SetCookie + } + + #[expect( + clippy::inline_always, + reason = "generic header forwarding should monomorphize into each source call site" + )] + #[inline(always)] + fn view_with(source: &S, _mode: crate::DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + decode_view_values(source.lines(Self::name())) + } + + #[expect( + clippy::inline_always, + reason = "generic header forwarding should monomorphize into each source call site" + )] + #[inline(always)] + fn owned_with(source: &S, _mode: crate::DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + decode_owned_values(source.lines(Self::name())) + } + + #[expect( + clippy::inline_always, + reason = "generic header forwarding should monomorphize into each sink call site" + )] + #[inline(always)] + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), value.into_sensitive_encoded_values()?) + } +} + +fn decode_view_values(values: Option>) -> Result>, DecodeError> { + let Some(values) = values else { + return Ok(None); + }; + values.validate_custom_source()?; + let mut iter = values.repeated(); + if let Some(value) = iter.next() { + validate(value)?; + } + iter.try_for_each(validate)?; + Ok(Some(SetCookieView { values })) +} + +fn decode_owned_values(values: Option>) -> Result, DecodeError> { + let Some(values) = values else { + return Ok(None); + }; + let mut decoded = SetCookieOwned { + values: FieldLinesIter::empty(), + }; + for (value, mut owned) in values.repeated_owned()? { + validate(value)?; + if !owned.is_sensitive() { + owned.set_sensitive(true); + } + decoded.values.push(owned); + } + Ok(Some(decoded)) +} + +fn validate(value: FieldValueRef<'_>) -> Result<(), DecodeError> { + if value.as_bytes().is_empty() { + Err(super::invalid_syntax(&FieldName::SetCookie)) + } else { + Ok(()) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #![expect( + clippy::assertions_on_result_states, + reason = "tests classify parser outcomes without needing successful values" + )] + + use super::{SetCookie, SetCookieOwned}; + use crate::sink::{EncodedValues, FieldSink, InsertError, InsertErrorKind}; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeErrorKind, FieldName, FieldValue, TestSink}; + + struct Source(Vec); + + impl FieldSource for Source { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::SetCookie) + .then(|| FieldLines::from_slice(name, &self.0)) + .flatten() + } + } + + #[test] + fn owned_collection_covers_storage_and_iterator_forms() { + let mut cookies = SetCookieOwned::default(); + assert!(cookies.is_empty()); + assert_eq!(cookies.len(), 0); + cookies.push(FieldValue::from_static("a=1")).expect("nonempty cookie"); + cookies.push_str("b=2").expect("valid cookie string"); + assert_eq!(cookies.len(), 2); + assert!(cookies.iter().all(FieldValue::is_sensitive)); + cookies.iter_mut().for_each(|value| value.set_sensitive(true)); + assert_eq!((&cookies).into_iter().count(), 2); + assert_eq!((&mut cookies).into_iter().count(), 2); + assert!(format!("{cookies:?}").contains("value_count")); + assert_eq!(cookies.clone().into_iter().count(), 2); + + let parsed = "c=3".parse::().expect("one cookie parses"); + assert_eq!(parsed.len(), 1); + assert_eq!( + "".parse::().expect_err("empty cookie must fail").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert!(cookies.push_str("\n").is_err()); + } + + #[test] + fn header_decoding_and_encoding_cover_absent_empty_and_repeated() { + let source = Source(vec![FieldValue::from_static("a=1"), FieldValue::from_static("b=2")]); + let view = SetCookie::view(&source).expect("valid repeated cookies").expect("present"); + assert_eq!(view.len(), 2); + assert!(!view.is_empty()); + assert_eq!(view.iter().count(), 2); + assert!(format!("{view:?}").contains("value_count")); + + let owned = SetCookie::owned(&source).expect("valid owned cookies").expect("present"); + assert!(owned.iter().all(FieldValue::is_sensitive)); + + let mut first = FieldValue::from_static("a=1"); + first.set_sensitive(true); + let mut second = FieldValue::from_static("b=2"); + second.set_sensitive(true); + assert!( + SetCookie::owned(&Source(vec![first, second])) + .expect("sensitive cookies remain valid") + .expect("present") + .iter() + .all(FieldValue::is_sensitive) + ); + + let mut table = TestSink::new(); + SetCookie::insert(&mut table, owned).expect("table accepts cookies"); + assert_eq!(table.lines(&FieldName::SetCookie).expect("cookies stored").repeated().count(), 2); + table.remove_values(&FieldName::SetCookie); + assert!(SetCookie::view(&table).expect("absence is valid").is_none()); + assert!(SetCookie::owned(&table).expect("absence is valid").is_none()); + + let invalid = Source(vec![FieldValue::from_static("a=1"), FieldValue::from_static("")]); + assert!(SetCookie::view(&invalid).is_err()); + assert!(SetCookie::owned(&invalid).is_err()); + + let empty = Source(vec![FieldValue::from_static("")]); + assert!(SetCookie::view(&empty).is_err()); + assert!(SetCookie::owned(&empty).is_err()); + + let _ = EncodedValues::new(); + } + + #[test] + fn emission_rejects_values_invalidated_through_mutable_iteration() { + let mut cookies = SetCookieOwned::new(); + cookies.push_str("a=1").expect("valid cookie"); + *cookies.iter_mut().next().expect("one cookie") = FieldValue::from_static(""); + + let mut sink = TestSink::new(); + assert_eq!( + SetCookie::insert(&mut sink, cookies), + Err(InsertError::new(InsertErrorKind::InvalidValue)) + ); + assert!(sink.lines(&FieldName::SetCookie).is_none()); + + let mut cookies = SetCookieOwned::new(); + cookies.push_str("a=1").expect("valid cookie"); + *(&mut cookies).into_iter().next().expect("one cookie") = FieldValue::from_static(""); + assert_eq!( + crate::sink::FieldSinkExt::append_set_cookie(&mut sink, cookies).err(), + Some(InsertError::new(InsertErrorKind::InvalidValue)) + ); + assert!(sink.lines(&FieldName::SetCookie).is_none()); + } +} diff --git a/crates/http_headers/src/headers/shared.rs b/crates/http_headers/src/headers/shared.rs new file mode 100644 index 000000000..4dc05bedf --- /dev/null +++ b/crates/http_headers/src/headers/shared.rs @@ -0,0 +1,479 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Cross-family blanket trait implementations and small parsing helpers +//! reused by more than one header family. + +#![allow( + dead_code, + unused_imports, + unused_macros, + reason = "shared helpers are selected by independent header-family features" +)] + +use std::{fmt, mem, slice, str}; + +use super::*; +use crate::sink::EncodedValues; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef}; + +#[derive(Clone, Eq, Hash, PartialEq)] +pub(super) enum FieldLinesIter { + Empty, + One(FieldValue), + Many(Vec), +} + +impl FieldLinesIter { + pub(super) const fn empty() -> Self { + Self::Empty + } + + pub(super) const fn one(value: FieldValue) -> Self { + Self::One(value) + } + + pub(super) fn push(&mut self, value: FieldValue) { + match self { + Self::Empty => *self = Self::One(value), + Self::One(first) => { + let first = mem::replace(first, FieldValue::from_static("")); + *self = Self::Many(vec![first, value]); + } + Self::Many(values) => values.push(value), + } + } + + pub(super) fn len(&self) -> usize { + self.as_slice().len() + } + + pub(super) fn iter(&self) -> slice::Iter<'_, FieldValue> { + self.as_slice().iter() + } + + pub(super) fn iter_mut(&mut self) -> slice::IterMut<'_, FieldValue> { + match self { + Self::Empty => [].iter_mut(), + Self::One(value) => slice::from_mut(value).iter_mut(), + Self::Many(values) => values.iter_mut(), + } + } + + pub(super) fn is_empty(&self) -> bool { + self.as_slice().is_empty() + } + + pub(super) fn into_encoded(self) -> EncodedValues { + match self { + Self::Empty => EncodedValues::new(), + Self::One(value) => EncodedValues::single(value), + Self::Many(values) => EncodedValues::from_vec(values), + } + } + + pub(super) fn into_vec(self) -> Vec { + match self { + Self::Empty => Vec::new(), + Self::One(value) => vec![value], + Self::Many(values) => values, + } + } + + fn as_slice(&self) -> &[FieldValue] { + match self { + Self::Empty => &[], + Self::One(value) => slice::from_ref(value), + Self::Many(values) => values, + } + } +} + +macro_rules! impl_from_str { + ($($(#[$meta:meta])* ($owned:ty, $descriptor:ty)),+ $(,)?) => { + $( + $(#[$meta])* + impl std::str::FromStr for $owned { + type Err = DecodeError; + + #[inline(always)] + fn from_str(value: &str) -> Result { + parse_field_value(<$descriptor as crate::Field>::name(), value) + .and_then(Self::try_from) + } + } + )+ + }; +} + +macro_rules! impl_ascii_display { + ($($(#[$meta:meta])* ($owned:ty, $descriptor:ty)),+ $(,)?) => { + $( + $(#[$meta])* + impl std::fmt::Display for $owned { + #[inline(always)] + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + let value = <$descriptor as crate::SingleValueField>::as_field_value(self); + fmt_ascii_value(value, f) + } + } + )+ + }; +} + +macro_rules! impl_field_value_conversion { + ($owned:ty, |$value:ident| $field_value:expr) => { + impl From<$owned> for $crate::FieldValue { + #[inline] + fn from($value: $owned) -> Self { + $field_value + } + } + }; +} + +pub(super) use impl_field_value_conversion; + +macro_rules! impl_string_conversions { + ($owned:ty, $name:expr, $invalid:path, $input:ident) => { + impl TryFrom<&str> for $owned { + type Error = $crate::DecodeError; + + fn try_from($input: &str) -> Result { + let value = $crate::FieldValue::from_str($input).map_err(|_invalid| $invalid($name))?; + Self::try_from(value) + } + } + + impl TryFrom for $owned { + type Error = $crate::DecodeError; + + fn try_from($input: String) -> Result { + let value = $crate::FieldValue::try_from($input).map_err(|_invalid| $invalid($name))?; + Self::try_from(value) + } + } + }; +} + +pub(super) use impl_string_conversions; + +macro_rules! impl_value_count_debug { + ($($value:ty => $name:literal),+ $(,)?) => { + $( + impl std::fmt::Debug for $value { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct($name) + .field("value_count", &self.values.len()) + .finish() + } + } + )+ + }; +} + +pub(super) use impl_value_count_debug; + +fn parse_field_value(name: &'static FieldName, value: &str) -> Result { + match FieldValue::from_str(value) { + Ok(value) => Ok(value), + Err(_invalid) => Err(invalid_syntax(name)), + } +} + +fn ascii_value(value: &FieldValue) -> Result<&str, fmt::Error> { + match str::from_utf8(value.as_bytes()) { + Ok(value) => Ok(value), + Err(_invalid) => Err(fmt::Error), + } +} + +fn fmt_ascii_value(value: &FieldValue, f: &mut fmt::Formatter<'_>) -> fmt::Result { + ascii_value(value).and_then(|value| f.write_str(value)) +} + +pub(super) fn fmt_ascii_values<'a>(values: impl IntoIterator>, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut values = values.into_iter(); + if let Some(value) = values.next() { + f.write_str(str::from_utf8(value.as_bytes()).map_err(|_invalid| fmt::Error)?)?; + } + for value in values { + f.write_str(", ")?; + f.write_str(str::from_utf8(value.as_bytes()).map_err(|_invalid| fmt::Error)?)?; + } + Ok(()) +} + +impl_from_str!( + #[cfg(any(test, feature = "headers-negotiation"))] + (AcceptOwned, Accept), + #[cfg(any(test, feature = "headers-negotiation"))] + (AcceptEncodingOwned, AcceptEncoding), + #[cfg(any(test, feature = "headers-negotiation"))] + (AcceptLanguageOwned, AcceptLanguage), + #[cfg(any(test, feature = "headers-range"))] + (AcceptRangesOwned, AcceptRanges), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowCredentialsOwned, AccessControlAllowCredentials), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowHeadersOwned, AccessControlAllowHeaders), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowMethodsOwned, AccessControlAllowMethods), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowOriginOwned, AccessControlAllowOrigin), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlExposeHeadersOwned, AccessControlExposeHeaders), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlMaxAgeOwned, AccessControlMaxAge), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlRequestHeadersOwned, AccessControlRequestHeaders), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlRequestMethodOwned, AccessControlRequestMethod), + #[cfg(any(test, feature = "headers-negotiation"))] + (AllowOwned, Allow), + #[cfg(any(test, feature = "headers-cache-control"))] + (CacheControlOwned, CacheControl), + #[cfg(any(test, feature = "headers-range"))] + (ContentRangeOwned, ContentRange), + #[cfg(any(test, feature = "headers-security"))] + (ContentSecurityPolicyOwned, ContentSecurityPolicy), + #[cfg(any(test, feature = "headers-content-type"))] + (ContentTypeOwned, ContentType), + #[cfg(any(test, feature = "headers-etag"))] + (ETagOwned, ETag), + #[cfg(any(test, feature = "headers-negotiation"))] + (HostOwned, Host), + #[cfg(any(test, feature = "headers-conditional"))] + (IfMatchOwned, IfMatch), + #[cfg(any(test, feature = "headers-conditional"))] + (IfModifiedSinceOwned, IfModifiedSince), + #[cfg(any(test, feature = "headers-conditional"))] + (IfNoneMatchOwned, IfNoneMatch), + #[cfg(any(test, feature = "headers-conditional"))] + (IfRangeOwned, IfRange), + #[cfg(any(test, feature = "headers-conditional"))] + (IfUnmodifiedSinceOwned, IfUnmodifiedSince), + #[cfg(any(test, feature = "headers-conditional"))] + (LastModifiedOwned, LastModified), + #[cfg(any(test, feature = "headers-location"))] + (LocationOwned, Location), + #[cfg(any(test, feature = "headers-range"))] + (RangeOwned, Range), + #[cfg(any(test, feature = "headers-security"))] + (ReferrerPolicyOwned, ReferrerPolicy), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketAcceptOwned, SecWebSocketAccept), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketExtensionsOwned, SecWebSocketExtensions), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketKeyOwned, SecWebSocketKey), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketProtocolOwned, SecWebSocketProtocol), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketVersionOwned, SecWebSocketVersion), + #[cfg(any(test, feature = "headers-negotiation"))] + (ServerOwned, Server), + #[cfg(any(test, feature = "headers-security"))] + (StrictTransportSecurityOwned, StrictTransportSecurity), + #[cfg(any(test, feature = "headers-user-agent"))] + (UserAgentOwned, UserAgent), + #[cfg(any(test, feature = "headers-negotiation"))] + (VaryOwned, Vary), + #[cfg(any(test, feature = "headers-security"))] + (XContentTypeOptionsOwned, XContentTypeOptions), +); + +impl_ascii_display!( + #[cfg(any(test, feature = "headers-range"))] + (ContentRangeOwned, ContentRange), + #[cfg(any(test, feature = "headers-negotiation"))] + (HostOwned, Host), + #[cfg(any(test, feature = "headers-conditional"))] + (IfModifiedSinceOwned, IfModifiedSince), + #[cfg(any(test, feature = "headers-conditional"))] + (IfUnmodifiedSinceOwned, IfUnmodifiedSince), + #[cfg(any(test, feature = "headers-conditional"))] + (LastModifiedOwned, LastModified), + #[cfg(any(test, feature = "headers-range"))] + (RangeOwned, Range), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketAcceptOwned, SecWebSocketAccept), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketKeyOwned, SecWebSocketKey), + #[cfg(any(test, feature = "headers-security"))] + (XContentTypeOptionsOwned, XContentTypeOptions), +); + +macro_rules! impl_ascii_values_display { + ($($(#[$meta:meta])* $owned:ty),+ $(,)?) => { + $( + $(#[$meta])* + impl std::fmt::Display for $owned { + #[inline] + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + fmt_ascii_values(self.values(), f) + } + } + )+ + }; +} + +impl_ascii_values_display!( + #[cfg(any(test, feature = "headers-negotiation"))] + AllowOwned, + #[cfg(any(test, feature = "headers-negotiation"))] + VaryOwned, +); + +fn error(name: &'static FieldName, kind: DecodeErrorKind) -> DecodeError { + DecodeError::new(name, kind) +} + +pub(super) fn invalid_syntax(name: &'static FieldName) -> DecodeError { + error(name, DecodeErrorKind::InvalidSyntax) +} + +/// Trims leading and trailing optional whitespace (`' '` and `'\t'`). +pub(super) fn trim_ows(bytes: &[u8]) -> &[u8] { + let start = bytes.iter().position(|byte| !matches!(byte, b' ' | b'\t')).unwrap_or(bytes.len()); + let end = bytes + .iter() + .rposition(|byte| !matches!(byte, b' ' | b'\t')) + .map_or(start, |index| index + 1); + &bytes[start..end] +} + +#[inline] +pub(super) fn has_non_ows(bytes: &[u8]) -> bool { + match bytes.first() { + Some(b' ' | b'\t') => bytes[1..].iter().any(|byte| !matches!(byte, b' ' | b'\t')), + Some(_) => true, + None => false, + } +} + +pub(super) fn value_from_bytes(name: &'static FieldName, bytes: Vec) -> Result { + match FieldValue::try_from(bytes) { + Ok(value) => Ok(value), + Err(_invalid) => Err(invalid_syntax(name)), + } +} + +// External monomorphizations are integration-tested; LLVM also emits an uncallable template. +#[cfg_attr(coverage_nightly, coverage(off))] +pub(super) fn normalized_comma_value<'a>( + name: &'static FieldName, + mut items: impl Iterator, +) -> Result, DecodeError> { + normalized_comma_value_impl(name, &mut items) +} + +fn normalized_comma_value_impl( + name: &'static FieldName, + items: &mut dyn Iterator, +) -> Result, DecodeError> { + let mut wire = Vec::new(); + for item in items { + if item.is_empty() { + continue; + } + if !wire.is_empty() { + wire.extend_from_slice(b", "); + } + wire.extend_from_slice(item); + } + if wire.is_empty() { + Ok(None) + } else { + value_from_bytes(name, wire).map(Some) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #![expect( + clippy::assertions_on_result_states, + reason = "tests classify parser outcomes without needing successful values" + )] + + use std::fmt; + + use super::{FieldLinesIter, ascii_value, has_non_ows, normalized_comma_value, trim_ows, value_from_bytes}; + use crate::{FieldName, FieldValue, FieldValueRef}; + + struct DisplayValues<'a>(&'a [FieldValueRef<'a>]); + + impl fmt::Display for DisplayValues<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + super::fmt_ascii_values(self.0.iter().copied(), f) + } + } + + #[test] + fn field_lines_cover_every_storage_shape() { + let mut lines = FieldLinesIter::empty(); + assert!(lines.is_empty()); + assert_eq!(lines.len(), 0); + assert!(lines.iter().next().is_none()); + assert!(lines.iter_mut().next().is_none()); + + lines.push(FieldValue::from_static("a")); + assert_eq!(lines.len(), 1); + lines.iter_mut().next().expect("one value").set_sensitive(true); + assert!(lines.iter().next().expect("one value").is_sensitive()); + + lines.push(FieldValue::from_static("b")); + lines.push(FieldValue::from_static("c")); + assert_eq!( + lines.iter().map(FieldValue::as_bytes).collect::>(), + [b"a".as_slice(), b"b".as_slice(), b"c".as_slice()] + ); + assert_eq!(lines.clone().into_encoded().len(), 3); + assert_eq!(lines.into_vec().len(), 3); + + assert_eq!(FieldLinesIter::one(FieldValue::from_static("x")).into_vec().len(), 1); + assert!(FieldLinesIter::empty().into_vec().is_empty()); + assert_eq!(FieldLinesIter::one(FieldValue::from_static("x")).into_encoded().len(), 1); + assert!(FieldLinesIter::empty().into_encoded().is_empty()); + } + + #[test] + fn repeated_ascii_formatting_handles_empty_and_invalid_values() { + assert_eq!(DisplayValues(&[]).to_string(), ""); + + let invalid = [FieldValueRef::new(b"\xff")]; + let mut output = String::new(); + assert_eq!( + fmt::write(&mut output, format_args!("{}", DisplayValues(&invalid))), + Err(fmt::Error) + ); + } + + #[test] + fn whitespace_and_normalization_helpers_cover_edges() { + assert_eq!(trim_ows(b"\t value \t"), b"value"); + assert_eq!(trim_ows(b" \t"), b""); + assert!(has_non_ows(b"value")); + assert!(has_non_ows(b" \tvalue")); + assert!(!has_non_ows(b" \t")); + assert!(!has_non_ows(b"")); + + let normalized = normalized_comma_value(&FieldName::CacheControl, [b"".as_slice(), b"one", b"", b"two"].into_iter()) + .expect("valid field bytes") + .expect("nonempty normalized value"); + assert_eq!(normalized, "one, two"); + assert_eq!( + normalized_comma_value(&FieldName::CacheControl, [b"".as_slice()].into_iter()), + Ok(None) + ); + assert!(value_from_bytes(&FieldName::CacheControl, vec![b'\n']).is_err()); + + let content_range = "bytes 0-1/2".parse::().expect("valid range"); + assert_eq!(content_range.to_string(), "bytes 0-1/2"); + assert!("\n".parse::().is_err()); + + let opaque = FieldValue::try_from(vec![0xff]).expect("obs-text is valid field data"); + assert!(ascii_value(&opaque).is_err()); + } +} diff --git a/crates/http_headers/src/headers/tokens.rs b/crates/http_headers/src/headers/tokens.rs new file mode 100644 index 000000000..5e56cd924 --- /dev/null +++ b/crates/http_headers/src/headers/tokens.rs @@ -0,0 +1,54 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::str; + +#[expect( + clippy::inline_always, + reason = "preserves measured CORS method projection after sharing the helper" +)] +#[inline(always)] +pub(super) fn common_method(bytes: &[u8]) -> Option<&'static str> { + match bytes { + b"GET" => Some("GET"), + b"PUT" => Some("PUT"), + b"HEAD" => Some("HEAD"), + b"POST" => Some("POST"), + b"PATCH" => Some("PATCH"), + b"TRACE" => Some("TRACE"), + b"DELETE" => Some("DELETE"), + b"CONNECT" => Some("CONNECT"), + b"OPTIONS" => Some("OPTIONS"), + _ => None, + } +} + +#[inline] +pub(super) fn method_text(bytes: &[u8]) -> &str { + common_method(bytes).unwrap_or_else(|| str::from_utf8(bytes).expect("validated method tokens contain only ASCII")) +} + +#[expect( + clippy::inline_always, + reason = "preserves measured CORS field-name projection after sharing the helper" +)] +#[inline(always)] +pub(super) fn common_header_name(bytes: &[u8]) -> Option<&'static str> { + match bytes { + b"content-type" => Some("content-type"), + b"authorization" => Some("authorization"), + b"x-request-id" => Some("x-request-id"), + b"etag" => Some("etag"), + b"origin" => Some("origin"), + b"accept" => Some("accept"), + b"x-requested-with" => Some("x-requested-with"), + b"content-length" => Some("content-length"), + b"cache-control" => Some("cache-control"), + _ => None, + } +} + +#[inline] +pub(super) fn header_name_text(bytes: &[u8]) -> &str { + common_header_name(bytes).unwrap_or_else(|| str::from_utf8(bytes).expect("validated field-name tokens contain only ASCII")) +} diff --git a/crates/http_headers/src/headers/user_agent.rs b/crates/http_headers/src/headers/user_agent.rs new file mode 100644 index 000000000..146e2941a --- /dev/null +++ b/crates/http_headers/src/headers/user_agent.rs @@ -0,0 +1,359 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Validated opaque `User-Agent` values. + +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +/// Defines the `User-Agent` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 10.1.5](https://www.rfc-editor.org/rfc/rfc9110#section-10.1.5). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{UserAgent, UserAgentOwned}; +/// +/// let mut map = HeaderMap::new(); +/// UserAgent::insert(&mut map, UserAgentOwned::try_from_static("client/1")?)?; +/// assert!(UserAgent::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct UserAgent { + _private: (), +} + +/// Owned value for the `User-Agent` header. +/// +/// # Specification +/// +/// Defined by [RFC 9110 section 10.1.5]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::UserAgentOwned::try_from("example-client/1.0")?; +/// assert_eq!(value.as_bytes(), b"example-client/1.0"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `User-Agent: curl/8.5.0` identifies a command-line client, while +/// `User-Agent: example-client/1.0 (integration test)` includes a comment. +/// +/// [RFC 9110 section 10.1.5]: https://www.rfc-editor.org/rfc/rfc9110#section-10.1.5 +#[derive(Clone, Debug, Eq, Hash, Ord, PartialEq, PartialOrd)] +pub struct UserAgentOwned(FieldValue); + +/// Borrowed value for the `User-Agent` header. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ``` +/// use http_headers::headers::{UserAgent, UserAgentView}; +/// use http_headers::{FieldValue, SingleValueField}; +/// +/// let field = FieldValue::from_static("curl/8.4.0"); +/// let view: UserAgentView<'_> = +/// ::decode_view(field.as_field_value_ref())?; +/// assert_eq!(view.as_str()?, "curl/8.4.0"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct UserAgentView<'a>(FieldValueRef<'a>); + +impl UserAgentOwned { + /// Validates a static string without using a panicking constructor. + /// + /// # Errors + /// + /// Returns an error for an empty or invalid field value. + /// # Examples + /// + /// ``` + /// use http_headers::headers::UserAgentOwned; + /// + /// let value = UserAgentOwned::try_from_static("curl/8.4.0")?; + /// assert_eq!(value.as_bytes(), b"curl/8.4.0"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn try_from_static(value: &'static str) -> Result { + Self::try_from(value) + } + + /// Returns the wire bytes. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::UserAgentOwned::try_from("client/1")?; + /// assert_eq!(value.as_bytes(), b"client/1"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_bytes(&self) -> &[u8] { + self.0.as_bytes() + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ``` + /// use http_headers::headers::UserAgentOwned; + /// + /// let value = UserAgentOwned::try_from_static("example-client/1.0")?; + /// let field = value.into_field_value(); + /// assert_eq!(field.as_bytes(), b"example-client/1.0"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::shared::impl_field_value_conversion!(UserAgentOwned, |value| value.0); + +impl<'a> UserAgentView<'a> { + pub(crate) const fn field_value(self) -> FieldValueRef<'a> { + self.0 + } + + /// Returns the wire bytes. + #[must_use] + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::UserAgentOwned::try_from("client/1")?; + /// assert_eq!(value.as_bytes(), b"client/1"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_bytes(self) -> &'a [u8] { + self.0.as_bytes() + } + + /// Returns the value as UTF-8. + /// + /// # Errors + /// + /// Returns an error when the field contains non-UTF-8 bytes. + /// # Examples + /// + /// ``` + /// use http_headers::headers::UserAgent; + /// use http_headers::{FieldValue, SingleValueField}; + /// + /// let field = FieldValue::from_static("Mozilla/5.0 (X11; Linux x86_64)"); + /// let view = ::decode_view(field.as_field_value_ref())?; + /// assert_eq!(view.as_str()?, "Mozilla/5.0 (X11; Linux x86_64)"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_str(self) -> Result<&'a str, DecodeError> { + self.0 + .to_str() + .map_err(|_invalid| DecodeError::new(&FieldName::UserAgent, DecodeErrorKind::InvalidUtf8)) + } +} + +impl SingleValueField for UserAgent { + type View<'a> = UserAgentView<'a>; + type Owned = UserAgentOwned; + + fn name() -> &'static FieldName { + &FieldName::UserAgent + } + + #[inline] + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + validate(value)?; + Ok(UserAgentView(value)) + } + + #[inline] + fn decode_owned(value: FieldValue) -> Result { + UserAgentOwned::try_from(value) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.0 + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.0 + } +} + +super::shared::impl_string_conversions!(UserAgentOwned, &FieldName::UserAgent, super::invalid_syntax, value); + +impl TryFrom for UserAgentOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + validate(value.as_field_value_ref())?; + Ok(Self(value)) + } +} + +#[inline] +fn validate(value: FieldValueRef<'_>) -> Result<(), DecodeError> { + if crate::validate::field_value(value.as_bytes()) && super::has_non_ows(value.as_bytes()) { + Ok(()) + } else { + Err(super::invalid_syntax(&FieldName::UserAgent)) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #![expect( + clippy::assertions_on_result_states, + reason = "tests classify parser outcomes without needing successful values" + )] + + use super::{UserAgent, UserAgentOwned}; + use crate::sink::FieldSink; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef, SingleValueField, TestSink}; + + struct RawSource<'a>(FieldValueRef<'a>); + + impl FieldSource for RawSource<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_borrowed(name, std::slice::from_ref(&self.0)) + } + } + + #[test] + fn direct_and_source_decoders_reject_every_forbidden_field_byte() { + for byte in (0..=0x1f).chain(std::iter::once(0x7f)).filter(|byte| *byte != b'\t') { + let wire = [b'x', byte, b'y']; + let value = FieldValueRef::new(&wire); + let source = RawSource(value); + assert_eq!(UserAgent::decode_view(value).unwrap_err().kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!( + UserAgentOwned::try_from(String::from_utf8(wire.to_vec()).unwrap()) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!( + UserAgent::decode_view_with(value, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + UserAgent::view_with(&source, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + UserAgent::owned_with(&source, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + } + } + } + + #[test] + fn opaque_wire_bytes_and_sensitivity_survive_both_decode_modes() { + for wire in [b" \tclient/1 (opaque)\t ".as_slice(), b" \t\x80\xff\t "] { + for sensitive in [false, true] { + let value = FieldValueRef::new(wire).with_sensitive(sensitive); + let source = RawSource(value); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let direct = UserAgent::decode_view_with(value, mode).unwrap(); + let view = UserAgent::view_with(&source, mode).unwrap().unwrap(); + let owned = UserAgent::decode_owned_with(value.try_to_field_value().unwrap(), mode).unwrap(); + let sourced = UserAgent::owned_with(&source, mode).unwrap().unwrap(); + assert_eq!(direct.as_bytes(), wire); + assert_eq!(view.as_bytes(), wire); + assert_eq!(direct.field_value().is_sensitive(), sensitive); + assert_eq!(view.field_value().is_sensitive(), sensitive); + assert_eq!(owned, sourced); + assert_eq!(UserAgent::as_field_value(&owned).as_bytes(), wire); + assert_eq!(UserAgent::as_field_value(&owned).is_sensitive(), sensitive); + if wire.contains(&0xff) { + assert_eq!(direct.as_str().unwrap_err().kind(), DecodeErrorKind::InvalidUtf8); + } + #[cfg(feature = "http")] + { + let mut map = http::HeaderMap::new(); + map.insert(http::header::USER_AGENT, http::HeaderValue::try_from(value).unwrap()); + let view = UserAgent::view_with(&map, mode).unwrap().unwrap(); + let mapped = UserAgent::owned_with(&map, mode).unwrap().unwrap(); + assert_eq!(view.as_bytes(), wire); + assert_eq!(view.field_value().is_sensitive(), sensitive); + assert_eq!(mapped, owned); + assert_eq!(UserAgent::as_field_value(&mapped).is_sensitive(), sensitive); + } + } + } + } + } + + #[test] + fn constructors_accessors_and_header_round_trip() { + let static_value = UserAgentOwned::try_from_static("client/1").expect("valid user agent"); + assert_eq!(static_value.as_bytes(), b"client/1"); + assert_eq!( + UserAgentOwned::try_from(String::from("client/2")) + .expect("valid owned string") + .as_bytes(), + b"client/2" + ); + assert_eq!( + "client/3" + .parse::() + .expect("shared FromStr implementation") + .as_bytes(), + b"client/3" + ); + assert_eq!(static_value.clone().into_field_value(), FieldValue::from_static("client/1")); + + let mut table = TestSink::new(); + UserAgent::insert(&mut table, static_value).expect("table accepts user agent"); + let view = UserAgent::view(&table).expect("valid user agent").expect("present"); + assert_eq!(view.as_bytes(), b"client/1"); + assert_eq!(view.as_str(), Ok("client/1")); + assert_eq!( + UserAgent::owned(&table) + .expect("valid owned user agent") + .expect("present") + .as_bytes(), + b"client/1" + ); + table.remove_values(&FieldName::UserAgent); + assert!(UserAgent::view(&table).expect("absence is valid").is_none()); + } + + #[test] + fn validation_and_utf8_failures_are_reported() { + for value in ["", " ", "\t", " \t"] { + assert!(UserAgentOwned::try_from(value).is_err(), "{value:?}"); + } + assert!(UserAgentOwned::try_from(String::from("\n")).is_err()); + + let obs = FieldValue::try_from(vec![0xff]).expect("obs-text is valid field data"); + let view = ::decode_view(obs.as_field_value_ref()).expect("nonempty obs-text user agent"); + assert_eq!( + view.as_str().expect_err("obs-text is not UTF-8").kind(), + DecodeErrorKind::InvalidUtf8 + ); + assert!(::decode_owned(FieldValue::from_static(" ")).is_err()); + + let valid = FieldValue::from_static("direct/1"); + let view = ::decode_view(valid.as_field_value_ref()).expect("valid borrowed user agent"); + assert_eq!(view.as_bytes(), b"direct/1"); + let owned = ::decode_owned(valid.clone()).expect("valid owned user agent"); + assert_eq!(::as_field_value(&owned), &valid); + assert!(UserAgentOwned::try_from("\n").is_err()); + assert!(UserAgentOwned::try_from(String::from("\n")).is_err()); + } +} diff --git a/crates/http_headers/src/headers/websocket/mod.rs b/crates/http_headers/src/headers/websocket/mod.rs new file mode 100644 index 000000000..a74493c66 --- /dev/null +++ b/crates/http_headers/src/headers/websocket/mod.rs @@ -0,0 +1,25 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! WebSocket handshake keys, versions, protocols, and extensions. + +mod sec_web_socket_accept; +mod sec_web_socket_extensions; +mod sec_web_socket_key; +mod sec_web_socket_protocol; +mod sec_web_socket_version; +mod shared; + +#[doc(inline)] +pub use sec_web_socket_accept::{SecWebSocketAccept, SecWebSocketAcceptOwned, SecWebSocketAcceptView}; +#[doc(inline)] +pub use sec_web_socket_extensions::{ + SecWebSocketExtensions, SecWebSocketExtensionsBuilder, SecWebSocketExtensionsOwned, SecWebSocketExtensionsView, + WebSocketExtensionParameterView, WebSocketExtensionParameters, WebSocketExtensionView, +}; +#[doc(inline)] +pub use sec_web_socket_key::{SecWebSocketKey, SecWebSocketKeyOwned, SecWebSocketKeyView}; +#[doc(inline)] +pub use sec_web_socket_protocol::{SecWebSocketProtocol, SecWebSocketProtocolOwned, SecWebSocketProtocolView}; +#[doc(inline)] +pub use sec_web_socket_version::{SecWebSocketVersion, SecWebSocketVersionOwned}; diff --git a/crates/http_headers/src/headers/websocket/sec_web_socket_accept.rs b/crates/http_headers/src/headers/websocket/sec_web_socket_accept.rs new file mode 100644 index 000000000..b25045db6 --- /dev/null +++ b/crates/http_headers/src/headers/websocket/sec_web_socket_accept.rs @@ -0,0 +1,368 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use sha1::{Digest as _, Sha1}; + +use super::super::invalid_syntax; +use super::sec_web_socket_key::SecWebSocketKeyOwned; +use super::shared::{encode_fixed_base64, validate_canonical_base64}; +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +const WEBSOCKET_GUID: &[u8] = b"258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; + +/// Defines the `Sec-WebSocket-Accept` header. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.3](https://www.rfc-editor.org/rfc/rfc6455#section-11.3.3). +/// +/// # Examples +/// +/// ```rust +/// use http_headers::headers::SecWebSocketAccept; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let view = +/// SecWebSocketAccept::decode_view(FieldValueRef::new(b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="))?; +/// assert_eq!(view.encoded(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +#[derive(Debug)] +pub struct SecWebSocketAccept { + _private: (), +} + +/// Owned value for the `Sec-WebSocket-Accept` header. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.3]. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::headers::SecWebSocketAcceptOwned; +/// +/// let value = SecWebSocketAcceptOwned::try_from("s3pPLMBiTxaQ9kYGzzhZRbK+xOo=")?; +/// assert_eq!(value.encoded(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Sec-WebSocket-Accept: s3pPLMBiTxaQ9kYGzzhZRbK+xOo=` is the response +/// corresponding to the RFC example key. +/// +/// [RFC 6455 section 11.3.3]: https://www.rfc-editor.org/rfc/rfc6455#section-11.3.3 +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub struct SecWebSocketAcceptOwned { + value: FieldValue, +} + +/// Borrowed value for the `Sec-WebSocket-Accept` header. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{SecWebSocketAccept, SecWebSocketAcceptView}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let view: SecWebSocketAcceptView<'_> = +/// SecWebSocketAccept::decode_view(FieldValueRef::new(b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="))?; +/// assert_eq!(view.encoded(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct SecWebSocketAcceptView<'a> { + value: FieldValueRef<'a>, +} + +impl SecWebSocketAcceptOwned { + /// Encodes a 20-byte SHA-1 handshake digest. + /// + /// The digest is the binary result of hashing the client key's wire bytes + /// followed by the WebSocket GUID. + /// + /// # Errors + /// + /// The fixed-size digest always has a representable canonical encoding. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketAcceptOwned; + /// + /// let value = SecWebSocketAcceptOwned::from_digest([0; 20])?; + /// assert_eq!(value.encoded(), b"AAAAAAAAAAAAAAAAAAAAAAAAAAA="); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::unnecessary_wraps, + reason = "WebSocket value constructors consistently report validation through DecodeError" + )] + pub fn from_digest(digest: [u8; 20]) -> Result { + Ok(Self { + value: encode_fixed_base64(&digest), + }) + } + + /// Computes the server handshake value for a validated client key. + /// + /// # Errors + /// + /// Returns an error if the resulting field value cannot be represented. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketAcceptOwned; + /// + /// let key = "dGhlIHNhbXBsZSBub25jZQ==".try_into()?; + /// let value = SecWebSocketAcceptOwned::from_key(&key)?; + /// assert_eq!(value.encoded(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn from_key(key: &SecWebSocketKeyOwned) -> Result { + let mut digest = Sha1::new(); + digest.update(key.encoded()); + digest.update(WEBSOCKET_GUID); + Self::from_digest(digest.finalize().into()) + } + + /// Returns the canonical base64 bytes. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketAcceptOwned; + /// + /// let value = SecWebSocketAcceptOwned::try_from("s3pPLMBiTxaQ9kYGzzhZRbK+xOo=")?; + /// assert_eq!(value.encoded(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn encoded(&self) -> &[u8] { + self.value.as_bytes() + } + + /// Returns the stored field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketAcceptOwned; + /// + /// let value = SecWebSocketAcceptOwned::try_from("s3pPLMBiTxaQ9kYGzzhZRbK+xOo=")?; + /// assert_eq!( + /// value.as_field_value().as_bytes(), + /// b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketAcceptOwned; + /// + /// let value = SecWebSocketAcceptOwned::try_from("s3pPLMBiTxaQ9kYGzzhZRbK+xOo=")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(SecWebSocketAcceptOwned, |value| value.value); + +impl<'a> SecWebSocketAcceptView<'a> { + /// Returns the canonical base64 bytes. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketAccept; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = + /// SecWebSocketAccept::decode_view(FieldValueRef::new(b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="))?; + /// assert_eq!(view.encoded(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn encoded(self) -> &'a [u8] { + self.value.as_bytes() + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketAccept; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = + /// SecWebSocketAccept::decode_view(FieldValueRef::new(b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="))?; + /// assert_eq!( + /// view.as_field_value().as_bytes(), + /// b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } +} + +impl SingleValueField for SecWebSocketAccept { + type View<'a> = SecWebSocketAcceptView<'a>; + type Owned = SecWebSocketAcceptOwned; + + fn name() -> &'static FieldName { + &FieldName::SecWebSocketAccept + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + validate_canonical_base64(value.as_bytes(), 20, &FieldName::SecWebSocketAccept)?; + Ok(SecWebSocketAcceptView { value }) + } + + fn decode_owned(value: FieldValue) -> Result { + SecWebSocketAcceptOwned::try_from(value) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +super::super::shared::impl_string_conversions!(SecWebSocketAcceptOwned, &FieldName::SecWebSocketAccept, invalid_syntax, value); + +impl TryFrom for SecWebSocketAcceptOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + validate_canonical_base64(value.as_bytes(), 20, &FieldName::SecWebSocketAccept)?; + Ok(Self { value }) + } +} + +impl TryFrom<&SecWebSocketKeyOwned> for SecWebSocketAcceptOwned { + type Error = DecodeError; + + fn try_from(key: &SecWebSocketKeyOwned) -> Result { + Self::from_key(key) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{SecWebSocketAccept, SecWebSocketAcceptOwned}; + use crate::sink::{EncodedValues, FieldSink}; + use crate::{DecodeErrorKind, Field, FieldValue, FieldValueRef, TestSink}; + + #[test] + fn digest_key_and_accessors_produce_canonical_handshake_value() { + let zero = SecWebSocketAcceptOwned::from_digest([0; 20]).expect("digest encodes as base64"); + assert_eq!(zero.encoded(), b"AAAAAAAAAAAAAAAAAAAAAAAAAAA="); + assert_eq!(zero.as_field_value().as_bytes(), zero.encoded()); + assert_eq!(zero.clone().into_field_value().as_bytes(), zero.encoded()); + + let key = super::SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ==").expect("RFC key"); + let accept = SecWebSocketAcceptOwned::from_key(&key).expect("key hashes"); + assert_eq!(accept.encoded(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + assert_eq!( + ::as_field_value(&accept).as_bytes(), + accept.encoded() + ); + assert_eq!( + ::into_field_value(accept.clone()).as_bytes(), + accept.encoded() + ); + assert_eq!( + ::decode_owned(FieldValue::from_static("s3pPLMBiTxaQ9kYGzzhZRbK+xOo="),) + .expect("direct owned decode"), + accept + ); + assert_eq!( + ::decode_view(FieldValueRef::new(b"s3pPLMBiTxaQ9kYGzzhZRbK+xO!=",)) + .expect_err("invalid borrowed digest") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!(SecWebSocketAcceptOwned::try_from(&key).expect("TryFrom hashes"), accept); + } + + #[test] + fn conversions_and_header_round_trip_validate_exact_base64() { + const WIRE: &str = "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="; + for value in [ + SecWebSocketAcceptOwned::try_from(WIRE), + SecWebSocketAcceptOwned::try_from(String::from(WIRE)), + SecWebSocketAcceptOwned::try_from(FieldValue::from_static(WIRE)), + ] { + assert_eq!(value.expect("canonical digest").encoded(), WIRE.as_bytes()); + } + + let mut table = TestSink::new(); + assert!(SecWebSocketAccept::view(&table).expect("absent header succeeds").is_none()); + SecWebSocketAccept::insert(&mut table, SecWebSocketAcceptOwned::try_from(WIRE).expect("canonical digest")).expect("header inserts"); + let view = SecWebSocketAccept::view(&table).expect("view decodes").expect("header is present"); + assert_eq!(view.encoded(), WIRE.as_bytes()); + assert_eq!(view.as_field_value().as_bytes(), WIRE.as_bytes()); + assert_eq!( + SecWebSocketAccept::owned(&table) + .expect("owned decode succeeds") + .expect("header is present") + .encoded(), + WIRE.as_bytes() + ); + + for raw in [ + "s3pPLMBiTxaQ9kYGzzhZRbK+xOo", + "s3pPLMBiTxaQ9kYGzzhZRbK+xO!=", + "s3pPLMBiTxaQ9kYGzzhZRbK+xOp=", + ] { + assert_eq!( + SecWebSocketAcceptOwned::try_from(raw).expect_err("noncanonical digest").kind(), + DecodeErrorKind::InvalidSyntax + ); + } + assert_eq!( + SecWebSocketAcceptOwned::try_from(String::from("digest\n")) + .expect_err("invalid field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + SecWebSocketAcceptOwned::try_from("digest\n") + .expect_err("invalid borrowed field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn singleton_decoder_rejects_repeated_accept_values() { + let mut table = TestSink::new(); + table + .set_values( + SecWebSocketAccept::name(), + EncodedValues::from_vec(vec![ + FieldValue::from_static("s3pPLMBiTxaQ9kYGzzhZRbK+xOo="), + FieldValue::from_static("s3pPLMBiTxaQ9kYGzzhZRbK+xOo="), + ]), + ) + .expect("table accepts raw values"); + assert_eq!( + SecWebSocketAccept::view(&table).expect_err("singleton rejects duplicates").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + } +} diff --git a/crates/http_headers/src/headers/websocket/sec_web_socket_extensions.rs b/crates/http_headers/src/headers/websocket/sec_web_socket_extensions.rs new file mode 100644 index 000000000..7f27f6310 --- /dev/null +++ b/crates/http_headers/src/headers/websocket/sec_web_socket_extensions.rs @@ -0,0 +1,1419 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::ops::Range; +use std::{fmt, str}; + +use super::super::shared::FieldLinesIter; +use super::super::{ExtensionValue, invalid_syntax, trim_ows, value_from_bytes}; +use super::shared::{CommaItems, validate_header_value_list, validate_list}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Defines the `Sec-WebSocket-Extensions` header. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.2](https://www.rfc-editor.org/rfc/rfc6455#section-11.3.2). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{SecWebSocketExtensions, SecWebSocketExtensionsOwned}; +/// +/// let mut map = HeaderMap::new(); +/// SecWebSocketExtensions::insert( +/// &mut map, +/// SecWebSocketExtensionsOwned::try_from("permessage-deflate")?, +/// )?; +/// assert!(SecWebSocketExtensions::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct SecWebSocketExtensions { + _private: (), +} + +/// Owned value for the `Sec-WebSocket-Extensions` header. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.2]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::SecWebSocketExtensionsOwned::try_from("permessage-deflate")?; +/// assert_eq!(value.extensions().count(), 1); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Sec-WebSocket-Extensions: permessage-deflate` offers one extension. +/// `Sec-WebSocket-Extensions: permessage-deflate; client_max_window_bits=15` +/// also supplies an extension parameter. +/// +/// [RFC 6455 section 11.3.2]: https://www.rfc-editor.org/rfc/rfc6455#section-11.3.2 +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct SecWebSocketExtensionsOwned { + values: FieldLinesIter, +} + +/// Borrowed value for the `Sec-WebSocket-Extensions` header. +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{ +/// SecWebSocketExtensions, SecWebSocketExtensionsOwned, SecWebSocketExtensionsView, +/// }; +/// +/// let mut map = HeaderMap::new(); +/// let value = SecWebSocketExtensionsOwned::try_from("permessage-deflate")?; +/// SecWebSocketExtensions::insert(&mut map, value)?; +/// let view: SecWebSocketExtensionsView<'_> = +/// SecWebSocketExtensions::view(&map)?.expect("header is present"); +/// let extension = view +/// .extensions() +/// .next() +/// .transpose()? +/// .expect("one extension"); +/// assert_eq!(extension.name(), "permessage-deflate"); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +pub struct SecWebSocketExtensionsView<'a> { + values: FieldLines<'a>, +} + +/// Builder for one canonical `Sec-WebSocket-Extensions` field value. +#[derive(Clone, Debug, Default)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{SecWebSocketExtensionsBuilder, SecWebSocketExtensionsOwned}; +/// +/// let builder: SecWebSocketExtensionsBuilder = SecWebSocketExtensionsOwned::builder(); +/// let value = builder.extension("permessage-deflate").build()?; +/// assert_eq!(value.extensions().count(), 1); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct SecWebSocketExtensionsBuilder { + wire: Vec, + extension_count: usize, + pending_error: Option, +} + +/// One borrowed WebSocket extension. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{SecWebSocketExtensionsOwned, WebSocketExtensionView}; +/// +/// let value = +/// SecWebSocketExtensionsOwned::try_from("permessage-deflate; client_max_window_bits")?; +/// let extension: WebSocketExtensionView<'_> = value +/// .extensions() +/// .next() +/// .transpose()? +/// .expect("one extension"); +/// assert_eq!(extension.name(), "permessage-deflate"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct WebSocketExtensionView<'a> { + raw: &'a [u8], + name: &'a str, + parameter_start: usize, +} + +/// One borrowed WebSocket extension parameter. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{SecWebSocketExtensionsOwned, WebSocketExtensionParameterView}; +/// +/// let value = +/// SecWebSocketExtensionsOwned::try_from("permessage-deflate; client_max_window_bits")?; +/// let extension = value +/// .extensions() +/// .next() +/// .transpose()? +/// .expect("one extension"); +/// let parameter: WebSocketExtensionParameterView<'_> = extension +/// .parameters() +/// .next() +/// .transpose()? +/// .expect("one parameter"); +/// assert_eq!(parameter.name(), "client_max_window_bits"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct WebSocketExtensionParameterView<'a> { + raw: &'a [u8], + name: &'a str, + value: Option<&'a [u8]>, + quoted: bool, +} + +/// Iterator over the parameters of one WebSocket extension. +#[derive(Debug)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{SecWebSocketExtensionsOwned, WebSocketExtensionParameters}; +/// +/// let value = SecWebSocketExtensionsOwned::try_from( +/// "permessage-deflate; client_max_window_bits; server_max_window_bits=15", +/// )?; +/// let extension = value +/// .extensions() +/// .next() +/// .transpose()? +/// .expect("one extension"); +/// let mut parameters: WebSocketExtensionParameters<'_> = extension.parameters(); +/// assert_eq!( +/// parameters.next().transpose()?.expect("first").name(), +/// "client_max_window_bits" +/// ); +/// assert_eq!( +/// parameters +/// .next() +/// .transpose()? +/// .expect("second") +/// .value_str()?, +/// Some("15") +/// ); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct WebSocketExtensionParameters<'a> { + bytes: &'a [u8], + position: usize, +} + +super::super::shared::impl_value_count_debug!( + SecWebSocketExtensionsOwned => "SecWebSocketExtensionsOwned", + SecWebSocketExtensionsView<'_> => "SecWebSocketExtensionsView", +); + +impl fmt::Display for SecWebSocketExtensionsOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + super::super::shared::fmt_ascii_values(self.values.iter().map(FieldValue::as_field_value_ref), f) + } +} + +impl SecWebSocketExtensionsOwned { + #[cfg(all(feature = "serde", feature = "headers-websocket"))] + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.iter().map(FieldValue::as_field_value_ref) + } + + /// Creates an extension builder. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = SecWebSocketExtensionsOwned::builder() + /// .extension("permessage-deflate") + /// .build()?; + /// assert_eq!(value.extensions().count(), 1); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn builder() -> SecWebSocketExtensionsBuilder { + SecWebSocketExtensionsBuilder { + wire: Vec::new(), + extension_count: 0, + pending_error: None, + } + } + + /// Iterates extensions in wire order. + /// + /// # Errors + /// + /// An item contains [`DecodeErrorKind::InvalidSyntax`] when a stored + /// extension is malformed. Iteration resumes with the following item. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = SecWebSocketExtensionsOwned::try_from( + /// "permessage-deflate; client_max_window_bits=15, x-test", + /// )?; + /// let names = value + /// .extensions() + /// .map(|extension| Ok(extension?.name())) + /// .collect::, http_headers::DecodeError>>()?; + /// assert_eq!(names, ["permessage-deflate", "x-test"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn extensions(&self) -> impl Iterator, DecodeError>> { + self.values + .iter() + .flat_map(|value| CommaItems::new(value.as_bytes())) + .map(parse_extension) + } +} + +impl SecWebSocketExtensionsView<'_> { + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.repeated() + } + + /// Iterates extensions in wire order. + /// + /// # Errors + /// + /// An item contains [`DecodeErrorKind::InvalidSyntax`] when a stored + /// extension is malformed. Iteration resumes with the following item. + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), Box> { + /// use http::HeaderMap; + /// use http_headers::Field; + /// use http_headers::headers::{SecWebSocketExtensions, SecWebSocketExtensionsOwned}; + /// + /// let mut map = HeaderMap::new(); + /// let value = SecWebSocketExtensionsOwned::try_from("permessage-deflate, x-test")?; + /// SecWebSocketExtensions::insert(&mut map, value)?; + /// let view = SecWebSocketExtensions::view(&map)?.expect("header is present"); + /// let names = view + /// .extensions() + /// .map(|extension| Ok(extension?.name())) + /// .collect::, http_headers::DecodeError>>()?; + /// assert_eq!(names, ["permessage-deflate", "x-test"]); + /// # Ok::<(), Box>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn extensions(&self) -> impl Iterator, DecodeError>> { + self.values.comma_items().map(|item| item.and_then(parse_extension)) + } +} + +impl SecWebSocketExtensionsBuilder { + /// Starts another extension for validation by [`Self::build`]. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = SecWebSocketExtensionsOwned::builder() + /// .extension("permessage-deflate") + /// .extension("x-test") + /// .build()?; + /// let names = value + /// .extensions() + /// .map(|extension| Ok(extension?.name())) + /// .collect::, http_headers::DecodeError>>()?; + /// assert_eq!(names, ["permessage-deflate", "x-test"]); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[must_use] + pub fn extension(mut self, name: impl AsRef) -> Self { + let name = name.as_ref(); + if !validate::token(name.as_bytes()) { + self.record_error(DecodeErrorKind::InvalidToken); + } + if self.extension_count != 0 { + self.wire.extend_from_slice(b", "); + } + self.wire.extend_from_slice(name.as_bytes()); + if let Some(count) = self.extension_count.checked_add(1) { + self.extension_count = count; + } else { + self.record_error(DecodeErrorKind::InvalidNumber); + } + self + } + + /// Adds a flag parameter for validation by [`Self::build`]. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = SecWebSocketExtensionsOwned::builder() + /// .extension("permessage-deflate") + /// .parameter_flag("client_max_window_bits") + /// .build()?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// let parameter = extension + /// .parameters() + /// .next() + /// .transpose()? + /// .expect("one parameter"); + /// assert_eq!(parameter.name(), "client_max_window_bits"); + /// assert_eq!(parameter.value(), None); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[must_use] + pub fn parameter_flag(self, name: impl AsRef) -> Self { + self.parameter(name, ExtensionValue::Flag) + } + + /// Adds a parameter for validation by [`Self::build`]. + /// + /// A value-bearing parameter must contain an HTTP token. Validation, + /// including whether an extension precedes the parameter, occurs in + /// [`Self::build`]. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::{ExtensionValue, SecWebSocketExtensionsOwned}; + /// + /// let value = SecWebSocketExtensionsOwned::builder() + /// .extension("permessage-deflate") + /// .parameter("server_max_window_bits", ExtensionValue::Value("15")) + /// .build()?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// let parameter = extension + /// .parameters() + /// .next() + /// .transpose()? + /// .expect("one parameter"); + /// assert_eq!(parameter.name(), "server_max_window_bits"); + /// assert_eq!(parameter.value_str()?, Some("15")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[must_use] + pub fn parameter(mut self, name: impl AsRef, value: ExtensionValue<'_>) -> Self { + let name = name.as_ref(); + if self.extension_count == 0 || !validate::token(name.as_bytes()) { + self.record_error(DecodeErrorKind::InvalidToken); + } + self.wire.extend_from_slice(b"; "); + self.wire.extend_from_slice(name.as_bytes()); + match value { + ExtensionValue::Flag => {} + ExtensionValue::Value(value) => { + if !validate::token(value.as_bytes()) { + self.record_error(DecodeErrorKind::InvalidToken); + } + self.wire.push(b'='); + self.wire.extend_from_slice(value.as_bytes()); + } + } + self + } + + /// Adds a token-valued parameter to the most recently added extension. + #[must_use] + pub fn parameter_value(self, name: &str, value: &str) -> Self { + self.parameter(name, ExtensionValue::Value(value)) + } + + /// Adds a quoted parameter for validation by [`Self::build`]. + /// + /// WebSocket extension quoted strings must decode to an HTTP token. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = SecWebSocketExtensionsOwned::builder() + /// .extension("permessage-deflate") + /// .quoted_parameter("mode", "fast") + /// .build()?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// let parameter = extension + /// .parameters() + /// .next() + /// .transpose()? + /// .expect("one parameter"); + /// assert_eq!(parameter.name(), "mode"); + /// assert!(parameter.is_quoted()); + /// assert_eq!(parameter.value(), Some(&b"\"fast\""[..])); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[must_use] + pub fn quoted_parameter(mut self, name: &str, value: &str) -> Self { + if self.extension_count == 0 || !validate::token(name.as_bytes()) || !validate::token(value.as_bytes()) { + self.record_error(DecodeErrorKind::InvalidToken); + } + self.wire.extend_from_slice(b"; "); + self.wire.extend_from_slice(name.as_bytes()); + self.wire.extend_from_slice(b"=\""); + self.wire.extend_from_slice(value.as_bytes()); + self.wire.push(b'"'); + self + } + + /// Builds a nonempty extension list. + /// + /// # Errors + /// + /// Returns an error when no extension was added, an item is malformed, a + /// parameter has no preceding extension, or item counting overflowed. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let empty = SecWebSocketExtensionsOwned::builder().build(); + /// assert!(empty.is_err()); + /// + /// let value = SecWebSocketExtensionsOwned::builder() + /// .extension("permessage-deflate") + /// .build()?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// assert_eq!(extension.name(), "permessage-deflate"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn build(self) -> Result { + if let Some(kind) = self.pending_error { + return Err(DecodeError::new(&FieldName::SecWebSocketExtensions, kind)); + } + if self.extension_count == 0 { + return Err(invalid_syntax(&FieldName::SecWebSocketExtensions)); + } + let value = value_from_bytes(&FieldName::SecWebSocketExtensions, self.wire)?; + SecWebSocketExtensionsOwned::try_from(value) + } + + fn record_error(&mut self, kind: DecodeErrorKind) { + if self.pending_error.is_none() { + self.pending_error = Some(kind); + } + } +} + +impl<'a> WebSocketExtensionView<'a> { + /// Returns the extension token. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = + /// SecWebSocketExtensionsOwned::try_from("permessage-deflate; client_max_window_bits")?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// assert_eq!(extension.name(), "permessage-deflate"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn name(self) -> &'a str { + self.name + } + + /// Returns the complete extension bytes after surrounding OWS trimming. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = + /// SecWebSocketExtensionsOwned::try_from("permessage-deflate; client_max_window_bits")?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// assert_eq!( + /// extension.as_bytes(), + /// b"permessage-deflate; client_max_window_bits" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_bytes(self) -> &'a [u8] { + self.raw + } + + /// Iterates extension parameters. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = SecWebSocketExtensionsOwned::try_from( + /// "permessage-deflate; client_max_window_bits; server_max_window_bits=15", + /// )?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// let parameter_names = extension + /// .parameters() + /// .map(|parameter| Ok(parameter?.name())) + /// .collect::, http_headers::DecodeError>>()?; + /// assert_eq!( + /// parameter_names, + /// ["client_max_window_bits", "server_max_window_bits"] + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn parameters(self) -> WebSocketExtensionParameters<'a> { + WebSocketExtensionParameters { + bytes: self.raw, + position: self.parameter_start, + } + } +} + +impl<'a> WebSocketExtensionParameterView<'a> { + /// Returns the parameter name. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = + /// SecWebSocketExtensionsOwned::try_from("permessage-deflate; server_max_window_bits=15")?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// let parameter = extension + /// .parameters() + /// .next() + /// .transpose()? + /// .expect("one parameter"); + /// assert_eq!(parameter.name(), "server_max_window_bits"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn name(self) -> &'a str { + self.name + } + + /// Returns the raw token or quoted-string value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = SecWebSocketExtensionsOwned::try_from( + /// "permessage-deflate; client_max_window_bits; server_max_window_bits=15", + /// )?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// let mut parameters = extension.parameters(); + /// let flag = parameters.next().transpose()?.expect("flag parameter"); + /// assert_eq!(flag.value(), None); + /// let valued = parameters.next().transpose()?.expect("valued parameter"); + /// assert_eq!(valued.value(), Some(&b"15"[..])); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn value(self) -> Option<&'a [u8]> { + self.value + } + + /// Returns whether the value is a quoted string. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let token_value = SecWebSocketExtensionsOwned::try_from("permessage-deflate; mode=fast")?; + /// let token_parameter = token_value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension") + /// .parameters() + /// .next() + /// .transpose()? + /// .expect("one parameter"); + /// assert!(!token_parameter.is_quoted()); + /// + /// let quoted_value = SecWebSocketExtensionsOwned::try_from("permessage-deflate; mode=\"fast\"")?; + /// let quoted_parameter = quoted_value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension") + /// .parameters() + /// .next() + /// .transpose()? + /// .expect("one parameter"); + /// assert!(quoted_parameter.is_quoted()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn is_quoted(self) -> bool { + self.quoted + } + + /// Returns the value as UTF-8 without unescaping quoted strings. + /// + /// # Errors + /// + /// Returns an error when the value contains non-UTF-8 `obs-text`. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = + /// SecWebSocketExtensionsOwned::try_from("permessage-deflate; server_max_window_bits=15")?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// let parameter = extension + /// .parameters() + /// .next() + /// .transpose()? + /// .expect("one parameter"); + /// assert_eq!(parameter.value_str()?, Some("15")); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn value_str(self) -> Result, DecodeError> { + self.value + .map(str::from_utf8) + .transpose() + .map_err(|_invalid| DecodeError::new(&FieldName::SecWebSocketExtensions, DecodeErrorKind::InvalidUtf8)) + } + + /// Returns the complete parameter bytes after surrounding OWS trimming. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketExtensionsOwned; + /// + /// let value = + /// SecWebSocketExtensionsOwned::try_from("permessage-deflate; client_max_window_bits")?; + /// let extension = value + /// .extensions() + /// .next() + /// .transpose()? + /// .expect("one extension"); + /// let parameter = extension + /// .parameters() + /// .next() + /// .transpose()? + /// .expect("one parameter"); + /// assert_eq!(parameter.as_bytes(), b"client_max_window_bits"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_bytes(self) -> &'a [u8] { + self.raw + } +} + +impl<'a> Iterator for WebSocketExtensionParameters<'a> { + type Item = Result, DecodeError>; + + fn next(&mut self) -> Option { + match take_extension_parameter(self.bytes, &mut self.position) { + Ok(Some(parameter)) => Some(Ok(parameter)), + Ok(None) => None, + Err(error) => { + self.position = self.bytes.len(); + Some(Err(error)) + } + } + } +} + +impl Field for SecWebSocketExtensions { + type View<'a> = SecWebSocketExtensionsView<'a>; + type Owned = SecWebSocketExtensionsOwned; + + fn name() -> &'static FieldName { + &FieldName::SecWebSocketExtensions + } + + fn view_with(source: &S, _mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + validate_extensions(&lines)?; + Ok(Some(SecWebSocketExtensionsView { values: lines })) + } + + fn owned_with(source: &S, _mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_list_item_limit(b',', true)?; + lines.validate_list_item_limit(b';', false)?; + let mut copied = FieldLinesIter::empty(); + let mut present = false; + for (value_index, (value, owned)) in lines.repeated_owned()?.enumerate() { + present |= validate_extension_header_value(value).map_err(|error| { + if error.kind() == DecodeErrorKind::UnterminatedQuote { + error.at_value(value_index) + } else { + error + } + })?; + copied.push(owned); + } + if !present { + return Err(invalid_syntax(&FieldName::SecWebSocketExtensions)); + } + Ok(Some(SecWebSocketExtensionsOwned { values: copied })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), value.values.into_encoded()) + } +} + +super::super::shared::impl_string_conversions!( + SecWebSocketExtensionsOwned, + &FieldName::SecWebSocketExtensions, + invalid_syntax, + value +); + +impl TryFrom for SecWebSocketExtensionsOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + if !validate_extension_header_value(value.as_field_value_ref())? { + return Err(invalid_syntax(&FieldName::SecWebSocketExtensions)); + } + Ok(Self { + values: FieldLinesIter::one(value), + }) + } +} + +fn validate_extensions(values: &FieldLines<'_>) -> Result<(), DecodeError> { + values.validate_list_item_limit(b',', true)?; + values.validate_list_item_limit(b';', false)?; + let mut present = false; + for value in values.repeated() { + let Some(line_present) = validate_plain_extension_line(value.as_bytes())? else { + return validate_list(values, &mut validate_extension); + }; + present |= line_present; + } + present + .then_some(()) + .ok_or_else(|| DecodeError::new(&FieldName::SecWebSocketExtensions, DecodeErrorKind::MissingValue)) +} + +fn validate_extension_header_value(value: FieldValueRef<'_>) -> Result { + match validate_plain_extension_line(value.as_bytes())? { + Some(present) => Ok(present), + None => validate_header_value_list(value, &FieldName::SecWebSocketExtensions, validate_extension), + } +} + +/// Validates an unquoted extension list in one pass. +/// +/// `None` selects the structured parser when quoted-string syntax appears. +fn validate_plain_extension_line(bytes: &[u8]) -> Result, DecodeError> { + const WINDOW: &[u8; 42] = b"permessage-deflate; client_max_window_bits"; + const TAKEOVER: &[u8; 74] = b"permessage-deflate; server_no_context_takeover; client_no_context_takeover"; + + // Keep early prefix rejection; bounded suffix comparisons avoid an outlined byte comparison. + match bytes.len() { + 74 if bytes[..2] == *b"pe" + && bytes[2..18] == TAKEOVER[2..18] + && bytes[18..26] == TAKEOVER[18..26] + && bytes[26..42] == TAKEOVER[26..42] + && bytes[42..50] == TAKEOVER[42..50] + && bytes[50..66] == TAKEOVER[50..66] + && bytes[66..] == TAKEOVER[66..] => + { + return Ok(Some(true)); + } + 42 if bytes[..2] == *b"pe" + && bytes[2..18] == WINDOW[2..18] + && bytes[18..22] == WINDOW[18..22] + && bytes[22..38] == WINDOW[22..38] + && bytes[38..] == WINDOW[38..] => + { + return Ok(Some(true)); + } + 18 if bytes[..2] == *b"pe" && bytes[2..] == *b"rmessage-deflate" => return Ok(Some(true)), + _ => {} + } + if http_headers_simd::find_either(bytes, b'"', b'"').is_some() { + return Ok(None); + } + let mut position = 0_usize; + let mut present = false; + loop { + skip_ows(bytes, &mut position); + while bytes.get(position) == Some(&b',') { + position += 1; + skip_ows(bytes, &mut position); + } + if position == bytes.len() { + break; + } + if matches!(bytes.get(position), Some(b'"' | b'\\')) { + return Ok(None); + } + if !take_plain_extension_token(bytes, &mut position) { + return Err(DecodeError::new(&FieldName::SecWebSocketExtensions, DecodeErrorKind::InvalidToken)); + } + present = true; + + loop { + skip_ows(bytes, &mut position); + match bytes.get(position) { + None => return Ok(Some(present)), + Some(b',') => { + position += 1; + break; + } + Some(b'"' | b'\\') => return Ok(None), + Some(b';') => { + position += 1; + skip_ows(bytes, &mut position); + if matches!(bytes.get(position), Some(b'"' | b'\\')) { + return Ok(None); + } + if !take_plain_extension_token(bytes, &mut position) { + return Err(DecodeError::new(&FieldName::SecWebSocketExtensions, DecodeErrorKind::InvalidToken)); + } + if !take_plain_parameter_value(bytes, &mut position)? { + return Ok(None); + } + } + Some(_) => return Err(invalid_syntax(&FieldName::SecWebSocketExtensions)), + } + } + } + Ok(Some(present)) +} + +fn take_plain_parameter_value(bytes: &[u8], position: &mut usize) -> Result { + if bytes.get(*position) != Some(&b'=') { + return Ok(true); + } + *position += 1; + if matches!(bytes.get(*position), Some(b'"' | b'\\')) { + return Ok(false); + } + if !take_plain_extension_token(bytes, position) { + return Err(DecodeError::new(&FieldName::SecWebSocketExtensions, DecodeErrorKind::InvalidToken)); + } + Ok(true) +} + +fn validate_extension(bytes: &[u8]) -> Result<(), DecodeError> { + parse_extension(bytes).map(drop) +} + +#[inline] +fn take_plain_extension_token(bytes: &[u8], position: &mut usize) -> bool { + let start = *position; + while bytes.get(*position).is_some_and(|byte| validate::token_byte(*byte)) { + *position += 1; + } + start != *position +} + +fn parse_extension(bytes: &[u8]) -> Result, DecodeError> { + let mut position = 0_usize; + let name_range = take_token(bytes, &mut position)?; + let parameter_start = position; + while take_extension_parameter(bytes, &mut position)?.is_some() {} + let name = str::from_utf8(&bytes[name_range]).expect("HTTP token parsing guarantees an in-bounds ASCII range"); + Ok(WebSocketExtensionView { + raw: bytes, + name, + parameter_start, + }) +} + +fn take_extension_parameter<'a>(bytes: &'a [u8], position: &mut usize) -> Result>, DecodeError> { + skip_ows(bytes, position); + if *position == bytes.len() { + return Ok(None); + } + if bytes.get(*position) != Some(&b';') { + return Err(invalid_syntax(&FieldName::SecWebSocketExtensions)); + } + let raw_start = *position; + *position += 1; + skip_ows(bytes, position); + let name_range = take_token(bytes, position)?; + let (value, quoted) = if bytes.get(*position) == Some(&b'=') { + *position += 1; + if bytes.get(*position) == Some(&b'"') { + (Some(take_quoted_string(bytes, position)?), true) + } else { + (Some(take_token(bytes, position)?), false) + } + } else { + (None, false) + }; + let raw_end = *position; + let name = str::from_utf8(&bytes[name_range]).expect("HTTP token parsing guarantees an in-bounds ASCII range"); + let value = value.map(|range| &bytes[range]); + Ok(Some(WebSocketExtensionParameterView { + raw: trim_ows(&bytes[raw_start + 1..raw_end]), + name, + value, + quoted, + })) +} + +fn take_token(bytes: &[u8], position: &mut usize) -> Result, DecodeError> { + let start = *position; + while bytes.get(*position).is_some_and(|byte| validate::token_byte(*byte)) { + *position += 1; + } + if start == *position { + Err(DecodeError::new(&FieldName::SecWebSocketExtensions, DecodeErrorKind::InvalidToken)) + } else { + Ok(start..*position) + } +} + +fn take_quoted_string(bytes: &[u8], position: &mut usize) -> Result, DecodeError> { + let start = *position; + *position += 1; + let mut escaped = false; + let mut value_length = 0_usize; + while let Some(byte) = bytes.get(*position).copied() { + *position += 1; + if escaped { + if !validate::token_byte(byte) { + return Err(invalid_syntax(&FieldName::SecWebSocketExtensions)); + } + escaped = false; + value_length += 1; + } else if byte == b'\\' { + escaped = true; + } else if byte == b'"' { + return if value_length == 0 { + Err(DecodeError::new(&FieldName::SecWebSocketExtensions, DecodeErrorKind::InvalidToken)) + } else { + Ok(start..*position) + }; + } else if !validate::token_byte(byte) { + return Err(invalid_syntax(&FieldName::SecWebSocketExtensions)); + } else { + value_length += 1; + } + } + Err(DecodeError::new( + &FieldName::SecWebSocketExtensions, + DecodeErrorKind::UnterminatedQuote, + )) +} + +fn skip_ows(bytes: &[u8], position: &mut usize) { + while bytes.get(*position).is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + *position += 1; + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{ + CommaItems, SecWebSocketExtensions, SecWebSocketExtensionsOwned, WebSocketExtensionParameterView, WebSocketExtensionParameters, + WebSocketExtensionView, parse_extension, take_quoted_string, validate_plain_extension_line, + }; + use crate::sink::{EncodedValues, FieldSink}; + use crate::source::FieldSource; + use crate::{DecodeErrorKind, Field, FieldValue, TestSink}; + + #[test] + fn cached_extension_substitutions_match_structured_parsing() { + for literal in [ + b"permessage-deflate".as_slice(), + b"permessage-deflate; client_max_window_bits", + b"permessage-deflate; server_no_context_takeover; client_no_context_takeover", + ] { + let mut bytes = literal.to_vec(); + for index in 0..bytes.len() { + for replacement in crate::test_support::substitution_bytes(literal[index], index, literal.len()) { + bytes[index] = replacement; + let plain = validate_plain_extension_line(&bytes); + if plain != Ok(None) { + assert_eq!( + plain == Ok(Some(true)), + CommaItems::new(&bytes).all(|item| parse_extension(item).is_ok()), + "{literal:?}, index {index}, replacement {replacement}" + ); + } + } + bytes[index] = literal[index]; + } + } + } + + #[test] + fn builder_and_accessors_cover_flags_tokens_quotes_and_multiple_extensions() { + let extensions = SecWebSocketExtensionsOwned::builder() + .extension("permessage-deflate") + .parameter_flag("client_max_window_bits") + .parameter_value("server_max_window_bits", "15") + .quoted_parameter("mode", "fast") + .extension("x-test") + .build() + .expect("nonempty extension list"); + assert_eq!(format!("{extensions:?}"), "SecWebSocketExtensionsOwned { value_count: 1 }"); + + let parsed = extensions + .extensions() + .collect::, _>>() + .expect("built extensions parse"); + assert_eq!(parsed.len(), 2); + assert_eq!(parsed[0].name(), "permessage-deflate"); + assert_eq!( + parsed[0].as_bytes(), + b"permessage-deflate; client_max_window_bits; server_max_window_bits=15; mode=\"fast\"" + ); + let parameters = parsed[0].parameters().collect::, _>>().expect("parameters parse"); + assert_eq!(parameters[0].name(), "client_max_window_bits"); + assert_eq!(parameters[0].value(), None); + assert!(!parameters[0].is_quoted()); + assert_eq!(parameters[0].value_str(), Ok(None)); + assert_eq!(parameters[0].as_bytes(), b"client_max_window_bits"); + assert_eq!(parameters[1].name(), "server_max_window_bits"); + assert_eq!(parameters[1].value(), Some(b"15".as_slice())); + assert_eq!(parameters[1].value_str(), Ok(Some("15"))); + assert!(!parameters[1].is_quoted()); + assert_eq!(parameters[2].name(), "mode"); + assert_eq!(parameters[2].value(), Some(b"\"fast\"".as_slice())); + assert_eq!(parameters[2].value_str(), Ok(Some("\"fast\""))); + assert!(parameters[2].is_quoted()); + assert_eq!(parsed[1].name(), "x-test"); + assert_eq!(parsed[1].parameters().next(), None); + } + + #[test] + fn builder_rejects_missing_extensions_and_invalid_tokens() { + assert_eq!( + SecWebSocketExtensionsOwned::builder().build().expect_err("empty builder").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + SecWebSocketExtensionsOwned::builder() + .parameter_flag("flag") + .build() + .expect_err("parameter needs extension") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + SecWebSocketExtensionsOwned::builder() + .extension("not valid") + .build() + .expect_err("invalid extension token") + .kind(), + DecodeErrorKind::InvalidToken + ); + let builder = SecWebSocketExtensionsOwned::builder().extension("valid"); + assert_eq!( + builder + .clone() + .parameter_flag("bad name") + .build() + .expect_err("invalid parameter name") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + builder + .clone() + .parameter_value("mode", "bad value") + .build() + .expect_err("invalid parameter value") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + builder + .quoted_parameter("mode", "bad value") + .build() + .expect_err("quoted value must decode to token") + .kind(), + DecodeErrorKind::InvalidToken + ); + + let overflow = super::SecWebSocketExtensionsBuilder { + wire: Vec::new(), + extension_count: usize::MAX, + pending_error: None, + }; + assert_eq!( + overflow.extension("x").build().expect_err("extension count overflow").kind(), + DecodeErrorKind::InvalidNumber + ); + + let invalid_wire = super::SecWebSocketExtensionsBuilder { + wire: vec![b'\n'], + extension_count: 1, + pending_error: None, + }; + assert_eq!( + invalid_wire.build().expect_err("invalid private builder wire").kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn conversions_views_owned_and_insert_preserve_repeated_field_lines() { + for value in [ + SecWebSocketExtensionsOwned::try_from("permessage-deflate; client_max_window_bits=15"), + SecWebSocketExtensionsOwned::try_from(String::from("permessage-deflate; client_max_window_bits=15")), + SecWebSocketExtensionsOwned::try_from(FieldValue::from_static("permessage-deflate; client_max_window_bits=15")), + ] { + assert_eq!( + value + .expect("valid extension") + .extensions() + .next() + .expect("one extension") + .expect("extension parses") + .name(), + "permessage-deflate" + ); + } + assert_eq!( + SecWebSocketExtensionsOwned::try_from(String::from("extension\n")) + .expect_err("invalid field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + SecWebSocketExtensionsOwned::try_from("extension\n") + .expect_err("invalid borrowed field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + SecWebSocketExtensionsOwned::try_from(FieldValue::from_static("x; =bad")) + .expect_err("invalid stored parameter") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + SecWebSocketExtensionsOwned::try_from(FieldValue::from_static(",,,")) + .expect_err("empty stored list") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let mut table = TestSink::new(); + assert!(SecWebSocketExtensions::view(&table).expect("absent view succeeds").is_none()); + assert!(SecWebSocketExtensions::owned(&table).expect("absent owned succeeds").is_none()); + table + .set_values( + SecWebSocketExtensions::name(), + EncodedValues::from_vec(vec![ + FieldValue::from_static("permessage-deflate"), + FieldValue::from_static("x-test; mode=\"fast\", x-other"), + ]), + ) + .expect("table accepts extensions"); + let view = SecWebSocketExtensions::view(&table) + .expect("view decodes") + .expect("header is present"); + assert_eq!(view.field_values().count(), 2); + assert_eq!( + view.extensions() + .map(|extension| extension.map(WebSocketExtensionView::name)) + .collect::, _>>(), + Ok(vec!["permessage-deflate", "x-test", "x-other"]) + ); + assert_eq!(format!("{view:?}"), "SecWebSocketExtensionsView { value_count: 2 }"); + + let owned = SecWebSocketExtensions::owned(&table) + .expect("owned decode succeeds") + .expect("header is present"); + let mut output = TestSink::new(); + SecWebSocketExtensions::insert(&mut output, owned).expect("extensions insert"); + assert_eq!(output.lines(SecWebSocketExtensions::name()).expect("inserted extensions").len(), 2); + } + + #[test] + fn plain_and_structured_parsers_classify_errors_and_quotes() { + assert_eq!(validate_plain_extension_line(b"x-test; flag"), Ok(Some(true))); + assert_eq!(validate_plain_extension_line(b"permessage-deflate"), Ok(Some(true))); + assert_eq!(validate_plain_extension_line(b",, \t"), Ok(Some(false))); + assert_eq!(validate_plain_extension_line(b"x; p=\"v\""), Ok(None)); + assert_eq!(validate_plain_extension_line(b"\\x"), Ok(None)); + assert_eq!(validate_plain_extension_line(b"x; \\p"), Ok(None)); + assert_eq!(validate_plain_extension_line(b"x; p=\\v"), Ok(None)); + assert_eq!(validate_plain_extension_line(b"x\\"), Ok(None)); + assert_eq!(validate_plain_extension_line(b"x,y"), Ok(Some(true))); + assert_eq!( + validate_plain_extension_line(b"=x").expect_err("missing extension name").kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + validate_plain_extension_line(b"x; =v").expect_err("missing parameter name").kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + validate_plain_extension_line(b"x; p=").expect_err("missing parameter value").kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + validate_plain_extension_line(b"x p").expect_err("missing delimiter").kind(), + DecodeErrorKind::InvalidSyntax + ); + + let quoted = parse_extension(b"x; mode=\"fa\\st\"").expect("escaped token parses"); + let parameter = quoted.parameters().next().expect("one parameter").expect("parameter parses"); + assert!(parameter.is_quoted()); + assert_eq!(parameter.value(), Some(b"\"fa\\st\"".as_slice())); + + for raw in [ + b"".as_slice(), + b"x; mode=\"\"".as_slice(), + b"x; mode=\"bad value\"".as_slice(), + b"x; mode=\"unterminated".as_slice(), + b"x; mode=\"trail\\".as_slice(), + b"x; mode=\"bad\\ \"".as_slice(), + b"x; mode=".as_slice(), + ] { + let _error = parse_extension(raw).expect_err("malformed extension must fail"); + } + + let mut position = 0; + let _range = take_quoted_string(b"\"ok\"", &mut position).expect("quoted token is valid"); + + let invalid_utf8 = WebSocketExtensionParameterView { + raw: b"mode=\xff", + name: "mode", + value: Some(&[0xff]), + quoted: false, + }; + assert_eq!( + invalid_utf8.value_str().expect_err("non-UTF-8 value").kind(), + DecodeErrorKind::InvalidUtf8 + ); + } + + #[test] + fn parameter_iterator_stops_after_error_and_owned_quote_errors_keep_line_index() { + let mut parameters = WebSocketExtensionParameters { + bytes: b"x; =bad", + position: 1, + }; + assert_eq!( + parameters.next().expect("one error").expect_err("missing name").kind(), + DecodeErrorKind::InvalidToken + ); + assert!(parameters.next().is_none()); + + let mut parameters = WebSocketExtensionParameters { + bytes: b"x bad", + position: 1, + }; + assert_eq!( + parameters + .next() + .expect("one error") + .expect_err("parameter delimiter is required") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let mut table = TestSink::new(); + table + .set_values( + SecWebSocketExtensions::name(), + EncodedValues::from_vec(vec![ + FieldValue::from_static("x"), + FieldValue::from_static("y; mode=\"unterminated"), + ]), + ) + .expect("table accepts raw extensions"); + let error = SecWebSocketExtensions::owned(&table).expect_err("unterminated quote is rejected"); + assert_eq!(error.kind(), DecodeErrorKind::UnterminatedQuote); + assert_eq!(error.value_index(), Some(1)); + + table + .set_values( + SecWebSocketExtensions::name(), + EncodedValues::single(FieldValue::from_static("x; =bad")), + ) + .expect("table accepts malformed extension"); + assert_eq!( + SecWebSocketExtensions::owned(&table) + .expect_err("invalid parameter is rejected") + .kind(), + DecodeErrorKind::InvalidToken + ); + + table + .set_values( + SecWebSocketExtensions::name(), + EncodedValues::single(FieldValue::from_static(",,, \t")), + ) + .expect("table accepts empty list"); + assert_eq!( + SecWebSocketExtensions::view(&table).expect_err("empty list is rejected").kind(), + DecodeErrorKind::MissingValue + ); + + table + .set_values( + SecWebSocketExtensions::name(), + EncodedValues::single(FieldValue::from_static("x; =bad")), + ) + .expect("table accepts invalid plain extension"); + assert_eq!( + SecWebSocketExtensions::view(&table) + .expect_err("plain parser error propagates") + .kind(), + DecodeErrorKind::InvalidToken + ); + } +} diff --git a/crates/http_headers/src/headers/websocket/sec_web_socket_key.rs b/crates/http_headers/src/headers/websocket/sec_web_socket_key.rs new file mode 100644 index 000000000..b6f66c3ae --- /dev/null +++ b/crates/http_headers/src/headers/websocket/sec_web_socket_key.rs @@ -0,0 +1,332 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::super::invalid_syntax; +use super::shared::{encode_fixed_base64, validate_canonical_base64}; +use crate::{DecodeError, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +/// Defines the `Sec-WebSocket-Key` header. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.1](https://www.rfc-editor.org/rfc/rfc6455#section-11.3.1). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{SecWebSocketKey, SecWebSocketKeyOwned}; +/// +/// let mut map = HeaderMap::new(); +/// SecWebSocketKey::insert( +/// &mut map, +/// SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ==")?, +/// )?; +/// assert!(SecWebSocketKey::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct SecWebSocketKey { + _private: (), +} + +/// Owned value for the `Sec-WebSocket-Key` header. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.1]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ==")?; +/// assert_eq!(value.encoded(), b"dGhlIHNhbXBsZSBub25jZQ=="); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Sec-WebSocket-Key: dGhlIHNhbXBsZSBub25jZQ==` is the RFC example of a +/// base64-encoded 16-byte nonce. +/// +/// [RFC 6455 section 11.3.1]: https://www.rfc-editor.org/rfc/rfc6455#section-11.3.1 +#[derive(Clone, Debug, Eq, Hash, PartialEq)] +pub struct SecWebSocketKeyOwned { + value: FieldValue, +} + +/// Borrowed value for the `Sec-WebSocket-Key` header. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +/// # Examples +/// +/// ```rust +/// use http_headers::headers::{SecWebSocketKey, SecWebSocketKeyView}; +/// use http_headers::{FieldValueRef, SingleValueField}; +/// +/// let view: SecWebSocketKeyView<'_> = +/// SecWebSocketKey::decode_view(FieldValueRef::new(b"dGhlIHNhbXBsZSBub25jZQ=="))?; +/// assert_eq!(view.encoded(), b"dGhlIHNhbXBsZSBub25jZQ=="); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct SecWebSocketKeyView<'a> { + value: FieldValueRef<'a>, +} + +impl SecWebSocketKeyOwned { + /// Encodes a 16-byte client nonce. + /// + /// The nonce must be generated independently for every connection by a + /// cryptographically secure random number generator. Predictable or reused + /// nonces defeat the handshake's protection against intermediary cache + /// poisoning; this function encodes the nonce but does not generate it. + /// + /// # Errors + /// + /// The fixed-size nonce always has a representable canonical encoding. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketKeyOwned; + /// + /// let value = SecWebSocketKeyOwned::from_nonce(*b"the sample nonce")?; + /// assert_eq!(value.encoded(), b"dGhlIHNhbXBsZSBub25jZQ=="); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::unnecessary_wraps, + reason = "WebSocket value constructors consistently report validation through DecodeError" + )] + pub fn from_nonce(nonce: [u8; 16]) -> Result { + Ok(Self { + value: encode_fixed_base64(&nonce), + }) + } + + /// Returns the canonical base64 bytes. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketKeyOwned; + /// + /// let value = SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ==")?; + /// assert_eq!(value.encoded(), b"dGhlIHNhbXBsZSBub25jZQ=="); + /// assert!(SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ=").is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn encoded(&self) -> &[u8] { + self.value.as_bytes() + } + + /// Returns the stored field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketKeyOwned; + /// + /// let value = SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ==")?; + /// assert_eq!( + /// value.as_field_value().as_bytes(), + /// b"dGhlIHNhbXBsZSBub25jZQ==" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn as_field_value(&self) -> &FieldValue { + &self.value + } + + /// Returns reusable wire storage. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketKeyOwned; + /// + /// let value = SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ==")?; + /// let field_value = value.into_field_value(); + /// assert_eq!(field_value.as_bytes(), b"dGhlIHNhbXBsZSBub25jZQ=="); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn into_field_value(self) -> FieldValue { + self.into() + } +} + +super::super::shared::impl_field_value_conversion!(SecWebSocketKeyOwned, |value| value.value); + +impl<'a> SecWebSocketKeyView<'a> { + /// Returns the canonical base64 bytes. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketKey; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = SecWebSocketKey::decode_view(FieldValueRef::new(b"dGhlIHNhbXBsZSBub25jZQ=="))?; + /// assert_eq!(view.encoded(), b"dGhlIHNhbXBsZSBub25jZQ=="); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn encoded(self) -> &'a [u8] { + self.value.as_bytes() + } + + /// Returns the original field value. + #[must_use] + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketKey; + /// use http_headers::{FieldValueRef, SingleValueField}; + /// + /// let view = SecWebSocketKey::decode_view(FieldValueRef::new(b"dGhlIHNhbXBsZSBub25jZQ=="))?; + /// assert_eq!( + /// view.as_field_value().as_bytes(), + /// b"dGhlIHNhbXBsZSBub25jZQ==" + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub const fn as_field_value(self) -> FieldValueRef<'a> { + self.value + } +} + +impl SingleValueField for SecWebSocketKey { + type View<'a> = SecWebSocketKeyView<'a>; + type Owned = SecWebSocketKeyOwned; + + fn name() -> &'static FieldName { + &FieldName::SecWebSocketKey + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + validate_canonical_base64(value.as_bytes(), 16, &FieldName::SecWebSocketKey)?; + Ok(SecWebSocketKeyView { value }) + } + + fn decode_owned(value: FieldValue) -> Result { + SecWebSocketKeyOwned::try_from(value) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.value + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.value + } +} + +super::super::shared::impl_string_conversions!(SecWebSocketKeyOwned, &FieldName::SecWebSocketKey, invalid_syntax, value); + +impl TryFrom for SecWebSocketKeyOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + validate_canonical_base64(value.as_bytes(), 16, &FieldName::SecWebSocketKey)?; + Ok(Self { value }) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{SecWebSocketKey, SecWebSocketKeyOwned}; + use crate::sink::{EncodedValues, FieldSink}; + use crate::{DecodeErrorKind, Field, FieldValue, FieldValueRef, TestSink}; + + #[test] + fn nonce_and_accessors_produce_canonical_base64() { + let zero = SecWebSocketKeyOwned::from_nonce([0; 16]).expect("nonce encodes"); + assert_eq!(zero.encoded(), b"AAAAAAAAAAAAAAAAAAAAAA=="); + assert_eq!(zero.as_field_value().as_bytes(), zero.encoded()); + assert_eq!(zero.clone().into_field_value().as_bytes(), zero.encoded()); + assert_eq!( + ::as_field_value(&zero).as_bytes(), + zero.encoded() + ); + assert_eq!( + ::into_field_value(zero.clone()).as_bytes(), + zero.encoded() + ); + assert_eq!( + ::decode_owned(FieldValue::from_static("AAAAAAAAAAAAAAAAAAAAAA==",)) + .expect("direct owned decode"), + zero + ); + assert_eq!( + ::decode_view(FieldValueRef::new(b"AAAAAAAAAAAAAAAAAAAAA!==",)) + .expect_err("invalid borrowed nonce") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn conversions_and_header_round_trip_validate_exact_base64() { + const WIRE: &str = "dGhlIHNhbXBsZSBub25jZQ=="; + for value in [ + SecWebSocketKeyOwned::try_from(WIRE), + SecWebSocketKeyOwned::try_from(String::from(WIRE)), + SecWebSocketKeyOwned::try_from(FieldValue::from_static(WIRE)), + ] { + assert_eq!(value.expect("canonical nonce").encoded(), WIRE.as_bytes()); + } + + let mut table = TestSink::new(); + assert!(SecWebSocketKey::view(&table).expect("absent header succeeds").is_none()); + SecWebSocketKey::insert(&mut table, SecWebSocketKeyOwned::try_from(WIRE).expect("canonical nonce")).expect("header inserts"); + let view = SecWebSocketKey::view(&table).expect("view decodes").expect("header is present"); + assert_eq!(view.encoded(), WIRE.as_bytes()); + assert_eq!(view.as_field_value().as_bytes(), WIRE.as_bytes()); + assert_eq!( + SecWebSocketKey::owned(&table) + .expect("owned decode succeeds") + .expect("header is present") + .encoded(), + WIRE.as_bytes() + ); + + for raw in ["dGhlIHNhbXBsZSBub25jZQ=", "dGhlIHNhbXBsZSBub25jZ!==", "dGhlIHNhbXBsZSBub25jZR=="] { + assert_eq!( + SecWebSocketKeyOwned::try_from(raw).expect_err("noncanonical nonce").kind(), + DecodeErrorKind::InvalidSyntax + ); + } + assert_eq!( + SecWebSocketKeyOwned::try_from(String::from("nonce\n")) + .expect_err("invalid field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + SecWebSocketKeyOwned::try_from("nonce\n") + .expect_err("invalid borrowed field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn singleton_decoder_rejects_repeated_keys() { + let mut table = TestSink::new(); + table + .set_values( + SecWebSocketKey::name(), + EncodedValues::from_vec(vec![ + FieldValue::from_static("dGhlIHNhbXBsZSBub25jZQ=="), + FieldValue::from_static("dGhlIHNhbXBsZSBub25jZQ=="), + ]), + ) + .expect("table accepts raw values"); + assert_eq!( + SecWebSocketKey::view(&table).expect_err("singleton rejects duplicates").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + } +} diff --git a/crates/http_headers/src/headers/websocket/sec_web_socket_protocol.rs b/crates/http_headers/src/headers/websocket/sec_web_socket_protocol.rs new file mode 100644 index 000000000..f90027c7d --- /dev/null +++ b/crates/http_headers/src/headers/websocket/sec_web_socket_protocol.rs @@ -0,0 +1,604 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::{fmt, str}; + +use http_headers_simd::{EmptyMembers, TokenListScan}; + +use super::super::invalid_syntax; +use super::super::shared::FieldLinesIter; +use super::shared::{CommaItems, validate_bare_list, validate_header_value_list}; +use crate::sink::{FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef, validate}; + +/// Defines the `Sec-WebSocket-Protocol` header. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.4](https://www.rfc-editor.org/rfc/rfc6455#section-11.3.4). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{SecWebSocketProtocol, SecWebSocketProtocolOwned}; +/// +/// let mut map = HeaderMap::new(); +/// SecWebSocketProtocol::insert(&mut map, SecWebSocketProtocolOwned::new("chat")?)?; +/// assert!(SecWebSocketProtocol::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct SecWebSocketProtocol { + _private: (), +} + +/// Owned value for the `Sec-WebSocket-Protocol` header. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.4]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::SecWebSocketProtocolOwned::try_from("chat, superchat")?; +/// assert_eq!(value.protocols().count(), 2); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +/// +/// `Sec-WebSocket-Protocol: chat, superchat` offers two protocols in a +/// request; `Sec-WebSocket-Protocol: chat` selects one in a response. +/// +/// [RFC 6455 section 11.3.4]: https://www.rfc-editor.org/rfc/rfc6455#section-11.3.4 +#[derive(Clone, Eq, Hash, PartialEq)] +pub struct SecWebSocketProtocolOwned { + values: FieldLinesIter, +} + +/// Borrowed value for the `Sec-WebSocket-Protocol` header. +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http_headers::Field; +/// use http_headers::headers::{ +/// SecWebSocketProtocol, SecWebSocketProtocolOwned, SecWebSocketProtocolView, +/// }; +/// +/// let mut map = http::HeaderMap::new(); +/// SecWebSocketProtocol::insert(&mut map, SecWebSocketProtocolOwned::new("chat")?)?; +/// let view: SecWebSocketProtocolView<'_> = +/// SecWebSocketProtocol::view(&map)?.expect("protocol header is present"); +/// assert_eq!( +/// view.protocols().collect::, _>>()?, +/// vec!["chat"] +/// ); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +pub struct SecWebSocketProtocolView<'a> { + values: FieldLines<'a>, +} + +super::super::shared::impl_value_count_debug!( + SecWebSocketProtocolOwned => "SecWebSocketProtocolOwned", + SecWebSocketProtocolView<'_> => "SecWebSocketProtocolView", +); + +impl fmt::Display for SecWebSocketProtocolOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + super::super::shared::fmt_ascii_values(self.values.iter().map(FieldValue::as_field_value_ref), f) + } +} + +impl SecWebSocketProtocolOwned { + #[cfg(all(feature = "serde", feature = "headers-websocket"))] + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.iter().map(FieldValue::as_field_value_ref) + } + + /// Constructs a list containing one subprotocol. + /// + /// # Errors + /// + /// Returns an error when `protocol` is not an HTTP token. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketProtocolOwned; + /// + /// let value = SecWebSocketProtocolOwned::new("chat")?; + /// assert_eq!(value.selected()?, "chat"); + /// assert!(SecWebSocketProtocolOwned::new("not valid").is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn new(protocol: impl AsRef) -> Result { + let protocol = protocol.as_ref(); + validate_protocol(protocol.as_bytes())?; + Ok(Self { + values: FieldLinesIter::one(token_field_value(protocol)), + }) + } + + /// Adds a subprotocol as a separate field line. + /// + /// # Errors + /// + /// Returns an error when `protocol` is not an HTTP token. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketProtocolOwned; + /// + /// let value = SecWebSocketProtocolOwned::new("chat")?.with_protocol("superchat")?; + /// assert_eq!( + /// value.protocols().collect::, _>>()?, + /// vec!["chat", "superchat"] + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn with_protocol(mut self, protocol: &str) -> Result { + validate_protocol(protocol.as_bytes())?; + self.values.push(token_field_value(protocol)); + Ok(self) + } + + /// Iterates subprotocol tokens in wire order. + /// + /// # Errors + /// + /// An item contains [`DecodeErrorKind::InvalidToken`] when a stored + /// protocol is malformed. Iteration resumes with the following protocol. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketProtocolOwned; + /// + /// let one = SecWebSocketProtocolOwned::new("chat")?; + /// assert_eq!( + /// one.protocols().collect::, _>>()?, + /// vec!["chat"] + /// ); + /// + /// let several = SecWebSocketProtocolOwned::try_from("chat, superchat")?; + /// assert_eq!( + /// several.protocols().collect::, _>>()?, + /// vec!["chat", "superchat"] + /// ); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn protocols(&self) -> impl Iterator> { + self.values + .iter() + .flat_map(|value| CommaItems::new(value.as_bytes())) + .map(protocol_str) + } + + /// Returns the single selected subprotocol. + /// + /// # Errors + /// + /// Returns an error when the header is empty, contains multiple + /// subprotocols, or stored wire data is invalid. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketProtocolOwned; + /// + /// let value = SecWebSocketProtocolOwned::new("chat")?; + /// assert_eq!(value.selected()?, "chat"); + /// + /// let offered = SecWebSocketProtocolOwned::try_from("chat, superchat")?; + /// assert!(offered.selected().is_err()); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn selected(&self) -> Result<&str, DecodeError> { + one_protocol(&mut self.protocols()) + } +} + +impl SecWebSocketProtocolView<'_> { + pub(crate) fn field_values(&self) -> impl Iterator> + '_ { + self.values.repeated() + } + + /// Iterates subprotocol tokens in wire order. + /// + /// # Errors + /// + /// An item contains [`DecodeErrorKind::InvalidToken`] when a stored + /// protocol is malformed. Iteration resumes with the following protocol. + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), Box> { + /// use http_headers::Field; + /// use http_headers::headers::{SecWebSocketProtocol, SecWebSocketProtocolOwned}; + /// + /// let mut map = http::HeaderMap::new(); + /// let offered = SecWebSocketProtocolOwned::new("chat")?.with_protocol("superchat")?; + /// SecWebSocketProtocol::insert(&mut map, offered)?; + /// let view = SecWebSocketProtocol::view(&map)?.expect("protocol header is present"); + /// assert_eq!( + /// view.protocols().collect::, _>>()?, + /// vec!["chat", "superchat"] + /// ); + /// # Ok::<(), Box>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn protocols(&self) -> impl Iterator> { + self.values.comma_items().map(|item| { + item.and_then(|item| { + validate_protocol(item)?; + Ok(token_str(item)) + }) + }) + } + + /// Returns the single selected subprotocol. + /// + /// # Errors + /// + /// Returns an error when the header is empty or contains multiple + /// subprotocols. + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() -> Result<(), Box> { + /// use http_headers::Field; + /// use http_headers::headers::{SecWebSocketProtocol, SecWebSocketProtocolOwned}; + /// + /// let mut map = http::HeaderMap::new(); + /// SecWebSocketProtocol::insert(&mut map, SecWebSocketProtocolOwned::new("chat")?)?; + /// let view = SecWebSocketProtocol::view(&map)?.expect("protocol header is present"); + /// assert_eq!(view.selected()?, "chat"); + /// # Ok::<(), Box>(()) + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + pub fn selected(&self) -> Result<&str, DecodeError> { + one_protocol(&mut self.protocols()) + } +} + +impl Field for SecWebSocketProtocol { + type View<'a> = SecWebSocketProtocolView<'a>; + type Owned = SecWebSocketProtocolOwned; + + fn name() -> &'static FieldName { + &FieldName::SecWebSocketProtocol + } + + fn view_with(source: &S, _mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + validate_bare_list(&lines, is_bare_protocol_list, validate_protocol)?; + Ok(Some(SecWebSocketProtocolView { values: lines })) + } + + fn owned_with(source: &S, _mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines.validate_list_item_limit(b',', true)?; + let mut copied = FieldLinesIter::empty(); + let mut present = false; + for (value, owned) in lines.repeated_owned()? { + present |= is_bare_protocol_list(value.as_bytes()) + || validate_header_value_list(value, &FieldName::SecWebSocketProtocol, validate_protocol)?; + copied.push(owned); + } + if !present { + return Err(invalid_syntax(&FieldName::SecWebSocketProtocol)); + } + Ok(Some(SecWebSocketProtocolOwned { values: copied })) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), value.values.into_encoded()) + } +} + +super::super::shared::impl_string_conversions!(SecWebSocketProtocolOwned, &FieldName::SecWebSocketProtocol, invalid_syntax, value); + +impl TryFrom for SecWebSocketProtocolOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + if !is_bare_protocol_list(value.as_bytes()) + && !validate_header_value_list(value.as_field_value_ref(), &FieldName::SecWebSocketProtocol, validate_protocol)? + { + return Err(invalid_syntax(&FieldName::SecWebSocketProtocol)); + } + Ok(Self { + values: FieldLinesIter::one(value), + }) + } +} + +fn one_protocol<'a>(protocols: &mut dyn Iterator>) -> Result<&'a str, DecodeError> { + let protocol = protocols + .next() + .ok_or_else(|| DecodeError::new(&FieldName::SecWebSocketProtocol, DecodeErrorKind::MissingValue))??; + if let Some(protocol) = protocols.next() { + let _protocol = protocol?; + Err(DecodeError::new( + &FieldName::SecWebSocketProtocol, + DecodeErrorKind::UnexpectedMultipleValues, + )) + } else { + Ok(protocol) + } +} + +/// Returns whether a field line is a list of bare subprotocol tokens. +/// +/// Subprotocols are tokens, so a line whose every comma-separated member is a +/// token needs no further parse: the general splitter would find the same +/// members, accept each of them, and skip the empty ones exactly as `#rule` +/// expansion does. Requiring at least one member keeps the missing-value +/// diagnostic with the splitter, which reports it for the whole header rather +/// than for one line. +fn is_bare_protocol_list(bytes: &[u8]) -> bool { + match bytes.len() { + 10 if bytes[..2] == *b"gr" && bytes[2..] == *b"aphql-ws" => return true, + 15 if bytes[..2] == *b"ch" && bytes[2..10] == *b"at, supe" && bytes[7..] == *b"uperchat" => return true, + 20 if bytes[..2] == *b"gr" && bytes[2..] == *b"aphql-transport-ws" => return true, + 32 if bytes[..2] == *b"gr" && bytes[2..] == *b"aphql-transport-ws, graphql-ws" => return true, + _ => {} + } + http_headers_simd::scan_token_list(bytes, EmptyMembers::Skip) == TokenListScan::Members +} + +fn validate_protocol(bytes: &[u8]) -> Result<(), DecodeError> { + if validate::token(bytes) { + Ok(()) + } else { + Err(DecodeError::new(&FieldName::SecWebSocketProtocol, DecodeErrorKind::InvalidToken)) + } +} + +fn protocol_str(bytes: &[u8]) -> Result<&str, DecodeError> { + validate_protocol(bytes)?; + Ok(token_str(bytes)) +} + +fn token_field_value(token: &str) -> FieldValue { + FieldValue::from_str(token).expect("HTTP token validation guarantees a safe field value") +} + +fn token_str(bytes: &[u8]) -> &str { + str::from_utf8(bytes).expect("HTTP token validation guarantees ASCII") +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use http_headers_simd::{EmptyMembers, TokenListScan, scan_token_list}; + + use super::{ + SecWebSocketProtocol, SecWebSocketProtocolOwned, SecWebSocketProtocolView, is_bare_protocol_list, one_protocol, protocol_str, + }; + use crate::sink::{EncodedValues, FieldSink}; + use crate::source::FieldSource; + use crate::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValue, TestSink}; + + #[test] + fn cached_protocol_substitutions_match_token_scanning() { + for literal in [ + b"graphql-ws".as_slice(), + b"chat, superchat", + b"graphql-transport-ws", + b"graphql-transport-ws, graphql-ws", + ] { + let mut bytes = literal.to_vec(); + for index in 0..bytes.len() { + for replacement in crate::test_support::substitution_bytes(literal[index], index, literal.len()) { + bytes[index] = replacement; + assert_eq!( + is_bare_protocol_list(&bytes), + scan_token_list(&bytes, EmptyMembers::Skip) == TokenListScan::Members, + "{literal:?}, index {index}, replacement {replacement}" + ); + } + bytes[index] = literal[index]; + } + } + } + + #[test] + fn constructors_accessors_and_conversions_preserve_protocol_order() { + let protocols = SecWebSocketProtocolOwned::new("chat") + .expect("token protocol") + .with_protocol("superchat") + .expect("second token protocol"); + assert_eq!(protocols.protocols().collect::, _>>(), Ok(vec!["chat", "superchat"])); + assert_eq!( + protocols.selected().expect_err("multiple protocols are not a selection").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!(format!("{protocols:?}"), "SecWebSocketProtocolOwned { value_count: 2 }"); + assert_eq!(SecWebSocketProtocolOwned::new("chat").expect("one protocol").selected(), Ok("chat")); + + for value in [ + SecWebSocketProtocolOwned::try_from("chat, superchat"), + SecWebSocketProtocolOwned::try_from(String::from("chat, superchat")), + SecWebSocketProtocolOwned::try_from(FieldValue::from_static("chat, superchat")), + ] { + assert_eq!( + value.expect("valid list").protocols().collect::, _>>(), + Ok(vec!["chat", "superchat"]) + ); + } + assert_eq!( + SecWebSocketProtocolOwned::new("not valid") + .expect_err("spaces are not token bytes") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + SecWebSocketProtocolOwned::new("chat") + .expect("valid first protocol") + .with_protocol("not valid") + .expect_err("invalid appended protocol") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + SecWebSocketProtocolOwned::try_from(String::from("chat\n")) + .expect_err("invalid field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + + #[test] + fn view_owned_insert_and_slow_list_validation_round_trip() { + let mut table = TestSink::new(); + assert!(SecWebSocketProtocol::view(&table).expect("absent view succeeds").is_none()); + assert!(SecWebSocketProtocol::owned(&table).expect("absent owned succeeds").is_none()); + table + .set_values( + SecWebSocketProtocol::name(), + EncodedValues::from_vec(vec![ + FieldValue::from_static("chat"), + FieldValue::from_static("graphql-ws, superchat"), + ]), + ) + .expect("table accepts protocols"); + + let view = SecWebSocketProtocol::view(&table) + .expect("view decodes") + .expect("header is present"); + assert_eq!(view.field_values().count(), 2); + assert_eq!( + view.protocols().collect::, _>>(), + Ok(vec!["chat", "graphql-ws", "superchat"]) + ); + assert_eq!( + view.selected().expect_err("multiple values are not selected").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!(format!("{view:?}"), "SecWebSocketProtocolView { value_count: 2 }"); + + let owned = SecWebSocketProtocol::owned(&table) + .expect("owned decode succeeds") + .expect("header is present"); + let mut output = TestSink::new(); + SecWebSocketProtocol::insert(&mut output, owned).expect("owned protocols insert"); + assert_eq!(output.lines(SecWebSocketProtocol::name()).expect("inserted protocols").len(), 2); + } + + #[test] + fn selection_and_validation_report_empty_invalid_and_following_errors() { + assert_eq!( + one_protocol(&mut std::iter::empty()).expect_err("empty iterator").kind(), + DecodeErrorKind::MissingValue + ); + let first_error = DecodeError::new(&FieldName::SecWebSocketProtocol, DecodeErrorKind::InvalidToken); + assert_eq!( + one_protocol(&mut std::iter::once(Err(first_error))) + .expect_err("first error propagates") + .kind(), + DecodeErrorKind::InvalidToken + ); + let late_error = DecodeError::new(&FieldName::SecWebSocketProtocol, DecodeErrorKind::InvalidToken); + assert_eq!( + one_protocol(&mut [Ok("chat"), Err(late_error)].into_iter()) + .expect_err("second error propagates") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert_eq!( + protocol_str(&[0xff]).expect_err("non-token byte").kind(), + DecodeErrorKind::InvalidToken + ); + let malformed_view = SecWebSocketProtocolView { + values: crate::source::FieldLines::single(&FieldName::SecWebSocketProtocol, b"not valid"), + }; + assert_eq!( + malformed_view + .protocols() + .next() + .expect("one item") + .expect_err("private malformed storage") + .kind(), + DecodeErrorKind::InvalidToken + ); + assert!(is_bare_protocol_list(b"chat, superchat")); + assert!(is_bare_protocol_list(b"custom, another")); + assert!(!is_bare_protocol_list(b"")); + assert!(!is_bare_protocol_list(b"chat, not valid")); + + let mut table = TestSink::new(); + for (raw, view_kind, owned_kind) in [ + (", ,", DecodeErrorKind::MissingValue, DecodeErrorKind::InvalidSyntax), + ("chat, not valid", DecodeErrorKind::InvalidToken, DecodeErrorKind::InvalidToken), + ( + "chat, \"unterminated", + DecodeErrorKind::UnterminatedQuote, + DecodeErrorKind::InvalidToken, + ), + ] { + table + .set_values( + SecWebSocketProtocol::name(), + EncodedValues::single(FieldValue::from_str(raw).expect("safe raw value")), + ) + .expect("table accepts raw value"); + assert_eq!(SecWebSocketProtocol::view(&table).expect_err("invalid view").kind(), view_kind); + assert_eq!( + SecWebSocketProtocol::owned(&table).expect_err("invalid owned value").kind(), + owned_kind + ); + } + assert_eq!( + SecWebSocketProtocolOwned::try_from("chat\n") + .expect_err("invalid borrowed field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + SecWebSocketProtocolOwned::try_from(FieldValue::from_static("chat, not valid")) + .expect_err("invalid stored list") + .kind(), + DecodeErrorKind::InvalidToken + ); + for raw in ["", " \t ", ",", ", ,", " ,\t, "] { + assert_eq!( + SecWebSocketProtocolOwned::try_from(FieldValue::from_static(raw)) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + let protocol = SecWebSocketProtocolOwned::try_from(FieldValue::from_static("chat")).unwrap(); + assert_eq!(protocol.selected(), Ok("chat")); + } +} diff --git a/crates/http_headers/src/headers/websocket/sec_web_socket_version.rs b/crates/http_headers/src/headers/websocket/sec_web_socket_version.rs new file mode 100644 index 000000000..f760c1efd --- /dev/null +++ b/crates/http_headers/src/headers/websocket/sec_web_socket_version.rs @@ -0,0 +1,762 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt::Write as _; +use std::{fmt, str}; + +use super::super::invalid_syntax; +use super::shared::{trim_ows, validate_list}; +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue}; + +/// Defines the `Sec-WebSocket-Version` header. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.5](https://www.rfc-editor.org/rfc/rfc6455#section-11.3.5). +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{SecWebSocketVersion, SecWebSocketVersionOwned}; +/// +/// let mut map = HeaderMap::new(); +/// SecWebSocketVersion::insert(&mut map, SecWebSocketVersionOwned::new(13))?; +/// assert!(SecWebSocketVersion::view(&map)?.is_some()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +#[derive(Debug)] +pub struct SecWebSocketVersion { + _private: (), +} + +/// Set of WebSocket versions, one bit per version. +/// +/// [RFC 6455 section 11.3.5] bounds a version to 0-255, so the entire domain +/// fits in 256 bits. Holding the set inline keeps every field-line shape, +/// including the version list a `426` response advertises, free of allocation +/// and makes reading a version a bit test rather than a reparse. +/// +/// [RFC 6455 section 11.3.5]: https://www.rfc-editor.org/rfc/rfc6455#section-11.3.5 +#[derive(Clone, Copy, Eq, Hash, PartialEq)] +struct VersionSet { + words: [u64; 4], +} + +impl VersionSet { + const EMPTY: Self = Self { words: [0; 4] }; + + const fn one(version: u8) -> Self { + let mut set = Self::EMPTY; + set.insert(version); + set + } + + const fn insert(&mut self, version: u8) { + self.words[(version >> 6) as usize] |= 1_u64 << (version & 63); + } + + const fn contains(self, version: u8) -> bool { + self.words[(version >> 6) as usize] & (1_u64 << (version & 63)) != 0 + } + + const fn is_empty(self) -> bool { + self.words[0] | self.words[1] | self.words[2] | self.words[3] == 0 + } + + fn len(self) -> usize { + self.words.iter().map(|word| word.count_ones() as usize).sum() + } + + /// Yields the members in ascending order. + fn iter(self) -> impl Iterator { + self.words + .into_iter() + .zip([0_u8, 64, 128, 192]) + .flat_map(|(word, base)| VersionBits { word, base }) + } +} + +/// Yields the versions one bitmap word holds, lowest first. +struct VersionBits { + word: u64, + base: u8, +} + +impl Iterator for VersionBits { + type Item = u8; + + fn next(&mut self) -> Option { + (self.word != 0).then(|| { + let bit = self.word.trailing_zeros(); + self.word &= self.word - 1; + self.base + u8::try_from(bit).expect("a set bit index of a word fits in one byte") + }) + } +} + +/// Owned value for the `Sec-WebSocket-Version` header. +/// +/// The field value carries nothing but version numbers, so the set of those +/// numbers is all this type keeps: surrounding whitespace, member order, and +/// repetition are accepted when decoding and dropped, and encoding renders one +/// ascending comma-separated list. +/// +/// # Specification +/// +/// Defined by [RFC 6455 section 11.3.5]. +/// +/// # Examples +/// +/// ```rust +/// let value = http_headers::headers::SecWebSocketVersionOwned::new(13); +/// assert_eq!(value.versions().next(), Some(13)); +/// ``` +/// +/// `Sec-WebSocket-Version: 13` requests the standard version. +/// A server can advertise alternatives with +/// `Sec-WebSocket-Version: 7, 8, 13`. +/// +/// [RFC 6455 section 11.3.5]: https://www.rfc-editor.org/rfc/rfc6455#section-11.3.5 +#[derive(Clone, Copy, Eq, Hash, PartialEq)] +pub struct SecWebSocketVersionOwned { + versions: VersionSet, +} + +impl fmt::Debug for SecWebSocketVersionOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("SecWebSocketVersionOwned") + .field("versions", &self.versions.len()) + .finish_non_exhaustive() + } +} + +impl fmt::Display for SecWebSocketVersionOwned { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut versions = self.versions(); + if let Some(version) = versions.next() { + version.fmt(f)?; + } + for version in versions { + f.write_str(", ")?; + version.fmt(f)?; + } + Ok(()) + } +} + +impl SecWebSocketVersionOwned { + /// Constructs one WebSocket version. + /// + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::SecWebSocketVersionOwned::new(13); + /// assert_eq!(value.versions().next(), Some(13)); + /// ``` + #[must_use] + pub const fn new(version: u8) -> Self { + Self { + versions: VersionSet::one(version), + } + } + + /// Adds another supported version. + /// + /// Adding a version the value already holds leaves it unchanged. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketVersionOwned; + /// + /// let value = SecWebSocketVersionOwned::new(13).with_version(8); + /// assert_eq!(value.versions().collect::>(), vec![8, 13]); + /// ``` + #[must_use] + pub const fn with_version(mut self, version: u8) -> Self { + self.versions.insert(version); + self + } + + /// Reports whether the value holds a version. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketVersionOwned; + /// + /// let value = SecWebSocketVersionOwned::new(13); + /// assert!(value.supports(13)); + /// assert!(!value.supports(8)); + /// ``` + #[must_use] + pub const fn supports(&self, version: u8) -> bool { + self.versions.contains(version) + } + + /// Iterates the versions in ascending order. + /// + /// # Examples + /// + /// ```rust + /// let value = http_headers::headers::SecWebSocketVersionOwned::new(13); + /// assert_eq!(value.versions().next(), Some(13)); + /// ``` + pub fn versions(&self) -> impl Iterator + '_ { + self.versions.iter() + } + + /// Returns the single requested version. + /// + /// # Errors + /// + /// Returns an error when the value holds no version or more than one. + /// # Examples + /// + /// ```rust + /// use http_headers::headers::SecWebSocketVersionOwned; + /// + /// let value = SecWebSocketVersionOwned::new(13); + /// assert_eq!(value.requested()?, 13); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + pub fn requested(&self) -> Result { + let mut versions = self.versions(); + let version = versions + .next() + .ok_or_else(|| DecodeError::new(&FieldName::SecWebSocketVersion, DecodeErrorKind::MissingValue))?; + if versions.next().is_some() { + return Err(DecodeError::new( + &FieldName::SecWebSocketVersion, + DecodeErrorKind::UnexpectedMultipleValues, + )); + } + Ok(version) + } + + /// Renders the versions as the one field value that carries them all. + pub(crate) fn field_value(&self) -> FieldValue { + let mut versions = self.versions(); + let Some(first) = versions.next() else { + return FieldValue::from_static(""); + }; + let Some(second) = versions.next() else { + return FieldValue::from(u64::from(first)); + }; + let mut rendered = format!("{first}, {second}"); + for version in versions { + write!(rendered, ", {version}").expect("writing to a string cannot fail"); + } + FieldValue::try_from(rendered).expect("decimal versions are valid field-value bytes") + } +} + +impl Field for SecWebSocketVersion { + type View<'a> = SecWebSocketVersionOwned; + type Owned = SecWebSocketVersionOwned; + + fn name() -> &'static FieldName { + &FieldName::SecWebSocketVersion + } + + #[expect( + clippy::inline_always, + reason = "folding the one-line fast path into the caller removes a full stack frame" + )] + #[inline(always)] + fn view_with(source: &S, _mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + decode_versions(source.lines(Self::name())) + } + + #[expect( + clippy::inline_always, + reason = "folding the one-line fast path into the caller removes a full stack frame" + )] + #[inline(always)] + fn owned_with(source: &S, _mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + decode_versions(source.lines(Self::name())) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::single(value.field_value())) + } +} + +/// Reads the version set a `Sec-WebSocket-Version` header holds. +#[expect( + clippy::inline_always, + reason = "keep the version 13 comparison and constant bitmap inside typed decode callers" +)] +#[inline(always)] +fn decode_versions(values: Option>) -> Result, DecodeError> { + let Some(values) = values else { + return Ok(None); + }; + values.validate_list_item_limit(b',', true)?; + let mut lines = values.repeated(); + let lone = match (lines.next(), lines.next()) { + (Some(first), None) => Some(first.as_bytes()), + _ => None, + }; + let versions = if lone == Some(b"13".as_slice()) { + VersionSet::one(13) + } else { + collect_version_lines(values, lone)? + }; + Ok(Some(SecWebSocketVersionOwned { versions })) +} + +super::super::shared::impl_string_conversions!(SecWebSocketVersionOwned, &FieldName::SecWebSocketVersion, invalid_syntax, value); + +impl TryFrom for SecWebSocketVersionOwned { + type Error = DecodeError; + + fn try_from(value: FieldValue) -> Result { + let values = FieldLines::single(&FieldName::SecWebSocketVersion, value.as_bytes()); + Ok(decode_versions(Some(values))?.expect("FieldLines::single always supplies one field value")) + } +} + +/// Decodes uncommon bare versions and advertised version lists. +#[expect( + clippy::needless_pass_by_value, + reason = "moving fallback storage keeps its stack materialization off the version 13 fast path" +)] +#[cold] +#[inline(never)] +fn collect_version_lines(values: FieldLines<'_>, lone: Option<&[u8]>) -> Result { + if let Some(version) = lone.and_then(parse_version_value) { + return Ok(VersionSet::one(version)); + } + let mut versions = VersionSet::EMPTY; + for value in values.repeated() { + if !collect_bare_version_line(value.as_bytes(), &mut versions)? { + return collect_quoted_version_lines(&values); + } + } + require_version(!versions.is_empty())?; + Ok(versions) +} + +/// Collects lines a quoted member forced the delimiter-aware walk to reread. +/// +/// Quoting never produces a valid version — the grammar admits only digits — +/// but it does move where members start and end, so the shared walk decides +/// the split and reports an unterminated quote. +#[cold] +#[inline(never)] +fn collect_quoted_version_lines(values: &FieldLines<'_>) -> Result { + let mut versions = VersionSet::EMPTY; + let mut collect = |bytes: &[u8]| -> Result<(), DecodeError> { + versions.insert(parse_version(bytes)?); + Ok(()) + }; + validate_list(values, &mut collect)?; + Ok(versions) +} + +/// Collects one comma-separated line of bare versions. +/// +/// Reports `false` when the line quotes a member, which only the general +/// walk can split correctly. +#[inline] +fn collect_bare_version_line(bytes: &[u8], versions: &mut VersionSet) -> Result { + if bytes.iter().any(|byte| matches!(byte, b'"' | b'\\')) { + return Ok(false); + } + for item in bytes.split(|byte| *byte == b',').map(trim_ows) { + if !item.is_empty() { + versions.insert(parse_version(item)?); + } + } + Ok(true) +} + +#[inline] +fn require_version(present: bool) -> Result<(), DecodeError> { + present + .then_some(()) + .ok_or_else(|| DecodeError::new(&FieldName::SecWebSocketVersion, DecodeErrorKind::MissingValue)) +} + +fn parse_version(bytes: &[u8]) -> Result { + parse_version_value(bytes).ok_or_else(|| invalid_syntax(&FieldName::SecWebSocketVersion)) +} + +fn parse_version_value(bytes: &[u8]) -> Option { + // One and two digits cover every version the protocol has ever defined, + // including the draft numbers a `426` response advertises, so both stay + // inline; only the three-digit form the grammar still permits is remote + // enough to pay for a call. + match *bytes { + [only] if only.is_ascii_digit() => Some(only - b'0'), + [first, second] if is_leading_digit(first) && second.is_ascii_digit() => Some((first - b'0') * 10 + (second - b'0')), + [_, _, _] => parse_uncommon_version_value(bytes), + _ => None, + } +} + +/// Parses the three-digit versions the wire rarely carries. +#[cold] +#[inline(never)] +fn parse_uncommon_version_value(bytes: &[u8]) -> Option { + let version = match *bytes { + [first, second, third] if is_leading_digit(first) && second.is_ascii_digit() && third.is_ascii_digit() => { + u16::from(first - b'0') * 100 + u16::from(second - b'0') * 10 + u16::from(third - b'0') + } + _ => return None, + }; + u8::try_from(version).ok() +} + +const fn is_leading_digit(byte: u8) -> bool { + byte.wrapping_sub(b'1') < 9 +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{ + SecWebSocketVersion, SecWebSocketVersionOwned, VersionSet, collect_quoted_version_lines, decode_versions, parse_version, + parse_version_value, + }; + use crate::sink::{EncodedValues, FieldSink}; + use crate::source::{FieldLines, FieldSource}; + use crate::{DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef, TestSink}; + + struct Source<'a>(&'a [u8]); + + impl FieldSource for Source<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::SecWebSocketVersion).then(|| FieldLines::single(name, self.0)) + } + } + + #[test] + fn a_version_set_round_trips_every_representable_version() { + for version in 0..=u8::MAX { + let set = VersionSet::one(version); + assert!(set.contains(version), "{version} is absent after insertion"); + assert!(!set.is_empty()); + assert_eq!(set.len(), 1); + assert_eq!(set.iter().collect::>(), vec![version]); + } + + let mut every = VersionSet::EMPTY; + assert!(every.is_empty()); + for version in 0..=u8::MAX { + every.insert(version); + } + assert_eq!(every.len(), 256); + assert_eq!(every.iter().collect::>(), (0..=u8::MAX).collect::>()); + } + + #[test] + fn the_canonical_version_decodes_without_allocating() { + let owned = SecWebSocketVersionOwned::new(13); + assert_eq!(owned.requested(), Ok(13)); + assert!(owned.supports(13)); + assert!(!owned.supports(8)); + assert_eq!(owned.versions().collect::>(), vec![13]); + assert_eq!(format!("{owned:?}"), "SecWebSocketVersionOwned { versions: 1, .. }"); + + let source = Source(b"13"); + let view = SecWebSocketVersion::view(&source) + .expect("version is valid") + .expect("version is present"); + assert_eq!(view.requested(), Ok(13)); + let owned = SecWebSocketVersion::owned(&source) + .expect("version is valid") + .expect("version is present"); + assert_eq!(owned.requested(), Ok(13)); + assert_eq!(view, owned); + } + + #[test] + fn an_advertised_version_list_collapses_to_an_ascending_set() { + let owned = SecWebSocketVersionOwned::new(8); + assert_eq!(owned.requested(), Ok(8)); + + let source = Source(b"7, 8, 13"); + let view = SecWebSocketVersion::view(&source) + .expect("versions are valid") + .expect("versions are present"); + assert_eq!(view.versions().collect::>(), vec![7, 8, 13]); + + // Order and repetition carry no meaning, so both are dropped. + let shuffled = Source(b"13, 7, 8, 13"); + assert_eq!( + SecWebSocketVersion::view(&shuffled) + .expect("versions are valid") + .expect("versions are present"), + view + ); + + let source = Source(b"8"); + let view = SecWebSocketVersion::view(&source) + .expect("uncommon version is valid") + .expect("version is present"); + assert_eq!(view.requested(), Ok(8)); + } + + #[test] + fn the_version_parser_matches_a_reference_implementation() { + /// Parses the `version` rule of RFC 6455 section 11.3.5 directly. + fn reference(bytes: &[u8]) -> Option { + let text = str::from_utf8(bytes).ok()?; + if text.len() > 1 && text.starts_with('0') { + return None; + } + if !text.bytes().all(|byte| byte.is_ascii_digit()) { + return None; + } + text.parse::().ok() + } + + for first in 0..=u8::MAX { + assert_eq!( + parse_version_value(&[first]), + reference(&[first]), + "one-byte disagreement for {first:#04x}" + ); + #[cfg(miri)] + if !crate::test_support::is_byte_case(first, b'0') { + continue; + } + for second in crate::test_support::byte_cases(b'0') { + let two = [first, second]; + assert_eq!(parse_version_value(&two), reference(&two), "two-byte disagreement for {two:?}"); + } + } + for value in 0_u32..1000 { + let three = format!("{value:03}"); + assert_eq!( + parse_version_value(three.as_bytes()), + reference(three.as_bytes()), + "three-byte disagreement for {three}" + ); + } + assert_eq!(parse_version_value(b""), None); + assert_eq!(parse_version_value(b"1234"), None); + } + + #[test] + fn version_line_substitutions_match_structured_parsing() { + for literal in [ + b"13".as_slice(), + b"7, 8, 13", + b" \t13\t ", + b" , , \t", + b"\"13", + b"\"13\"", + b"256, \"13", + b"13, \"13", + b"13\\13", + b"13, 256", + ] { + let mut bytes = literal.to_vec(); + for index in 0..bytes.len() { + for replacement in crate::test_support::substitution_bytes(literal[index], index, literal.len()) { + bytes[index] = replacement; + for prefix in [None, Some(b"13".as_slice()), Some(b"256".as_slice())] { + let repeated = [FieldValueRef::new(prefix.unwrap_or_default()), FieldValueRef::new(&bytes)]; + let lines = if prefix.is_some() { + FieldLines::from_borrowed(&FieldName::SecWebSocketVersion, &repeated).unwrap() + } else { + FieldLines::single(&FieldName::SecWebSocketVersion, &bytes) + }; + let expected = collect_quoted_version_lines(&lines).map(|versions| Some(versions.words)); + let actual = decode_versions(Some(lines)).map(|value| value.map(|value| value.versions.words)); + assert_eq!( + actual, expected, + "{literal:?}, index {index}, replacement {replacement}, prefix {prefix:?}" + ); + } + } + bytes[index] = literal[index]; + } + } + } + + #[test] + fn builders_conversions_and_insert_cover_canonical_and_general_storage() { + let versions = SecWebSocketVersionOwned::new(13).with_version(8).with_version(255); + assert_eq!(versions.versions().collect::>(), vec![8, 13, 255]); + assert_eq!(versions.with_version(8), versions); + assert_eq!( + versions.requested().expect_err("multiple versions").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!(format!("{versions:?}"), "SecWebSocketVersionOwned { versions: 3, .. }"); + + let uncommon = SecWebSocketVersionOwned::new(7).with_version(13); + assert_eq!(uncommon.versions().collect::>(), vec![7, 13]); + + assert_eq!( + SecWebSocketVersionOwned::try_from("13").expect("canonical conversion").requested(), + Ok(13) + ); + for value in [ + SecWebSocketVersionOwned::try_from("7, 8, 13"), + SecWebSocketVersionOwned::try_from(String::from("7, 8, 13")), + SecWebSocketVersionOwned::try_from(FieldValue::from_static("7, 8, 13")), + ] { + assert_eq!(value.expect("valid version list").versions().collect::>(), vec![7, 8, 13]); + } + assert_eq!( + SecWebSocketVersionOwned::try_from(String::from("13\n")) + .expect_err("invalid field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + SecWebSocketVersionOwned::try_from("13\n") + .expect_err("invalid borrowed field string") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + SecWebSocketVersionOwned::try_from(FieldValue::from_static("256")) + .expect_err("out-of-range stored version") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let mut table = TestSink::new(); + SecWebSocketVersion::insert(&mut table, SecWebSocketVersionOwned::new(13)).expect("canonical version inserts"); + assert_eq!( + table + .lines(SecWebSocketVersion::name()) + .expect("inserted version") + .exactly_one() + .expect("one line") + .as_bytes(), + b"13" + ); + // The whole set renders as one ascending list rather than one line + // per version. + SecWebSocketVersion::insert(&mut table, versions).expect("general versions insert"); + assert_eq!( + table + .lines(SecWebSocketVersion::name()) + .expect("inserted versions") + .exactly_one() + .expect("one line") + .as_bytes(), + b"8, 13, 255" + ); + } + + #[test] + fn repeated_lines_decode_owned_and_borrowed_and_invalid_lists_fail() { + let mut table = TestSink::new(); + assert!(SecWebSocketVersion::view(&table).expect("absent view succeeds").is_none()); + assert!(SecWebSocketVersion::owned(&table).expect("absent owned succeeds").is_none()); + table + .set_values( + SecWebSocketVersion::name(), + EncodedValues::from_vec(vec![FieldValue::from_static("7, 8"), FieldValue::from_static("13")]), + ) + .expect("table accepts versions"); + let view = SecWebSocketVersion::view(&table).expect("view decodes").expect("header is present"); + assert_eq!(view.versions().collect::>(), vec![7, 8, 13]); + assert_eq!( + SecWebSocketVersion::owned(&table) + .expect("owned decode succeeds") + .expect("header is present") + .versions() + .collect::>(), + vec![7, 8, 13] + ); + + for (raw, kind) in [ + (", ,", DecodeErrorKind::MissingValue), + ("256", DecodeErrorKind::InvalidSyntax), + ("01", DecodeErrorKind::InvalidSyntax), + ("\"13", DecodeErrorKind::UnterminatedQuote), + ] { + table + .set_values( + SecWebSocketVersion::name(), + EncodedValues::single(FieldValue::from_str(raw).expect("safe raw value")), + ) + .expect("table accepts raw value"); + assert_eq!(SecWebSocketVersion::view(&table).expect_err("invalid version list").kind(), kind); + assert_eq!( + SecWebSocketVersion::owned(&table).expect_err("invalid owned version list").kind(), + kind + ); + } + + table + .set_values(SecWebSocketVersion::name(), EncodedValues::single(FieldValue::from_static("8"))) + .expect("table accepts uncommon version"); + assert_eq!( + SecWebSocketVersion::owned(&table) + .expect("owned uncommon version decodes") + .expect("header is present") + .requested(), + Ok(8) + ); + } + + #[test] + fn private_version_helpers_report_empty_multiple_late_errors_and_bounds() { + let empty = SecWebSocketVersionOwned { + versions: VersionSet::EMPTY, + }; + assert_eq!(empty.requested().expect_err("empty set").kind(), DecodeErrorKind::MissingValue); + assert_eq!(empty.versions().next(), None); + assert_eq!(empty.field_value().as_bytes(), b""); + assert_eq!(empty.to_string(), ""); + assert_eq!( + SecWebSocketVersionOwned::new(13) + .with_version(8) + .requested() + .expect_err("two versions are not one request version") + .kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!(parse_version(b"255"), Ok(255)); + assert_eq!( + parse_version(b"256").expect_err("out-of-range version").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!(parse_version_value(b"0"), Some(0)); + assert_eq!(parse_version_value(b"9"), Some(9)); + assert_eq!(parse_version_value(b"99"), Some(99)); + assert_eq!(parse_version_value(b"100"), Some(100)); + assert_eq!(parse_version_value(b"255"), Some(255)); + assert_eq!(parse_version_value(b"256"), None); + assert_eq!(parse_version_value(b"01"), None); + assert_eq!(parse_version_value(b"000"), None); + assert_eq!(parse_version_value(b"1000"), None); + + let bare = FieldLines::single(&FieldName::SecWebSocketVersion, b"8, 13"); + assert_eq!( + super::collect_quoted_version_lines(&bare) + .expect("bare versions are accepted by the general collector") + .iter() + .collect::>(), + vec![8, 13] + ); + } +} diff --git a/crates/http_headers/src/headers/websocket/shared.rs b/crates/http_headers/src/headers/websocket/shared.rs new file mode 100644 index 000000000..f2e1a0012 --- /dev/null +++ b/crates/http_headers/src/headers/websocket/shared.rs @@ -0,0 +1,672 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use base64::Engine as _; +use base64::engine::general_purpose::STANDARD; + +use super::super::{invalid_syntax, trim_ows as header_trim_ows}; +use crate::source::FieldLines; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef}; + +pub(super) fn encode_fixed_base64(input: &[u8]) -> FieldValue { + FieldValue::try_from(STANDARD.encode(input)).expect("the standard base64 alphabet is always a valid field value") +} + +fn base64_encoded_length(decoded_length: usize) -> Option { + decoded_length.checked_add(2).and_then(|length| (length / 3).checked_mul(4)) +} + +pub(super) fn increment_item_count(value: usize, name: &'static FieldName) -> Result { + value + .checked_add(1) + .ok_or_else(|| DecodeError::new(name, DecodeErrorKind::InvalidNumber)) +} + +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the constant lengths into the scan" +)] +#[inline(always)] +pub(super) fn validate_canonical_base64(bytes: &[u8], decoded_length: usize, name: &'static FieldName) -> Result<(), DecodeError> { + let encoded_length = base64_encoded_length(decoded_length).ok_or_else(|| DecodeError::new(name, DecodeErrorKind::InvalidNumber))?; + let padding = (3 - decoded_length % 3) % 3; + if bytes.len() != encoded_length { + return Err(invalid_syntax(name)); + } + let data_end = encoded_length - padding; + let (data, pad) = bytes.split_at(data_end); + if !http_headers_simd::all_base64_alphabet(data) || !all_padding(pad) { + return Err(invalid_syntax(name)); + } + let canonical_tail = if padding == 0 { + true + } else if padding == 1 { + data.last().is_some_and(|byte| is_canonical_tail_pad1(*byte)) + } else { + data.last().is_some_and(|byte| is_canonical_tail_pad2(*byte)) + }; + if canonical_tail { Ok(()) } else { Err(invalid_syntax(name)) } +} + +/// Lets the nonce's two padding bytes use one fixed-width comparison. +#[expect( + clippy::inline_always, + reason = "specializing per call site folds the constant lengths into the scan" +)] +#[inline(always)] +fn all_padding(bytes: &[u8]) -> bool { + if bytes.len() == 2 { + bytes == b"==" + } else { + bytes.iter().fold(0_u8, |differs, byte| differs | (byte ^ b'=')) == 0 + } +} + +/// Classifies every byte as a canonical final data character. +/// +/// Bit zero is set when the byte may end a run padded with a single `=`, and +/// bit one when it may end a run padded with two. The table is derived from +/// [`base64_sextet`], so a probe is exactly the arithmetic test it replaces +/// while costing one load rather than a shift-and-mask chain. +static CANONICAL_TAIL: [u8; 256] = canonical_tail_table(); + +const fn canonical_tail_table() -> [u8; 256] { + let mut table = [0_u8; 256]; + let mut byte = 0_u8; + loop { + if let Some(sextet) = base64_sextet(byte) { + if sextet.trailing_zeros() >= 2 { + table[byte as usize] |= 1; + } + if sextet.trailing_zeros() >= 4 { + table[byte as usize] |= 2; + } + } + if byte == u8::MAX { + return table; + } + byte += 1; + } +} + +/// Accepts the last data character of a run padded with a single `=`. +/// +/// One pad character means the final sextet carries only two significant +/// bits beyond the byte boundary, so the encoding is canonical exactly when +/// the low two bits of the sextet are clear. +#[inline] +fn is_canonical_tail_pad1(byte: u8) -> bool { + CANONICAL_TAIL[usize::from(byte)] & 1 != 0 +} + +/// Accepts the last data character of a run padded with two `=`. +/// +/// Two pad characters leave only two significant bits in the final sextet, so +/// the encoding is canonical exactly when the low four bits are clear. +#[inline] +fn is_canonical_tail_pad2(byte: u8) -> bool { + CANONICAL_TAIL[usize::from(byte)] & 2 != 0 +} + +const fn base64_sextet(byte: u8) -> Option { + match byte { + b'A'..=b'Z' => Some(byte - b'A'), + b'a'..=b'z' => Some(byte - b'a' + 26), + b'0'..=b'9' => Some(byte - b'0' + 52), + b'+' => Some(62), + b'/' => Some(63), + _ => None, + } +} + +pub(super) fn validate_bare_list( + values: &FieldLines<'_>, + accepts_bare_item: fn(&[u8]) -> bool, + validate_item: fn(&[u8]) -> Result<(), DecodeError>, +) -> Result<(), DecodeError> { + values.validate_list_item_limit(b',', true)?; + for value in values.repeated() { + if !accepts_bare_item(value.as_bytes()) { + let mut validate_item = validate_item; + return validate_list(values, &mut validate_item); + } + } + Ok(()) +} + +/// Validates every comma-delimited member of a header's field lines. +/// +/// Quoting only ever affects where members start and end, so members holding +/// no quote or backslash are split by a tight scan that needs no per-item +/// iterator state. On the first quoted member, the delimiter-aware scan +/// resumes at that member and continues through the remaining field lines. +#[inline(never)] +pub(super) fn validate_list( + values: &FieldLines<'_>, + validate_item: &mut dyn FnMut(&[u8]) -> Result<(), DecodeError>, +) -> Result<(), DecodeError> { + values.validate_list_item_limit(b',', true)?; + let mut present = false; + let mut lines = values.repeated().enumerate(); + while let Some((value_index, value)) = lines.next() { + let mut item_start = 0_usize; + for raw_item in value.as_bytes().split(|byte| *byte == b',') { + if raw_item.iter().any(|byte| matches!(byte, b'"' | b'\\')) { + validate_quoted_line( + values.name(), + &value.as_bytes()[item_start..], + value_index, + validate_item, + &mut present, + )?; + for (value_index, value) in lines { + validate_quoted_line(values.name(), value.as_bytes(), value_index, validate_item, &mut present)?; + } + return Ok(()); + } + let item = trim_ows(raw_item); + if !item.is_empty() { + validate_item(item)?; + present = true; + } + item_start += raw_item.len() + 1; + } + } + if present { + Ok(()) + } else { + Err(DecodeError::new(values.name(), DecodeErrorKind::MissingValue)) + } +} + +#[inline(never)] +fn validate_quoted_line( + name: &'static FieldName, + bytes: &[u8], + value_index: usize, + validate_item: &mut dyn FnMut(&[u8]) -> Result<(), DecodeError>, + present: &mut bool, +) -> Result<(), DecodeError> { + let mut start = 0_usize; + let mut position = 0_usize; + let mut quoted = false; + let mut escaped = false; + while position < bytes.len() { + let byte = bytes[position]; + if escaped { + escaped = false; + } else if quoted && byte == b'\\' { + escaped = true; + } else if byte == b'"' { + quoted = !quoted; + } else if !quoted && byte == b',' { + let item = trim_ows(&bytes[start..position]); + *present |= validate_present_item(item, validate_item)?; + start = position + 1; + } + position += 1; + } + if quoted || escaped { + return Err(DecodeError::new(name, DecodeErrorKind::UnterminatedQuote).at_value(value_index)); + } + let item = trim_ows(&bytes[start..]); + *present |= validate_present_item(item, validate_item)?; + Ok(()) +} + +fn validate_present_item(item: &[u8], validate_item: &mut dyn FnMut(&[u8]) -> Result<(), DecodeError>) -> Result { + if item.is_empty() { + return Ok(false); + } + validate_item(item)?; + Ok(true) +} + +/// Removes optional whitespace without the bounds checks of range slicing. +#[inline] +pub(super) fn trim_ows(bytes: &[u8]) -> &[u8] { + let mut trimmed = bytes; + while let [b' ' | b'\t', rest @ ..] = trimmed { + trimmed = rest; + } + while let [rest @ .., b' ' | b'\t'] = trimmed { + trimmed = rest; + } + trimmed +} + +/// Validates a physical line and reports whether it contributes any members. +pub(super) fn validate_header_value_list( + value: FieldValueRef<'_>, + name: &'static FieldName, + validate_item: fn(&[u8]) -> Result<(), DecodeError>, +) -> Result { + let mut count = 0_usize; + for item in CommaItems::new(value.as_bytes()) { + validate_item(item)?; + count = increment_item_count(count, name)?; + } + Ok(count != 0) +} + +pub(super) struct CommaItems<'a> { + bytes: &'a [u8], + start: usize, + position: usize, + finished: bool, +} + +impl<'a> CommaItems<'a> { + pub(super) const fn new(bytes: &'a [u8]) -> Self { + Self { + bytes, + start: 0, + position: 0, + finished: false, + } + } +} + +impl<'a> Iterator for CommaItems<'a> { + type Item = &'a [u8]; + + fn next(&mut self) -> Option { + while !self.finished { + let mut quoted = false; + let mut escaped = false; + while let Some(byte) = self.bytes.get(self.position).copied() { + if escaped { + escaped = false; + } else if quoted && byte == b'\\' { + escaped = true; + } else if byte == b'"' { + quoted = !quoted; + } else if !quoted && byte == b',' { + let item = header_trim_ows(&self.bytes[self.start..self.position]); + self.position += 1; + self.start = self.position; + if item.is_empty() { + continue; + } + return Some(item); + } + self.position += 1; + } + self.finished = true; + let item = header_trim_ows(&self.bytes[self.start..]); + if !item.is_empty() { + return Some(item); + } + } + None + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use base64::Engine as _; + use base64::engine::general_purpose::STANDARD; + + use super::{ + CommaItems, base64_encoded_length, base64_sextet, canonical_tail_table, encode_fixed_base64, increment_item_count, + is_canonical_tail_pad1, is_canonical_tail_pad2, trim_ows, validate_bare_list, validate_canonical_base64, + validate_header_value_list, validate_list, validate_present_item, + }; + use crate::headers::{SecWebSocketExtensions, SecWebSocketProtocol}; + use crate::source::{FieldLines, FieldSource, MAX_CUSTOM_FIELD_LINES, MAX_CUSTOM_LIST_ITEMS}; + use crate::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef, TestSink}; + + struct RawSource<'a>(&'a [FieldValueRef<'a>]); + + impl FieldSource for RawSource<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_borrowed(name, self.0) + } + } + + fn assert_preserved(source: &impl FieldSource, mode: DecodeMode, expected: &[FieldValueRef<'_>]) { + assert!(F::view_with(source, mode).unwrap().is_some()); + let owned = F::owned_with(source, mode).unwrap().unwrap(); + let mut output = TestSink::new(); + F::insert(&mut output, owned).unwrap(); + let lines = output.lines(F::name()).unwrap(); + assert_eq!(lines.repeated().collect::>(), expected); + assert_eq!( + lines.repeated().map(FieldValueRef::is_sensitive).collect::>(), + expected.iter().map(|value| value.is_sensitive()).collect::>() + ); + } + + fn assert_rejected(source: &impl FieldSource, mode: DecodeMode) { + F::view_with(source, mode).map(|_| ()).unwrap_err(); + F::owned_with(source, mode).map(|_| ()).unwrap_err(); + } + + #[test] + fn empty_physical_lines_preserve_nonempty_logical_lists() { + for empty in [b"".as_slice(), b" \t ", b",, \t,"] { + for wires in [ + vec![empty, b"chat", b"superchat"], + vec![b"chat".as_slice(), empty, b"superchat"], + vec![b"chat".as_slice(), b"superchat", empty], + vec![empty, b"chat", empty, b"superchat", empty], + ] { + let values: Vec<_> = wires + .iter() + .enumerate() + .map(|(index, bytes)| FieldValueRef::new(bytes).with_sensitive(index % 2 == 0)) + .collect(); + let raw = RawSource(&values); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_preserved::(&raw, mode, &values); + assert_preserved::(&raw, mode, &values); + let protocols = SecWebSocketProtocol::view_with(&raw, mode).unwrap().unwrap(); + assert_eq!(protocols.protocols().collect::, _>>().unwrap(), ["chat", "superchat"]); + let owned = SecWebSocketProtocol::owned_with(&raw, mode).unwrap().unwrap(); + assert_eq!(owned.protocols().collect::, _>>().unwrap(), ["chat", "superchat"]); + let extensions = SecWebSocketExtensions::view_with(&raw, mode).unwrap().unwrap(); + assert_eq!( + extensions.extensions().map(|item| item.unwrap().name()).collect::>(), + ["chat", "superchat"] + ); + let owned = SecWebSocketExtensions::owned_with(&raw, mode).unwrap().unwrap(); + assert_eq!( + owned.extensions().map(|item| item.unwrap().name()).collect::>(), + ["chat", "superchat"] + ); + + #[cfg(feature = "http")] + { + let mut map = http::HeaderMap::new(); + for name in [SecWebSocketProtocol::name(), SecWebSocketExtensions::name()] { + for value in &values { + map.append(name.as_str(), http::HeaderValue::try_from(*value).unwrap()); + } + } + assert_preserved::(&map, mode, &values); + assert_preserved::(&map, mode, &values); + } + } + } + } + } + + #[test] + fn empty_lines_do_not_hide_missing_or_malformed_members_or_source_limits() { + let empty = FieldValueRef::new(b", \t,"); + let valid = FieldValueRef::new(b"chat"); + let too_many_items = vec!["chat"; MAX_CUSTOM_LIST_ITEMS + 1].join(","); + let mut too_many_lines = vec![empty; MAX_CUSTOM_FIELD_LINES]; + too_many_lines.push(valid); + for values in [ + vec![FieldValueRef::new(b""), FieldValueRef::new(b" \t"), empty], + vec![empty, valid, FieldValueRef::new(b"not valid"), empty], + vec![empty, valid, FieldValueRef::new(b"bad\r\n"), empty], + vec![empty, valid, FieldValueRef::new(b"x; p=\"unterminated"), empty], + vec![empty, FieldValueRef::new(too_many_items.as_bytes()), empty], + too_many_lines, + ] { + let raw = RawSource(&values); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_rejected::(&raw, mode); + assert_rejected::(&raw, mode); + } + } + + #[cfg(feature = "http")] + for wires in [["", " \t", ",, "], ["", "chat", "not valid"], ["", "chat", "x; p=\"unterminated"]] { + let mut map = http::HeaderMap::new(); + for name in [SecWebSocketProtocol::name(), SecWebSocketExtensions::name()] { + for wire in wires { + map.append(name.as_str(), http::HeaderValue::from_str(wire).unwrap()); + } + } + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_rejected::(&map, mode); + assert_rejected::(&map, mode); + } + } + } + + fn validate_alpha(item: &[u8]) -> Result<(), crate::DecodeError> { + item.iter() + .all(u8::is_ascii_alphabetic) + .then_some(()) + .ok_or_else(|| crate::DecodeError::new(&FieldName::SecWebSocketProtocol, DecodeErrorKind::InvalidToken)) + } + + fn reject_named_item(item: &[u8]) -> Result<(), crate::DecodeError> { + if item == b"reject" { + Err(crate::DecodeError::new( + &FieldName::SecWebSocketExtensions, + DecodeErrorKind::InvalidToken, + )) + } else { + Ok(()) + } + } + + #[test] + fn canonical_tail_table_matches_the_sextet_arithmetic() { + let build: fn() -> [u8; 256] = canonical_tail_table; + let runtime_table = std::hint::black_box(build)(); + for byte in 0..=u8::MAX { + let sextet = base64_sextet(byte); + assert_eq!(runtime_table[usize::from(byte)], super::CANONICAL_TAIL[usize::from(byte)]); + assert_eq!( + is_canonical_tail_pad1(byte), + sextet.is_some_and(|value| value.trailing_zeros() >= 2), + "one-pad classification differs for {byte:#04x}" + ); + assert_eq!( + is_canonical_tail_pad2(byte), + sextet.is_some_and(|value| value.trailing_zeros() >= 4), + "two-pad classification differs for {byte:#04x}" + ); + } + } + + #[test] + fn fixed_base64_substitutions_match_standard_decoder() { + for decoded_length in [0, 1, 2, 3, 16, 20] { + let encoded = STANDARD.encode(vec![0x69; decoded_length]); + let mut bytes = encoded.into_bytes(); + for index in 0..bytes.len() { + let original = bytes[index]; + for byte in crate::test_support::substitution_bytes(original, index, bytes.len()) { + bytes[index] = byte; + let expected = STANDARD.decode(&bytes).is_ok_and(|decoded| decoded.len() == decoded_length); + assert_eq!( + validate_canonical_base64(&bytes, decoded_length, &FieldName::SecWebSocketKey).is_ok(), + expected, + "decoded length {decoded_length}, byte {index} replaced by {byte:#04x}" + ); + } + bytes[index] = original; + } + assert_eq!( + validate_canonical_base64(&bytes, decoded_length, &FieldName::SecWebSocketKey), + Ok(()) + ); + } + } + + #[test] + fn padding_comparison_matches_each_byte() { + for length in 0..=4 { + let mut bytes = vec![b'='; length]; + assert!(super::all_padding(&bytes)); + for index in 0..length { + for byte in 0..=u8::MAX { + bytes[index] = byte; + assert_eq!(super::all_padding(&bytes), byte == b'='); + } + bytes[index] = b'='; + } + } + } + + #[test] + fn fixed_base64_encoding_and_validation_cover_all_padding_shapes() { + assert_eq!(encode_fixed_base64(b"a").as_bytes(), b"YQ=="); + assert_eq!(encode_fixed_base64(b"ab").as_bytes(), b"YWI="); + assert_eq!(encode_fixed_base64(b"abc").as_bytes(), b"YWJj"); + assert_eq!(base64_encoded_length(0), Some(0)); + assert_eq!(base64_encoded_length(3), Some(4)); + assert_eq!(base64_encoded_length(usize::MAX - 2), None); + assert_eq!(base64_encoded_length(usize::MAX), None); + assert_eq!( + increment_item_count(usize::MAX - 1, &FieldName::SecWebSocketProtocol).expect("last count is representable"), + usize::MAX + ); + assert_eq!( + increment_item_count(usize::MAX, &FieldName::SecWebSocketProtocol) + .expect_err("count overflow is rejected") + .kind(), + DecodeErrorKind::InvalidNumber + ); + + assert_eq!(validate_canonical_base64(b"YWJj", 3, &FieldName::SecWebSocketKey), Ok(())); + + let first_slow = [FieldValue::from_static("alpha, beta")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketProtocol, &first_slow).expect("nonempty values"); + assert_eq!( + validate_bare_list(&values, |item| !item.contains(&b','), |_| Ok::<_, crate::DecodeError>(()),), + Ok(()) + ); + + let later_slow = [FieldValue::from_static("alpha"), FieldValue::from_static("beta, gamma")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketProtocol, &later_slow).expect("nonempty values"); + assert_eq!( + validate_bare_list(&values, |item| !item.contains(&b','), |_| Ok::<_, crate::DecodeError>(()),), + Ok(()) + ); + assert_eq!(validate_canonical_base64(b"YQ==", 1, &FieldName::SecWebSocketKey), Ok(())); + assert_eq!(validate_canonical_base64(b"YWI=", 2, &FieldName::SecWebSocketKey), Ok(())); + for (wire, decoded) in [ + (b"YR==".as_slice(), 1), + (b"YWJ=".as_slice(), 2), + (b"YW!j".as_slice(), 3), + (b"YQ=A".as_slice(), 1), + (b"short".as_slice(), 3), + ] { + let _error = validate_canonical_base64(wire, decoded, &FieldName::SecWebSocketKey).expect_err("noncanonical base64 must fail"); + } + assert_eq!( + validate_canonical_base64(b"", usize::MAX, &FieldName::SecWebSocketKey) + .expect_err("length arithmetic overflows") + .kind(), + DecodeErrorKind::InvalidNumber + ); + } + + #[test] + fn list_validation_switches_between_bare_and_quoted_paths() { + assert_eq!(validate_present_item(b"", &mut validate_alpha), Ok(false)); + + let stored = [FieldValue::from_static("alpha"), FieldValue::from_static("beta, gamma")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketProtocol, &stored).expect("nonempty values"); + assert_eq!( + validate_bare_list(&values, |item| item.iter().all(u8::is_ascii_alphabetic), validate_alpha,), + Ok(()) + ); + + let invalid = [FieldValue::from_static("alpha, bad value")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketProtocol, &invalid).expect("nonempty values"); + assert_eq!( + validate_bare_list(&values, |item| item.iter().all(u8::is_ascii_alphabetic), validate_alpha,) + .expect_err("invalid member propagates") + .kind(), + DecodeErrorKind::InvalidToken + ); + + let quoted = [FieldValue::from_static("alpha, ext=\"a,b\""), FieldValue::from_static("z")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketExtensions, "ed).expect("nonempty values"); + let mut seen = Vec::new(); + let mut collect_item = |item: &[u8]| { + seen.push(item.to_vec()); + Ok(()) + }; + validate_list(&values, &mut collect_item).expect("quoted comma remains in one item"); + assert_eq!(seen, [b"alpha".to_vec(), b"ext=\"a,b\"".to_vec(), b"z".to_vec()]); + + let unterminated = [FieldValue::from_static("alpha"), FieldValue::from_static("\"beta")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketExtensions, &unterminated).expect("nonempty values"); + let error = validate_list(&values, &mut |_| Ok::<_, crate::DecodeError>(())).expect_err("unterminated quote"); + assert_eq!(error.kind(), DecodeErrorKind::UnterminatedQuote); + assert_eq!(error.value_index(), Some(1)); + + let quoted_then_unterminated = [FieldValue::from_static("ext=\"a,b\""), FieldValue::from_static("\"unterminated")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketExtensions, "ed_then_unterminated).expect("nonempty values"); + let error = validate_list(&values, &mut |_| Ok::<_, crate::DecodeError>(())).expect_err("later unterminated quote"); + assert_eq!(error.kind(), DecodeErrorKind::UnterminatedQuote); + assert_eq!(error.value_index(), Some(1)); + + let quoted = [FieldValue::from_static("reject, ext=\"a,b\"")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketExtensions, "ed).expect("nonempty values"); + assert_eq!( + validate_list(&values, &mut reject_named_item) + .expect_err("delimiter item error propagates") + .kind(), + DecodeErrorKind::InvalidToken + ); + + let quoted = [FieldValue::from_static("ext=\"a,b\", tail")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketExtensions, "ed).expect("nonempty values"); + assert_eq!( + validate_list(&values, &mut |_| { + Err::<(), _>(crate::DecodeError::new( + &FieldName::SecWebSocketExtensions, + DecodeErrorKind::InvalidToken, + )) + }) + .expect_err("quoted delimiter item error propagates") + .kind(), + DecodeErrorKind::InvalidToken + ); + + let quoted = [FieldValue::from_static("ext=\"a,b\", reject")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketExtensions, "ed).expect("nonempty values"); + assert_eq!( + validate_list(&values, &mut reject_named_item) + .expect_err("final item error propagates") + .kind(), + DecodeErrorKind::InvalidToken + ); + } + + #[test] + fn comma_items_header_value_lists_and_ows_handle_empty_and_quoted_members() { + assert_eq!(trim_ows(b" \t value \t "), b"value"); + assert_eq!( + CommaItems::new(b", alpha, ext=\"a,b\",, omega ").collect::>(), + [b"alpha".as_slice(), b"ext=\"a,b\"".as_slice(), b"omega".as_slice()] + ); + + let escaped = [FieldValue::from_static("ext=\"a\\,b\", tail")]; + let values = FieldLines::from_slice(&FieldName::SecWebSocketExtensions, &escaped).expect("nonempty values"); + let mut seen = Vec::new(); + let mut collect_item = |item: &[u8]| { + seen.push(item.to_vec()); + Ok(()) + }; + validate_list(&values, &mut collect_item).expect("escaped comma remains quoted"); + assert_eq!(seen, [b"ext=\"a\\,b\"".to_vec(), b"tail".to_vec()]); + + let value = FieldValue::from_static("alpha, beta"); + assert_eq!( + CommaItems::new(b"ext=\"a\\,b\", tail").collect::>(), + [b"ext=\"a\\,b\"".as_slice(), b"tail".as_slice()] + ); + assert_eq!( + validate_header_value_list(value.as_field_value_ref(), &FieldName::SecWebSocketProtocol, validate_alpha,), + Ok(true) + ); + let empty = FieldValue::from_static(",,,"); + assert_eq!( + validate_header_value_list(empty.as_field_value_ref(), &FieldName::SecWebSocketProtocol, validate_alpha,), + Ok(false) + ); + } +} diff --git a/crates/http_headers/src/http_adapter.rs b/crates/http_headers/src/http_adapter.rs new file mode 100644 index 000000000..03c016b00 --- /dev/null +++ b/crates/http_headers/src/http_adapter.rs @@ -0,0 +1,627 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Optional integration with the `http` crate. +//! +//! This module adapts `http::HeaderMap` to [`FieldSource`] and [`FieldSink`] +//! and converts values at the integration boundary. The core types do not use +//! `http` types as their backing representation. + +use std::mem; + +use http::header::Entry; +use http::{HeaderMap, HeaderName, HeaderValue, Request, Response}; +use smallvec::SmallVec; + +use crate::sink::{ + EncodedValues, FieldEncodeOutput, FieldEncoder, FieldSensitivity, FieldSink, FieldValueWriter, InsertError, InsertErrorKind, +}; +use crate::source::{FieldLines, FieldSource}; +use crate::{FieldName, FieldValue}; + +impl FieldSource for HeaderMap { + #[expect(clippy::inline_always, reason = "header decoding must inline the external map adapter")] + #[inline(always)] + fn contains(&self, name: &'static FieldName) -> bool { + match name.http_name() { + Some(http_name) => self.contains_key(http_name), + None => self.contains_key(HeaderName::from_static(name.as_str())), + } + } + + #[expect(clippy::inline_always, reason = "every typed decode crosses this hot adapter boundary")] + #[inline(always)] + fn lines(&self, name: &'static FieldName) -> Option> { + // A well-known name carries the matching `http` constant, while a + // custom typed name supplies validated static bytes without allocation. + match name.http_name() { + Some(http_name) => FieldLines::from_http(name, self.get_all(http_name)), + None => FieldLines::from_http(name, self.get_all(HeaderName::from_static(name.as_str()))), + } + } +} + +impl FieldSink for HeaderMap { + fn set_encoded(&mut self, name: &'static FieldName, encoder: E) -> Result<(), InsertError> + where + E: FieldEncoder, + { + let mut output = HttpOutput::default(); + encoder.encode(&mut output)?; + insert_http_values(self, http_name(name), output.values) + } + + fn append_encoded(&mut self, name: &'static FieldName, encoder: E) -> Result<(), InsertError> + where + E: FieldEncoder, + { + let mut output = HttpOutput::default(); + encoder.encode(&mut output)?; + if output.values.len() > 1 { + return append_http_values(self, http_name(name), output.values); + } + self.try_reserve(output.values.len()) + .map_err(|_full| InsertError::new(InsertErrorKind::CapacityExceeded))?; + let name = http_name(name); + for value in output.values { + self.append(name.clone(), value); + } + Ok(()) + } + + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.set_encoded(name, values) + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.append_encoded(name, values) + } + + fn remove_values(&mut self, name: &'static FieldName) { + match name.http_name() { + Some(http_name) => { + drop(self.remove(http_name)); + } + None => { + drop(self.remove(HeaderName::from_static(name.as_str()))); + } + } + } +} + +macro_rules! impl_http_message { + ($message:ident) => { + impl FieldSource for $message { + fn contains(&self, name: &'static FieldName) -> bool { + FieldSource::contains(self.headers(), name) + } + + fn lines(&self, name: &'static FieldName) -> Option> { + FieldSource::lines(self.headers(), name) + } + } + + impl FieldSink for $message { + fn set_encoded(&mut self, name: &'static FieldName, encoder: E) -> Result<(), InsertError> + where + E: FieldEncoder, + { + self.headers_mut().set_encoded(name, encoder) + } + + fn append_encoded(&mut self, name: &'static FieldName, encoder: E) -> Result<(), InsertError> + where + E: FieldEncoder, + { + self.headers_mut().append_encoded(name, encoder) + } + + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.headers_mut().set_values(name, values) + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.headers_mut().append_values(name, values) + } + + fn remove_values(&mut self, name: &'static FieldName) { + self.headers_mut().remove_values(name); + } + } + }; +} + +impl_http_message!(Request); +impl_http_message!(Response); + +fn http_name(name: &'static FieldName) -> HeaderName { + match name.http_name() { + Some(known) => known.clone(), + None => HeaderName::from_static(name.as_str()), + } +} + +#[derive(Default)] +struct HttpOutput { + values: SmallVec<[HeaderValue; 1]>, +} + +struct HttpWriter<'a> { + values: &'a mut SmallVec<[HeaderValue; 1]>, + state: WriterState, + expected: usize, + sensitive: bool, +} + +/// What a writer has taken so far. +/// +/// A borrowed encoder hands over its whole field line in one `write_bytes` of +/// an already-contiguous slice, which becomes the header value directly: the +/// value is validated and copied once, and no intermediate buffer is ever +/// created. Anything that writes in pieces — or that writes a length other +/// than the one it announced — accumulates into a buffer instead, so the +/// length and validity checks still happen in `finish`. +enum WriterState { + Empty, + Single(HeaderValue), + Buffered(SmallVec<[u8; 32]>), +} + +fn buffered(expected: usize, bytes: &[u8]) -> Result, InsertError> { + let mut buffer = SmallVec::new(); + buffer + .try_reserve_exact(expected) + .map_err(|_error| InsertError::new(InsertErrorKind::AllocationFailed))?; + buffer.extend_from_slice(bytes); + Ok(buffer) +} + +fn try_push_http_value(values: &mut SmallVec<[HeaderValue; 1]>, value: HeaderValue) -> Result<(), InsertError> { + reserve_http_values(values, 1)?; + values.push(value); + Ok(()) +} + +fn reserve_http_values(values: &mut SmallVec<[HeaderValue; 1]>, additional: usize) -> Result<(), InsertError> { + values + .try_reserve(additional) + .map_err(|_error| InsertError::new(InsertErrorKind::AllocationFailed))?; + Ok(()) +} + +impl FieldEncodeOutput for HttpOutput { + type Writer<'a> = HttpWriter<'a>; + + fn begin_value(&mut self, length: usize, sensitivity: FieldSensitivity) -> Result, InsertError> { + if length > isize::MAX as usize { + return Err(InsertError::new(InsertErrorKind::CapacityExceeded)); + } + Ok(HttpWriter { + values: &mut self.values, + state: WriterState::Empty, + expected: length, + sensitive: sensitivity.is_sensitive(), + }) + } + + fn push_value(&mut self, value: FieldValue) -> Result<(), InsertError> { + let value = HeaderValue::try_from(value).map_err(|_invalid| InsertError::new(InsertErrorKind::InvalidValue))?; + try_push_http_value(&mut self.values, value) + } + + fn push_u64(&mut self, value: u64) -> Result<(), InsertError> { + try_push_http_value(&mut self.values, HeaderValue::from(value)) + } +} + +impl FieldValueWriter for HttpWriter<'_> { + fn write_bytes(&mut self, bytes: &[u8]) -> Result<(), InsertError> { + let written = match &self.state { + WriterState::Empty => 0, + WriterState::Single(value) => value.as_bytes().len(), + WriterState::Buffered(buffer) => buffer.len(), + }; + if bytes.len() > self.expected.saturating_sub(written) { + return Err(InsertError::new(InsertErrorKind::InvalidEncoding)); + } + self.state = match mem::replace(&mut self.state, WriterState::Empty) { + // The single-shot case: the announced length arrives whole, so it + // becomes the header value here and `finish` only has to push it. + // An invalid slice falls back to the buffer so that the error + // still surfaces from `finish`, as it does for every other writer. + WriterState::Empty if bytes.len() == self.expected => match HeaderValue::from_bytes(bytes) { + Ok(value) => WriterState::Single(value), + Err(_invalid) => WriterState::Buffered(buffered(self.expected, bytes)?), + }, + WriterState::Empty => WriterState::Buffered(buffered(self.expected, bytes)?), + WriterState::Single(value) => WriterState::Single(value), + WriterState::Buffered(mut buffer) => { + buffer.extend_from_slice(bytes); + WriterState::Buffered(buffer) + } + }; + Ok(()) + } + + fn finish(self) -> Result<(), InsertError> { + let mut value = match self.state { + WriterState::Single(value) => value, + WriterState::Empty => { + if self.expected != 0 { + return Err(InsertError::new(InsertErrorKind::InvalidEncoding)); + } + HeaderValue::from_static("") + } + WriterState::Buffered(bytes) => { + if bytes.len() != self.expected { + return Err(InsertError::new(InsertErrorKind::InvalidEncoding)); + } + if bytes.spilled() { + // `into_vec` retains the spilled allocation; doing this for + // inline bytes would allocate before constructing the value. + HeaderValue::try_from(bytes.into_vec()).map_err(|_invalid| InsertError::new(InsertErrorKind::InvalidValue))? + } else { + HeaderValue::from_bytes(&bytes).map_err(|_invalid| InsertError::new(InsertErrorKind::InvalidValue))? + } + } + }; + value.set_sensitive(self.sensitive); + try_push_http_value(self.values, value) + } +} + +fn append_http_values(map: &mut HeaderMap, name: HeaderName, encoded: SmallVec<[HeaderValue; 1]>) -> Result<(), InsertError> { + map.try_reserve(encoded.len()) + .map_err(|_full| InsertError::new(InsertErrorKind::CapacityExceeded))?; + let mut values = encoded.into_iter(); + let first = values + .next() + .expect("append_http_values is called only for multiple encoded values"); + match map + .try_entry(name) + .map_err(|_full| InsertError::new(InsertErrorKind::CapacityExceeded))? + { + Entry::Occupied(mut entry) => { + entry.append(first); + for value in values { + entry.append(value); + } + } + Entry::Vacant(entry) => { + let mut entry = entry + .try_insert_entry(first) + .map_err(|_full| InsertError::new(InsertErrorKind::CapacityExceeded))?; + for value in values { + entry.append(value); + } + } + } + Ok(()) +} + +/// Replaces every value stored under `name` with the encoded field lines. +fn insert_http_values(map: &mut HeaderMap, name: HeaderName, encoded: SmallVec<[HeaderValue; 1]>) -> Result<(), InsertError> { + map.try_reserve(encoded.len()) + .map_err(|_full| InsertError::new(InsertErrorKind::CapacityExceeded))?; + let mut values = encoded.into_iter(); + let Some(first) = values.next() else { + map.remove(&name); + return Ok(()); + }; + match map + .try_entry(name) + .map_err(|_full| InsertError::new(InsertErrorKind::CapacityExceeded))? + { + Entry::Occupied(mut entry) => { + entry.insert(first); + for value in values { + entry.append(value); + } + } + Entry::Vacant(entry) => { + let mut entry = entry + .try_insert_entry(first) + .map_err(|_full| InsertError::new(InsertErrorKind::CapacityExceeded))?; + for value in values { + entry.append(value); + } + } + } + Ok(()) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::sync::LazyLock; + + use super::{HttpOutput, append_http_values, buffered, insert_http_values, reserve_http_values, try_push_http_value}; + use crate::sink::{EncodedValues, FieldEncodeOutput, FieldSink, FieldValueWriter, InsertError, InsertErrorKind, U64Encoder}; + use crate::source::FieldSource; + use crate::{FieldValue, FieldValueRef}; + + static CUSTOM: LazyLock = LazyLock::new(|| crate::FieldName::from_static("x-trace-id")); + + const ENCODING_ERROR: InsertError = InsertError::new(InsertErrorKind::InvalidEncoding); + const VALUE_ERROR: InsertError = InsertError::new(InsertErrorKind::InvalidValue); + const CAPACITY_ERROR: InsertError = InsertError::new(InsertErrorKind::CapacityExceeded); + + struct RejectEncoder; + + impl crate::sink::FieldEncoder for RejectEncoder { + fn encode(self, _output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + Err(ENCODING_ERROR) + } + } + + #[test] + fn map_source_and_sink_cover_known_and_custom_names_and_entry_states() { + let mut map = http::HeaderMap::new(); + assert!(!map.contains(&crate::FieldName::UserAgent)); + assert!(!map.contains(&CUSTOM)); + assert!(FieldSource::lines(&map, &crate::FieldName::UserAgent).is_none()); + assert!(FieldSource::lines(&map, &CUSTOM).is_none()); + + map.set_values( + &crate::FieldName::UserAgent, + EncodedValues::from_vec(vec![FieldValue::from_static("client/1"), FieldValue::from_static("client/2")]), + ) + .expect("known values insert"); + assert!(map.contains(&crate::FieldName::UserAgent)); + assert_eq!( + FieldSource::lines(&map, &crate::FieldName::UserAgent) + .expect("known values") + .repeated() + .map(crate::FieldValueRef::as_bytes) + .collect::>(), + [b"client/1".as_slice(), b"client/2".as_slice()] + ); + + map.set_values( + &crate::FieldName::UserAgent, + EncodedValues::from_vec(vec![FieldValue::from_static("replacement"), FieldValue::from_static("additional")]), + ) + .expect("occupied entry replaces"); + assert_eq!(map[http::header::USER_AGENT], "replacement"); + assert_eq!(map.get_all(http::header::USER_AGENT).iter().count(), 2); + + map.set_encoded(&CUSTOM, FieldValueRef::new(b"trace").with_sensitive(true)) + .expect("custom value inserts"); + assert!(map.contains(&CUSTOM)); + assert!(map.get("x-trace-id").expect("custom value").is_sensitive()); + + map.append_encoded(&CUSTOM, U64Encoder::new(42)).expect("custom value appends"); + assert_eq!(map.get_all("x-trace-id").iter().count(), 2); + + map.remove_values(&crate::FieldName::UserAgent); + map.remove_values(&CUSTOM); + assert!(!map.contains(&crate::FieldName::UserAgent)); + assert!(!map.contains(&CUSTOM)); + + assert_eq!(map.set_encoded(&crate::FieldName::UserAgent, RejectEncoder), Err(ENCODING_ERROR)); + assert_eq!(map.append_encoded(&crate::FieldName::UserAgent, RejectEncoder), Err(ENCODING_ERROR)); + assert!(!map.contains(&crate::FieldName::UserAgent)); + + insert_http_values(&mut map, http::header::USER_AGENT, smallvec::SmallVec::new()).expect("empty values remove"); + assert!(!map.contains_key(http::header::USER_AGENT)); + + append_http_values( + &mut map, + http::header::SET_COOKIE, + smallvec::smallvec![http::HeaderValue::from_static("a=1"), http::HeaderValue::from_static("b=2"),], + ) + .expect("vacant entry accepts repeated values"); + append_http_values( + &mut map, + http::header::SET_COOKIE, + smallvec::smallvec![http::HeaderValue::from_static("c=3"), http::HeaderValue::from_static("d=4"),], + ) + .expect("occupied entry accepts repeated values"); + assert_eq!( + map.get_all(http::header::SET_COOKIE) + .iter() + .map(http::HeaderValue::as_bytes) + .collect::>(), + [b"a=1".as_slice(), b"b=2".as_slice(), b"c=3".as_slice(), b"d=4".as_slice(),] + ); + } + + #[test] + fn http_output_validates_lengths_bytes_sensitivity_and_native_values() { + let mut output = HttpOutput::default(); + let mut writer = output + .begin_value(3, crate::sink::FieldSensitivity::Sensitive) + .expect("writer starts"); + writer.write_bytes(b"abc").expect("bytes append"); + writer.finish().expect("exact valid value finishes"); + assert_eq!(output.values[0], "abc"); + assert!(output.values[0].is_sensitive()); + + let mut output = HttpOutput::default(); + let mut short = output + .begin_value(2, crate::sink::FieldSensitivity::NonSensitive) + .expect("writer starts"); + short.write_bytes(b"x").expect("bytes append"); + assert_eq!(short.finish(), Err(ENCODING_ERROR)); + + let mut output = HttpOutput::default(); + let mut invalid = output + .begin_value(1, crate::sink::FieldSensitivity::NonSensitive) + .expect("writer starts"); + invalid.write_bytes(b"\n").expect("bytes append"); + assert_eq!(invalid.finish(), Err(VALUE_ERROR)); + + output.push_value(FieldValue::from_static("owned")).expect("owned value transfers"); + output.push_u64(u64::MAX).expect("integer value formats"); + assert_eq!(output.values[0], "owned"); + assert_eq!(output.values[1], u64::MAX.to_string()); + } + + #[test] + fn http_writer_streams_overruns_and_empty_values_without_a_single_shot_buffer() { + // A streaming encoder writes in pieces, so the buffered path has to + // reassemble the announced length. + let mut output = HttpOutput::default(); + let mut writer = output + .begin_value(6, crate::sink::FieldSensitivity::NonSensitive) + .expect("writer starts"); + writer.write_bytes(b"abc").expect("first piece appends"); + writer.write_bytes(b"def").expect("second piece appends"); + writer.finish().expect("streamed value finishes"); + assert_eq!(output.values[0], "abcdef"); + assert!(!output.values[0].is_sensitive()); + + let bytes = [b'x'; 65]; + let mut output = HttpOutput::default(); + let mut writer = output + .begin_value(bytes.len(), crate::sink::FieldSensitivity::Sensitive) + .expect("writer starts"); + for chunk in bytes.chunks(7) { + writer.write_bytes(chunk).expect("chunk appends"); + } + writer.finish().expect("spilled streamed value finishes"); + assert_eq!(output.values[0].as_bytes(), bytes); + assert!(output.values[0].is_sensitive()); + + let mut invalid_bytes = [b'x'; 65]; + invalid_bytes[32] = b'\n'; + let mut output = HttpOutput::default(); + let mut writer = output + .begin_value(invalid_bytes.len(), crate::sink::FieldSensitivity::NonSensitive) + .expect("writer starts"); + for chunk in invalid_bytes.chunks(7) { + writer.write_bytes(chunk).expect("chunk appends"); + } + assert_eq!(writer.finish(), Err(VALUE_ERROR)); + assert!(output.values.is_empty()); + + // A complete single-shot write rejects more bytes before buffering. + let mut output = HttpOutput::default(); + let mut writer = output + .begin_value(3, crate::sink::FieldSensitivity::NonSensitive) + .expect("writer starts"); + writer.write_bytes(b"abc").expect("whole value appends"); + writer.write_bytes(b"").expect("an empty continuation is a no-op"); + assert_eq!(writer.write_bytes(b"d"), Err(ENCODING_ERROR)); + assert!(output.values.is_empty()); + + // Nothing written at all: legal only for a zero-length value. + let mut output = HttpOutput::default(); + let writer = output + .begin_value(0, crate::sink::FieldSensitivity::NonSensitive) + .expect("writer starts"); + writer.finish().expect("empty value finishes"); + assert_eq!(output.values[0], ""); + + let mut output = HttpOutput::default(); + let writer = output + .begin_value(1, crate::sink::FieldSensitivity::NonSensitive) + .expect("writer starts"); + assert_eq!(writer.finish(), Err(ENCODING_ERROR)); + assert!(output.values.is_empty()); + } + + #[test] + fn http_output_distinguishes_size_limits_from_reservation_failures() { + let mut output = HttpOutput::default(); + assert_eq!( + output.begin_value(usize::MAX, crate::sink::FieldSensitivity::NonSensitive).err(), + Some(CAPACITY_ERROR) + ); + assert!(output.values.is_empty()); + assert_eq!(buffered(usize::MAX, b"x"), Err(InsertError::new(InsertErrorKind::AllocationFailed))); + + let mut values = smallvec::smallvec![http::HeaderValue::from_static("original")]; + assert_eq!( + reserve_http_values(&mut values, usize::MAX), + Err(InsertError::new(InsertErrorKind::AllocationFailed)) + ); + assert_eq!(values.as_slice(), &[http::HeaderValue::from_static("original")]); + try_push_http_value(&mut values, http::HeaderValue::from_static("second")).unwrap(); + assert_eq!(values.len(), 2); + } + + /// The number of entries an `http::HeaderMap` holds once it can no longer + /// grow: `http` caps its index table at `1 << 15` slots and keeps three + /// quarters of them usable. + const MAXIMUM_ENTRIES: usize = (1 << 15) - (1 << 13); + + /// Builds a map that has reached `http`'s maximum entry count, so every + /// `try_reserve` for an additional entry fails. + fn saturated_map() -> http::HeaderMap { + let mut map = http::HeaderMap::with_capacity(MAXIMUM_ENTRIES); + map.insert(http::header::USER_AGENT, http::HeaderValue::from_static("original")); + #[cfg(not(miri))] + let mut filler = String::new(); + for index in 0..MAXIMUM_ENTRIES - 1 { + #[cfg(not(miri))] + let name = { + filler.clear(); + filler.push_str("x-fill-"); + filler.push_str(&index.to_string()); + http::HeaderName::from_bytes(filler.as_bytes()).expect("generated name is legal") + }; + #[cfg(miri)] + let name = http::HeaderName::from_static(crate::miri_http_map::name(index)); + drop(map.insert(name, http::HeaderValue::from_static("v"))); + } + assert_eq!(map.len(), MAXIMUM_ENTRIES); + assert!(map.try_reserve(1).is_err(), "the map must be unable to grow"); + map + } + + #[test] + fn capacity_exhaustion_reports_an_error_without_changing_the_map() { + let mut map = saturated_map(); + + // A new name needs a new entry the map cannot make room for. + assert_eq!(map.set_encoded(&CUSTOM, FieldValueRef::new(b"trace")), Err(CAPACITY_ERROR)); + assert!(!map.contains(&CUSTOM)); + assert_eq!( + map.set_values(&CUSTOM, EncodedValues::single(FieldValue::from_static("trace"))), + Err(CAPACITY_ERROR) + ); + assert!(!map.contains(&CUSTOM)); + assert_eq!(map.append_encoded(&CUSTOM, FieldValueRef::new(b"trace")), Err(CAPACITY_ERROR)); + assert!(!map.contains(&CUSTOM)); + + // An existing name is rejected too: `try_reserve` guards the whole + // insertion, and the values already stored stay exactly as they were. + assert_eq!( + map.set_values( + &crate::FieldName::UserAgent, + EncodedValues::from_vec(vec![FieldValue::from_static("replacement"), FieldValue::from_static("additional"),]), + ), + Err(CAPACITY_ERROR) + ); + assert_eq!( + map.append_encoded(&crate::FieldName::UserAgent, FieldValueRef::new(b"appended")), + Err(CAPACITY_ERROR) + ); + assert_eq!(map[http::header::USER_AGENT], "original"); + assert_eq!(map.get_all(http::header::USER_AGENT).iter().count(), 1); + assert_eq!(map.len(), MAXIMUM_ENTRIES); + + let mut values = smallvec::SmallVec::<[http::HeaderValue; 1]>::new(); + values.push(http::HeaderValue::from_static("appended")); + assert_eq!(append_http_values(&mut map, http::header::USER_AGENT, values), Err(CAPACITY_ERROR)); + assert_eq!(map[http::header::USER_AGENT], "original"); + + // The private helper is reached the same way through both sinks; call + // it directly so the failure is attributed to it and not to a caller. + // Its `try_entry`/`try_insert_entry` error arms stay defensive: the + // `try_reserve` above already fails for every additional entry the + // map cannot make room for, so nothing reaches them first. + let mut values = smallvec::SmallVec::<[http::HeaderValue; 1]>::new(); + values.push(http::HeaderValue::from_static("direct")); + assert_eq!( + insert_http_values(&mut map, http::HeaderName::from_static("x-direct"), values), + Err(CAPACITY_ERROR) + ); + assert!(!map.contains_key("x-direct")); + assert_eq!(map.len(), MAXIMUM_ENTRIES); + } +} diff --git a/crates/http_headers/src/lib.rs b/crates/http_headers/src/lib.rs new file mode 100644 index 000000000..ab2051e6b --- /dev/null +++ b/crates/http_headers/src/lib.rs @@ -0,0 +1,547 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +#![cfg_attr(coverage_nightly, feature(coverage_attribute))] +#![doc(html_logo_url = "https://media.githubusercontent.com/media/microsoft/oxidizer/refs/heads/main/crates/http_headers/logo.png")] +#![doc(html_favicon_url = "https://media.githubusercontent.com/media/microsoft/oxidizer/refs/heads/main/crates/http_headers/favicon.ico")] + +//! Efficient and robust HTTP header parsing and creation. +//! +//! This crate provides: +//! +//! - Highly optimized parsing of incoming HTTP headers which produce owned or borrowed +//! strongly-typed Rust structs. These parsers insulate your code from badly formed +//! headers. +//! +//! - Highly optimized production of headers, ensuring the headers are well-formed. +//! +//! Header parsing and production are abstracted over their source and destination. +//! The optional `http` feature integrates with +//! [`HeaderMap`](https://docs.rs/http/latest/http/header/struct.HeaderMap.html) +//! plus generic [`Request`](https://docs.rs/http/latest/http/request/struct.Request.html) +//! and [`Response`](https://docs.rs/http/latest/http/response/struct.Response.html) +//! values from the [`http`](https://crates.io/crates/http) crate. +//! +//! # Parsing headers +//! +//! Headers are parsed from an implementation of the [`source::FieldSource`] trait. The `http` crate feature +//! implements this trait for [`HeaderMap`](https://docs.rs/http/latest/http/header/struct.HeaderMap.html), +//! [`Request`](https://docs.rs/http/latest/http/request/struct.Request.html), and +//! [`Response`](https://docs.rs/http/latest/http/response/struct.Response.html). +//! Once you have a source, you can choose to parse into borrowed views or owned structs. +//! Prefer borrowed views when the decoded value does not need to outlive the +//! source as they are generally faster. Use owned structs when the parsed header +//! data needs to be retained (such as in a cache). +//! +//! Source and sink operations use static field-name descriptors, including +//! custom names stored in a `static LazyLock`. Locally constructed +//! runtime names are supported by [`FieldName`] for validation and conversion, +//! but dynamic lookup and mutation must use the container's native API. +//! +//! [`Field::view`] returns a header's +//! borrowed `*View` type, whose lifetime is tied to the source. +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-all"))] +//! # fn main() -> Result<(), Box> { +//! use http::HeaderMap; +//! use http_headers::Field; +//! use http_headers::headers::{ContentType, UserAgent}; +//! +//! // create a HeaderMap to show how to read from it +//! let mut headers = HeaderMap::new(); +//! headers.insert( +//! http::header::USER_AGENT, +//! http::HeaderValue::from_static("example-client/1.0"), +//! ); +//! headers.insert( +//! http::header::CONTENT_TYPE, +//! http::HeaderValue::from_static("application/json; charset=utf-8"), +//! ); +//! +//! if let Some(agent) = UserAgent::view(&headers)? { +//! assert_eq!(agent.as_str()?, "example-client/1.0"); +//! } +//! +//! if let Some(content_type) = ContentType::view(&headers)? { +//! assert_eq!(content_type.type_()?, "application"); +//! assert_eq!(content_type.subtype()?, "json"); +//! assert_eq!( +//! content_type.parameter("charset")?, +//! Some(b"utf-8".as_slice()) +//! ); +//! } +//! # Ok::<(), Box>(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! Prefer `view` unless the decoded value must outlive the source. Use +//! [`Field::owned`] when you need to retain, move, or independently +//! store the result: +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-all"))] +//! # fn main() -> Result<(), http_headers::DecodeError> { +//! use http::HeaderMap; +//! use http_headers::Field; +//! use http_headers::headers::UserAgent; +//! +//! let mut headers = HeaderMap::new(); +//! headers.insert( +//! http::header::USER_AGENT, +//! http::HeaderValue::from_static("example-client/1.0"), +//! ); +//! +//! let owned = UserAgent::owned(&headers)?.expect("User-Agent is present"); +//! drop(headers); +//! assert_eq!(owned.as_bytes(), b"example-client/1.0"); +//! # Ok::<(), http_headers::DecodeError>(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! Both methods return `Ok(None)` when the header is absent and `Err` when a +//! present value is malformed. +//! +//! ## Reading validated members +//! +//! Structured headers expose semantic values as well as their original wire +//! representation. `AllowOwned::methods()` yields case-sensitive method tokens; +//! `VaryOwned::entries()` distinguishes wildcard members from case-insensitive +//! field names. These borrowed member reads do not allocate. +//! +//! ```rust +//! # #[cfg(feature = "headers-negotiation")] +//! # fn main() -> Result<(), Box> { +//! use http_headers::headers::{AllowOwned, MethodView, VaryOwned}; +//! +//! let allow = AllowOwned::try_from("GET, HEAD, CUSTOM")?; +//! assert!(allow.methods().any(|method| method == MethodView::GET)); +//! assert!(!allow.methods().any(|method| method == MethodView::POST)); +//! +//! let vary = VaryOwned::try_from("Accept-Encoding, X-Tenant")?; +//! assert!(!vary.contains_wildcard()); +//! assert!(vary.entries().any(|entry| { +//! entry +//! .field_name() +//! .is_some_and(|name| name.eq_ignore_ascii_case("x-tenant")) +//! })); +//! # Ok(()) +//! # } +//! # #[cfg(not(feature = "headers-negotiation"))] +//! # fn main() {} +//! ``` +//! +//! The Accept family exposes typed ranges, parameters, and exact quality +//! weights through `entries()`. Location exposes URI-reference components +//! through `uri_reference()`. Host and Allow-Origin retain parsed authority +//! components. Semantic access does not sort lists or replace the original +//! field lines used for forwarding. +//! +//! # Producing headers +//! +//! You produce headers by populating an implementation of the [`sink::FieldSink`] trait. The +//! `http` cargo feature implements this trait for +//! [`HeaderMap`](https://docs.rs/http/latest/http/header/struct.HeaderMap.html), +//! [`Request`](https://docs.rs/http/latest/http/request/struct.Request.html), and +//! [`Response`](https://docs.rs/http/latest/http/response/struct.Response.html). +//! +//! Enabling a header-family feature exposes `sink::FieldSinkExt`, whose fluent +//! methods work for any sink. Core-only builds use [`sink::FieldSink`] directly. +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-all"))] +//! # fn main() -> Result<(), http_headers::sink::InsertError> { +//! use std::time::Duration; +//! +//! use http::HeaderMap; +//! use http_headers::headers::{CacheControl, ContentType}; +//! use http_headers::sink::FieldSinkExt; +//! +//! let mut headers = HeaderMap::new(); +//! headers +//! .set_content_type(ContentType::json())? +//! .set_content_length(1_024)? +//! .set_cache_control(CacheControl::public().max_age(Duration::from_secs(60)))?; +//! # Ok::<(), http_headers::sink::InsertError>(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! An owned value can insert itself when it has already been constructed: +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-all"))] +//! # fn main() -> Result<(), Box> { +//! use http::HeaderMap; +//! use http_headers::headers::LocationOwned; +//! +//! let mut headers = HeaderMap::new(); +//! LocationOwned::try_from("/next")?.insert_into(&mut headers)?; +//! # Ok::<(), Box>(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! A borrowed view can also be forwarded directly to another sink: +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-all"))] +//! # fn main() -> Result<(), Box> { +//! use http::HeaderMap; +//! use http_headers::Field; +//! use http_headers::headers::UserAgent; +//! +//! let mut incoming = HeaderMap::new(); +//! incoming.insert( +//! http::header::USER_AGENT, +//! http::HeaderValue::from_static("example-client/1.0"), +//! ); +//! +//! let mut outgoing = HeaderMap::new(); +//! if let Some(agent) = UserAgent::view(&incoming)? { +//! agent.insert_into(&mut outgoing)?; +//! } +//! # Ok::<(), Box>(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! [`Field::insert`] is the generic alternative when the descriptor type is +//! already known. It replaces all existing field lines for that header; +//! [`Field::remove`] removes them instead. +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-all"))] +//! # fn main() -> Result<(), Box> { +//! use http::HeaderMap; +//! use http_headers::Field; +//! use http_headers::headers::{UserAgent, UserAgentOwned}; +//! +//! let mut headers = HeaderMap::new(); +//! UserAgent::insert( +//! &mut headers, +//! UserAgentOwned::try_from_static("example-client/1.0")?, +//! )?; +//! UserAgent::remove(&mut headers); +//! # Ok::<(), Box>(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! Repeated field lines remain separate. In particular, `Set-Cookie` values are +//! never comma-joined: +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-all"))] +//! # fn main() -> Result<(), Box> { +//! use http::HeaderMap; +//! use http_headers::Field; +//! use http_headers::headers::{SetCookie, SetCookieOwned}; +//! +//! let mut cookies = SetCookieOwned::new(); +//! cookies.push_str("session=abc; Path=/; HttpOnly")?; +//! cookies.push_str("theme=dark; Path=/")?; +//! +//! let mut headers = HeaderMap::new(); +//! SetCookie::insert(&mut headers, cookies)?; +//! assert_eq!(headers.get_all(http::header::SET_COOKIE).iter().count(), 2); +//! # Ok::<(), Box>(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! # Serialization +//! +//! The `serde` cargo feature implements Serde serialization and deserialization for every owned header +//! struct, [`FieldName`], [`FieldValue`], and [`sink::EncodedValues`]. +//! Headers serialize as an ordered sequence of physical field values and +//! deserialize through relaxed validation, which includes strict syntax and +//! the documented interoperability deviations. This preserves round trips for +//! every owned value produced by the public API: +//! +//! ```rust +//! # #[cfg(all(feature = "serde", feature = "headers-all"))] +//! # fn main() -> Result<(), Box> { +//! use http_headers::headers::UserAgentOwned; +//! +//! let header = UserAgentOwned::try_from("example-client/1.0")?; +//! let json = serde_json::to_string(&header)?; +//! let decoded: UserAgentOwned = serde_json::from_str(&json)?; +//! assert_eq!(decoded.as_bytes(), header.as_bytes()); +//! # Ok(()) +//! # } +//! # #[cfg(not(all(feature = "serde", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! Repeated lines retain their boundaries: +//! +//! ```rust +//! # #[cfg(all(feature = "serde", feature = "headers-all"))] +//! # fn main() -> Result<(), Box> { +//! use http_headers::headers::SetCookieOwned; +//! +//! let mut cookies = SetCookieOwned::new(); +//! cookies.push_str("session=abc")?; +//! cookies.push_str("theme=dark")?; +//! let json = serde_json::to_string(&cookies)?; +//! let decoded: SetCookieOwned = serde_json::from_str(&json)?; +//! assert_eq!(decoded.len(), 2); +//! # Ok(()) +//! # } +//! # #[cfg(not(all(feature = "serde", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! Serialization is not redaction. Sensitive values include their original +//! bytes and an explicit sensitivity marker, so serialized data must be +//! protected like the header value itself: +//! +//! ```rust +//! # #[cfg(feature = "serde")] +//! # fn main() -> Result<(), Box> { +//! use http_headers::{FieldSensitivity, FieldValue}; +//! +//! let secret = +//! FieldValue::from_static("credential").with_sensitivity(FieldSensitivity::Sensitive); +//! let json = serde_json::to_string(&secret)?; +//! let decoded: FieldValue = serde_json::from_str(&json)?; +//! assert_eq!(decoded.as_bytes(), b"credential"); +//! assert!(decoded.is_sensitive()); +//! # Ok(()) +//! # } +//! # #[cfg(not(feature = "serde"))] +//! # fn main() {} +//! ``` +//! +//! # Strict and relaxed reads +//! +//! [`Field::view`] and [`Field::owned`] use strict syntax. Applications that +//! must accept specific common deviations can request [`DecodeMode::Relaxed`] +//! through [`Field::view_with`] or [`Field::owned_with`]: +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-all"))] +//! # fn main() -> Result<(), Box> { +//! use http_headers::headers::AcceptEncoding; +//! use http_headers::{DecodeMode, Field}; +//! +//! let mut headers = http::HeaderMap::new(); +//! headers.insert( +//! http::header::ACCEPT_ENCODING, +//! http::HeaderValue::from_static("gzip; q = .5"), +//! ); +//! +//! assert!(AcceptEncoding::view(&headers).is_err()); +//! assert!(AcceptEncoding::view_with(&headers, DecodeMode::Relaxed)?.is_some()); +//! # Ok::<(), Box>(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! Relaxed mode is not a general validation bypass. Each header documents the +//! additional forms it accepts, and the original field bytes are preserved. +//! +//! # Sensitive values +//! +//! Authorization, `Location`, and cookie values are marked sensitive so their +//! `Debug` representations and compatible sinks do not reveal their contents. +//! Basic authentication can be read through a borrowed view while reusing +//! caller-owned decode storage: +//! +//! ```rust +//! # #[cfg(all(feature = "http", feature = "headers-all"))] +//! # fn main() -> Result<(), http_headers::DecodeError> { +//! use http::HeaderMap; +//! use http_headers::Field; +//! use http_headers::headers::{Authorization, Basic, BasicCredentials}; +//! +//! let mut headers = HeaderMap::new(); +//! headers.insert( +//! http::header::AUTHORIZATION, +//! http::HeaderValue::from_static("Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ=="), +//! ); +//! +//! let authorization = Authorization::::view(&headers)?.expect("Authorization is present"); +//! let mut credentials = BasicCredentials::new(); +//! let decoded = authorization.extract(&mut credentials)?; +//! assert_eq!(decoded.username(), b"Aladdin"); +//! credentials.clear(); +//! # Ok::<(), http_headers::DecodeError>(()) +//! # } +//! # #[cfg(not(all(feature = "http", feature = "headers-all")))] +//! # fn main() {} +//! ``` +//! +//! `BasicCredentials` zeroizes decoded bytes when cleared, reused, or dropped. +//! +//! # Defining a custom single-value header +//! +//! Implement [`SingleValueField`] when a custom header is represented by exactly +//! one field line. The crate then supplies its [`Field`] implementation, +//! including borrowed and owned reads, singleton cardinality checks, insertion, +//! and removal. +//! +//! ```rust +//! use std::sync::LazyLock; +//! +//! use http_headers::{ +//! DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef, SingleValueField, +//! }; +//! +//! static REQUEST_ID: LazyLock = +//! LazyLock::new(|| FieldName::from_static("x-request-id")); +//! +//! struct RequestId; +//! +//! #[derive(Clone, Debug, Eq, PartialEq)] +//! struct RequestIdOwned(FieldValue); +//! +//! #[derive(Clone, Copy, Debug, Eq, PartialEq)] +//! struct RequestIdView<'a>(FieldValueRef<'a>); +//! +//! fn is_token(bytes: &[u8]) -> bool { +//! !bytes.is_empty() +//! && bytes +//! .iter() +//! .all(|byte| byte.is_ascii_alphanumeric() || b"!#$%&'*+-.^_`|~".contains(byte)) +//! } +//! +//! impl SingleValueField for RequestId { +//! type View<'a> = RequestIdView<'a>; +//! type Owned = RequestIdOwned; +//! +//! fn name() -> &'static FieldName { +//! &REQUEST_ID +//! } +//! +//! fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { +//! if is_token(value.as_bytes()) { +//! Ok(RequestIdView(value)) +//! } else { +//! Err(DecodeError::new(&REQUEST_ID, DecodeErrorKind::InvalidToken)) +//! } +//! } +//! +//! fn decode_owned(value: FieldValue) -> Result { +//! if is_token(value.as_bytes()) { +//! Ok(RequestIdOwned(value)) +//! } else { +//! Err(DecodeError::new(&REQUEST_ID, DecodeErrorKind::InvalidToken)) +//! } +//! } +//! +//! fn as_field_value(value: &Self::Owned) -> &FieldValue { +//! &value.0 +//! } +//! +//! fn into_field_value(value: Self::Owned) -> FieldValue { +//! value.0 +//! } +//! } +//! # Ok::<(), DecodeError>(()) +//! ``` +//! +//! # Performance +//! +//! [`docs/PERF.md`](https://github.com/microsoft/oxidizer/blob/main/crates/http_headers/docs/PERF.md) +//! records comparative typed-decode-and-read measurements against `headers 0.4.1`. +//! Results vary by header, ownership mode, and hardware, and include both faster +//! and slower cases. The table does not measure comparative header production or +//! end-to-end request processing. Borrowed reads generally avoid allocations. +//! +//! # Cargo features +//! +//! - `headers-all` (enabled by default): all built-in typed header families. +//! - `headers-authorization`, `headers-cache-control`, `headers-conditional`, +//! `headers-content-length`, `headers-content-type`, `headers-cors`, +//! `headers-etag`, `headers-location`, `headers-negotiation`, `headers-range`, +//! `headers-security`, `headers-set-cookie`, `headers-user-agent`, and +//! `headers-websocket`: individual built-in header families. +//! - `http`: optional adapter for `http::HeaderMap` and the `http` crate's name, +//! value, and method types. +//! - `serde`: serialization and deserialization for owned headers, +//! [`FieldName`], [`FieldValue`], and [`sink::EncodedValues`]. +//! +//! Disable default features to use only the core source, sink, name, and value +//! APIs, then enable only the header families an application needs. +//! +//! # What about trailers? +//! +//! Although this crate is named `http_headers`, it fully supports trailers as well. +//! The crate doesn't currently expose any trailer-specific structs however, so you +//! would need to define those structs and implement the parsers yourself as implementations +//! of the traits in this crate. +//! +//! # Alternate crates +//! +//! This crate is an alternative to the popular [`headers`](https://crates.io/crates/headers) crate. +//! `http_headers` has the following benefits: +//! +//! - Faster decoding for some headers in the measured configurations +//! - Supports more headers +//! - Performs more robust validation to avoid downstream surprises +//! - Supports explicit relaxed parsing options to support common malformed headers +//! - Supports serde + +#![forbid(unsafe_code)] + +mod decode_error; +mod field; +mod field_name; +mod field_value; +#[cfg(any( + test, + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +pub mod headers; +#[cfg(feature = "http")] +mod http_adapter; +#[cfg(all(test, miri, feature = "http"))] +mod miri_http_map; +#[cfg(feature = "serde")] +mod serde_impls; +pub mod sink; +pub mod source; +#[cfg(test)] +mod test_sink; +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod test_support; +mod validate; + +#[doc(inline)] +pub use decode_error::{DecodeError, DecodeErrorKind}; +#[doc(inline)] +pub use field::{DecodeMode, Field, SingleValueField}; +#[doc(inline)] +pub use field_name::{FieldName, InvalidFieldName}; +#[doc(inline)] +pub use field_value::{FieldValue, FieldValueRef, InvalidFieldValue}; +#[doc(inline)] +pub use sink::field_encoder::FieldSensitivity; +#[cfg(test)] +pub(crate) use test_sink::TestSink; diff --git a/crates/http_headers/src/miri_http_map.rs b/crates/http_headers/src/miri_http_map.rs new file mode 100644 index 000000000..78efaea6d --- /dev/null +++ b/crates/http_headers/src/miri_http_map.rs @@ -0,0 +1,30 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Static names avoid per-entry formatting and allocation in real-capacity Miri fixtures. + +#![cfg(test)] + +use std::str; + +pub(crate) const CAPACITY: usize = (1 << 15) - (1 << 13); + +static NAMES: [[u8; 6]; CAPACITY] = { + let mut names = [*b"x-aaaa"; CAPACITY]; + let mut index = 0; + while index < CAPACITY { + let mut value = index; + let mut position = 2; + while position < 6 { + names[index][position] = b"abcdefghijklmnopqrstuvwxyz"[value % 26]; + value /= 26; + position += 1; + } + index += 1; + } + names +}; + +pub(crate) fn name(index: usize) -> &'static str { + str::from_utf8(&NAMES[index]).unwrap() +} diff --git a/crates/http_headers/src/serde_impls.rs b/crates/http_headers/src/serde_impls.rs new file mode 100644 index 000000000..bbfabf438 --- /dev/null +++ b/crates/http_headers/src/serde_impls.rs @@ -0,0 +1,957 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::fmt; +#[cfg(any( + feature = "headers-authorization", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +use std::iter; + +use serde::de::{DeserializeSeed, Error as _, IgnoredAny, MapAccess, SeqAccess, Visitor}; +use serde::{Deserialize, Deserializer, Serialize, Serializer}; + +#[cfg(any( + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +use crate::DecodeMode; +#[cfg(any( + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +use crate::Field; +#[cfg(any( + feature = "headers-authorization", + feature = "headers-conditional", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +use crate::SingleValueField; +#[cfg(any( + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +use crate::headers::*; +use crate::sink::{EncodedValues, FieldSensitivity}; +#[cfg(any( + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +use crate::source::{FieldLines, FieldSource, MAX_CUSTOM_FIELD_BYTES, MAX_CUSTOM_FIELD_LINES}; +use crate::{DecodeErrorKind, FieldName, FieldValue, FieldValueRef}; + +const FIELD_VALUE_FIELDS: &[&str] = &["bytes", "sensitivity"]; +// Bound speculative allocation from untrusted Serde size hints. +const SIZE_HINT_RESERVE_LIMIT: usize = 1_024; + +#[derive(Serialize)] +struct FieldValueRepr<'a> { + bytes: &'a [u8], + sensitivity: FieldSensitivity, +} + +#[derive(Clone, Copy)] +enum FieldValueField { + Bytes, + Sensitivity, + Other, +} + +impl<'de> Deserialize<'de> for FieldValueField { + fn deserialize>(deserializer: D) -> Result { + struct FieldVisitor; + + impl Visitor<'_> for FieldVisitor { + type Value = FieldValueField; + + #[cfg_attr(coverage_nightly, coverage(off))] + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("`bytes` or `sensitivity`") + } + + fn visit_str(self, v: &str) -> Result { + Ok(match v { + "bytes" => FieldValueField::Bytes, + "sensitivity" => FieldValueField::Sensitivity, + _ => FieldValueField::Other, + }) + } + + fn visit_bytes(self, v: &[u8]) -> Result { + Ok(match v { + b"bytes" => FieldValueField::Bytes, + b"sensitivity" => FieldValueField::Sensitivity, + _ => FieldValueField::Other, + }) + } + } + + deserializer.deserialize_identifier(FieldVisitor) + } +} + +struct FieldBytesSeed { + limit: Option, +} + +impl<'de> DeserializeSeed<'de> for FieldBytesSeed { + type Value = Vec; + + fn deserialize>(self, deserializer: D) -> Result { + struct BytesVisitor { + limit: Option, + } + + impl<'de> Visitor<'de> for BytesVisitor { + type Value = Vec; + + #[cfg_attr(coverage_nightly, coverage(off))] + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self.limit { + Some(limit) => write!(formatter, "at most {limit} field-value bytes"), + None => formatter.write_str("field-value bytes"), + } + } + + fn visit_seq>(self, mut seq: A) -> Result { + let capacity = seq + .size_hint() + .unwrap_or(0) + .min(self.limit.unwrap_or(SIZE_HINT_RESERVE_LIMIT)) + .min(SIZE_HINT_RESERVE_LIMIT); + let mut bytes = Vec::with_capacity(capacity); + while self.limit.is_none_or(|limit| bytes.len() < limit) { + match seq.next_element()? { + Some(byte) => bytes.push(byte), + None => return Ok(bytes), + } + } + if seq.next_element::()?.is_some() { + return Err(A::Error::custom(format_args!( + "{}: field-value byte budget exceeded", + DecodeErrorKind::SourceLimitExceeded + ))); + } + Ok(bytes) + } + } + + deserializer.deserialize_seq(BytesVisitor { limit: self.limit }) + } +} + +struct FieldValueSeed { + byte_limit: Option, +} + +impl<'de> DeserializeSeed<'de> for FieldValueSeed { + type Value = FieldValue; + + fn deserialize>(self, deserializer: D) -> Result { + struct FieldValueVisitor { + byte_limit: Option, + } + + impl FieldValueVisitor { + fn finish(bytes: Option>, sensitivity: Option) -> Result { + let bytes = bytes.ok_or_else(|| E::missing_field("bytes"))?; + let sensitivity = sensitivity.ok_or_else(|| E::missing_field("sensitivity"))?; + FieldValue::try_from(bytes) + .map(|value| value.with_sensitivity(sensitivity)) + .map_err(E::custom) + } + } + + impl<'de> Visitor<'de> for FieldValueVisitor { + type Value = FieldValue; + + #[cfg_attr(coverage_nightly, coverage(off))] + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("a field value with bytes and sensitivity") + } + + fn visit_seq>(self, mut seq: A) -> Result { + let bytes = seq + .next_element_seed(FieldBytesSeed { limit: self.byte_limit })? + .ok_or_else(|| A::Error::invalid_length(0, &self))?; + let sensitivity = seq.next_element()?.ok_or_else(|| A::Error::invalid_length(1, &self))?; + Self::finish(Some(bytes), Some(sensitivity)) + } + + fn visit_map>(self, mut map: A) -> Result { + let mut bytes = None; + let mut sensitivity = None; + while let Some(field) = map.next_key()? { + match field { + FieldValueField::Bytes => { + if bytes.is_some() { + return Err(A::Error::duplicate_field("bytes")); + } + bytes = Some(map.next_value_seed(FieldBytesSeed { limit: self.byte_limit })?); + } + FieldValueField::Sensitivity => { + if sensitivity.is_some() { + return Err(A::Error::duplicate_field("sensitivity")); + } + sensitivity = Some(map.next_value()?); + } + FieldValueField::Other => { + map.next_value::()?; + } + } + } + Self::finish(bytes, sensitivity) + } + } + + deserializer.deserialize_struct( + "FieldValueRepr", + FIELD_VALUE_FIELDS, + FieldValueVisitor { + byte_limit: self.byte_limit, + }, + ) + } +} + +#[derive(Clone, Copy)] +struct ListCounter { + delimiter: u8, + skip_empty: bool, + backslash_escapes: bool, + count: usize, +} + +impl ListCounter { + const fn new(delimiter: u8, skip_empty: bool, backslash_escapes: bool) -> Self { + Self { + delimiter, + skip_empty, + backslash_escapes, + count: 0, + } + } + + fn observe(&mut self, name: &'static FieldName, bytes: &[u8]) -> Result<(), crate::DecodeError> { + crate::source::update_list_item_count( + name, + bytes, + self.delimiter, + self.skip_empty, + self.backslash_escapes, + &mut self.count, + ) + } +} + +enum ListBudgets { + None, + One(ListCounter), + Two(ListCounter, ListCounter), +} + +impl ListBudgets { + fn for_name(name: Option<&'static FieldName>) -> Self { + let comma = || ListCounter::new(b',', true, true); + let semicolon = || ListCounter::new(b';', false, true); + match name { + Some( + &FieldName::CacheControl + | &FieldName::Accept + | &FieldName::AcceptEncoding + | &FieldName::AcceptLanguage + | &FieldName::Allow + | &FieldName::Vary + | &FieldName::AcceptRanges + | &FieldName::Range + | &FieldName::AccessControlAllowHeaders + | &FieldName::AccessControlAllowMethods + | &FieldName::AccessControlExposeHeaders + | &FieldName::AccessControlRequestHeaders + | &FieldName::ReferrerPolicy + | &FieldName::SecWebSocketProtocol + | &FieldName::SecWebSocketVersion, + ) => Self::One(comma()), + Some(&FieldName::ContentLength) => Self::One(ListCounter::new(b',', false, true)), + Some(&FieldName::ContentType | &FieldName::StrictTransportSecurity) => Self::One(semicolon()), + Some(&FieldName::IfMatch | &FieldName::IfNoneMatch) => Self::One(ListCounter::new(b',', true, false)), + Some(&FieldName::SecWebSocketExtensions) => Self::Two(comma(), semicolon()), + _ => Self::None, + } + } + + fn observe(&mut self, name: Option<&'static FieldName>, bytes: &[u8]) -> Result<(), crate::DecodeError> { + let Some(name) = name else { + return Ok(()); + }; + match self { + Self::None => Ok(()), + Self::One(counter) => counter.observe(name, bytes), + Self::Two(first, second) => { + first.observe(name, bytes)?; + second.observe(name, bytes) + } + } + } +} + +struct FieldValuesVisitor { + name: Option<&'static FieldName>, + byte_limit: Option, + line_limit: Option, +} + +impl<'de> Visitor<'de> for FieldValuesVisitor { + type Value = Vec; + + #[cfg_attr(coverage_nightly, coverage(off))] + fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + match self.line_limit { + Some(limit) => write!(formatter, "at most {limit} field values"), + None => formatter.write_str("a sequence of field values"), + } + } + + fn visit_seq>(self, mut seq: A) -> Result { + let capacity = seq + .size_hint() + .unwrap_or(0) + .min(self.line_limit.unwrap_or(SIZE_HINT_RESERVE_LIMIT)) + .min(SIZE_HINT_RESERVE_LIMIT); + let mut values = Vec::with_capacity(capacity); + let mut total_bytes = 0_usize; + let mut list_budgets = ListBudgets::for_name(self.name); + while self.line_limit.is_none_or(|limit| values.len() < limit) { + let remaining_bytes = self.byte_limit.map(|limit| limit - total_bytes); + let Some(value) = seq.next_element_seed(FieldValueSeed { + byte_limit: remaining_bytes, + })? + else { + return Ok(values); + }; + list_budgets.observe(self.name, value.as_bytes()).map_err(A::Error::custom)?; + total_bytes = checked_total_bytes::(total_bytes, value.as_bytes().len())?; + values.push(value); + } + if seq.next_element::()?.is_some() { + return Err(A::Error::custom(format_args!( + "{}: field-value line budget exceeded", + DecodeErrorKind::SourceLimitExceeded + ))); + } + Ok(values) + } +} + +fn checked_total_bytes(total: usize, additional: usize) -> Result { + total.checked_add(additional).ok_or_else(|| { + E::custom(format_args!( + "{}: aggregate field-value size overflow", + DecodeErrorKind::SourceLimitExceeded + )) + }) +} + +fn deserialize_field_values<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + deserializer.deserialize_seq(FieldValuesVisitor { + name: None, + byte_limit: None, + line_limit: None, + }) +} + +#[cfg(any( + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +fn deserialize_field_values_for<'de, D: Deserializer<'de>>( + name: Option<&'static FieldName>, + deserializer: D, +) -> Result, D::Error> { + deserializer.deserialize_seq(FieldValuesVisitor { + name, + byte_limit: Some(MAX_CUSTOM_FIELD_BYTES), + line_limit: Some(MAX_CUSTOM_FIELD_LINES), + }) +} + +#[derive(Clone, Copy)] +struct SerializedFieldValue<'a> { + value: FieldValueRef<'a>, + sensitivity: FieldSensitivity, +} + +impl SerializedFieldValue<'_> { + fn preserving(value: FieldValueRef<'_>) -> SerializedFieldValue<'_> { + let sensitivity = if value.is_sensitive() { + FieldSensitivity::Sensitive + } else { + FieldSensitivity::NonSensitive + }; + SerializedFieldValue { value, sensitivity } + } +} + +impl Serialize for SerializedFieldValue<'_> { + fn serialize(&self, serializer: S) -> Result { + FieldValueRepr { + bytes: self.value.as_bytes(), + sensitivity: self.sensitivity, + } + .serialize(serializer) + } +} + +fn serialize_field_values<'a, S, I>(values: I, serializer: S) -> Result +where + S: Serializer, + I: IntoIterator>, +{ + serializer.collect_seq(values.into_iter().map(SerializedFieldValue::preserving)) +} + +impl Serialize for FieldName { + fn serialize(&self, serializer: S) -> Result { + serializer.serialize_str(self.as_str()) + } +} + +impl<'de> Deserialize<'de> for FieldName { + fn deserialize>(deserializer: D) -> Result { + let value = String::deserialize(deserializer)?; + if let Ok(name) = Self::try_from_bytes(&value) { + return Ok(name); + } + #[cfg(feature = "http")] + if let Ok(name) = http::HeaderName::from_lowercase(value.as_bytes()) { + return Ok(Self::from(name)); + } + Err(D::Error::custom(crate::InvalidFieldName)) + } +} + +impl Serialize for FieldValue { + fn serialize(&self, serializer: S) -> Result { + SerializedFieldValue::preserving(self.as_field_value_ref()).serialize(serializer) + } +} + +impl<'de> Deserialize<'de> for FieldValue { + fn deserialize>(deserializer: D) -> Result { + FieldValueSeed { byte_limit: None }.deserialize(deserializer) + } +} + +impl Serialize for EncodedValues { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(self.iter().map(FieldValue::as_field_value_ref), serializer) + } +} + +impl<'de> Deserialize<'de> for EncodedValues { + fn deserialize>(deserializer: D) -> Result { + deserialize_field_values(deserializer).map(Self::from_vec) + } +} + +#[cfg(any( + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +struct SerdeSource { + values: Vec, +} + +#[cfg(any( + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +impl FieldSource for SerdeSource { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_slice(name, &self.values) + } +} + +#[cfg(any( + feature = "headers-authorization", + feature = "headers-conditional", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +fn serialize_single_owned(value: &H::Owned, serializer: S) -> Result +where + H: SingleValueField, + S: Serializer, +{ + serialize_field_values(iter::once(H::as_field_value(value).as_field_value_ref()), serializer) +} + +#[cfg(any( + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +fn deserialize_owned<'de, H, D>(deserializer: D) -> Result +where + H: Field, + D: Deserializer<'de>, +{ + let source = SerdeSource { + values: deserialize_field_values_for(Some(H::name()), deserializer)?, + }; + H::owned_with(&source, DecodeMode::Relaxed) + .map_err(D::Error::custom)? + .ok_or_else(|| D::Error::custom("an owned header must contain at least one field value")) +} + +macro_rules! serde_deserialize_owned { + ($($(#[$meta:meta])* ($header:ty, $owned:ty)),+ $(,)?) => { + $( + $(#[$meta])* + impl<'de> Deserialize<'de> for $owned { + fn deserialize>(deserializer: D) -> Result { + deserialize_owned::<$header, D>(deserializer) + } + } + )+ + }; +} + +macro_rules! serde_single_owned { + ($($(#[$meta:meta])* ($header:ty, $owned:ty)),+ $(,)?) => { + $( + $(#[$meta])* + impl Serialize for $owned { + fn serialize(&self, serializer: S) -> Result { + serialize_single_owned::<$header, S>(self, serializer) + } + } + )+ + }; +} + +macro_rules! serde_values_owned { + ($($(#[$meta:meta])* ($owned:ty, $method:ident)),+ $(,)?) => { + $( + $(#[$meta])* + impl Serialize for $owned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(self.$method(), serializer) + } + } + )+ + }; +} + +serde_single_owned!( + #[cfg(feature = "headers-negotiation")] + (Host, HostOwned), + #[cfg(feature = "headers-negotiation")] + (Server, ServerOwned), + #[cfg(feature = "headers-range")] + (ContentRange, ContentRangeOwned), + #[cfg(feature = "headers-range")] + (Range, RangeOwned), + #[cfg(feature = "headers-etag")] + (ETag, ETagOwned), + #[cfg(feature = "headers-location")] + (Location, LocationOwned), + #[cfg(feature = "headers-user-agent")] + (UserAgent, UserAgentOwned), + #[cfg(feature = "headers-conditional")] + (IfModifiedSince, IfModifiedSinceOwned), + #[cfg(feature = "headers-conditional")] + (IfUnmodifiedSince, IfUnmodifiedSinceOwned), + #[cfg(feature = "headers-conditional")] + (IfRange, IfRangeOwned), + #[cfg(feature = "headers-conditional")] + (LastModified, LastModifiedOwned), + #[cfg(feature = "headers-security")] + (StrictTransportSecurity, StrictTransportSecurityOwned), + #[cfg(feature = "headers-security")] + (XContentTypeOptions, XContentTypeOptionsOwned), + #[cfg(feature = "headers-websocket")] + (SecWebSocketAccept, SecWebSocketAcceptOwned), + #[cfg(feature = "headers-websocket")] + (SecWebSocketKey, SecWebSocketKeyOwned), + #[cfg(feature = "headers-authorization")] + (Authorization, AuthorizationOwned), + #[cfg(feature = "headers-authorization")] + (Authorization, AuthorizationOwned), +); + +serde_values_owned!( + #[cfg(feature = "headers-negotiation")] + (AcceptOwned, values), + #[cfg(feature = "headers-negotiation")] + (AcceptEncodingOwned, values), + #[cfg(feature = "headers-negotiation")] + (AcceptLanguageOwned, values), + #[cfg(feature = "headers-negotiation")] + (AllowOwned, values), + #[cfg(feature = "headers-negotiation")] + (VaryOwned, values), + #[cfg(feature = "headers-cors")] + (AccessControlAllowHeadersOwned, field_values), + #[cfg(feature = "headers-cors")] + (AccessControlAllowMethodsOwned, field_values), + #[cfg(feature = "headers-cors")] + (AccessControlExposeHeadersOwned, field_values), + #[cfg(feature = "headers-cors")] + (AccessControlRequestHeadersOwned, field_values), + #[cfg(feature = "headers-security")] + (ContentSecurityPolicyOwned, field_values), + #[cfg(feature = "headers-security")] + (ReferrerPolicyOwned, field_values), +); + +#[cfg(feature = "headers-cache-control")] +impl Serialize for CacheControlOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(self.field_values(), serializer) + } +} + +#[cfg(feature = "headers-range")] +impl Serialize for AcceptRangesOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(self.field_values(), serializer) + } +} + +#[cfg(feature = "headers-content-type")] +impl Serialize for ContentTypeOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(iter::once(self.field_value()), serializer) + } +} + +#[cfg(feature = "headers-conditional")] +impl Serialize for IfMatchOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(self.field_values(), serializer) + } +} + +#[cfg(feature = "headers-conditional")] +impl Serialize for IfNoneMatchOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(self.field_values(), serializer) + } +} + +#[cfg(feature = "headers-websocket")] +impl Serialize for SecWebSocketExtensionsOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(self.field_values(), serializer) + } +} + +#[cfg(feature = "headers-websocket")] +impl Serialize for SecWebSocketProtocolOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(self.field_values(), serializer) + } +} + +#[cfg(feature = "headers-websocket")] +impl Serialize for SecWebSocketVersionOwned { + fn serialize(&self, serializer: S) -> Result { + let value = self.field_value(); + serialize_field_values(iter::once(value.as_field_value_ref()), serializer) + } +} + +serde_deserialize_owned!( + #[cfg(feature = "headers-cache-control")] + (CacheControl, CacheControlOwned), + #[cfg(feature = "headers-negotiation")] + (Accept, AcceptOwned), + #[cfg(feature = "headers-negotiation")] + (AcceptEncoding, AcceptEncodingOwned), + #[cfg(feature = "headers-negotiation")] + (AcceptLanguage, AcceptLanguageOwned), + #[cfg(feature = "headers-negotiation")] + (Allow, AllowOwned), + #[cfg(feature = "headers-negotiation")] + (Host, HostOwned), + #[cfg(feature = "headers-negotiation")] + (Server, ServerOwned), + #[cfg(feature = "headers-negotiation")] + (Vary, VaryOwned), + #[cfg(feature = "headers-range")] + (AcceptRanges, AcceptRangesOwned), + #[cfg(feature = "headers-range")] + (ContentRange, ContentRangeOwned), + #[cfg(feature = "headers-range")] + (Range, RangeOwned), + #[cfg(feature = "headers-etag")] + (ETag, ETagOwned), + #[cfg(feature = "headers-location")] + (Location, LocationOwned), + #[cfg(feature = "headers-user-agent")] + (UserAgent, UserAgentOwned), + #[cfg(feature = "headers-content-type")] + (ContentType, ContentTypeOwned), + #[cfg(feature = "headers-conditional")] + (IfMatch, IfMatchOwned), + #[cfg(feature = "headers-conditional")] + (IfNoneMatch, IfNoneMatchOwned), + #[cfg(feature = "headers-conditional")] + (IfModifiedSince, IfModifiedSinceOwned), + #[cfg(feature = "headers-conditional")] + (IfUnmodifiedSince, IfUnmodifiedSinceOwned), + #[cfg(feature = "headers-conditional")] + (IfRange, IfRangeOwned), + #[cfg(feature = "headers-conditional")] + (LastModified, LastModifiedOwned), + #[cfg(feature = "headers-cors")] + (AccessControlAllowHeaders, AccessControlAllowHeadersOwned), + #[cfg(feature = "headers-cors")] + (AccessControlAllowMethods, AccessControlAllowMethodsOwned), + #[cfg(feature = "headers-cors")] + (AccessControlExposeHeaders, AccessControlExposeHeadersOwned), + #[cfg(feature = "headers-cors")] + (AccessControlRequestHeaders, AccessControlRequestHeadersOwned), + #[cfg(feature = "headers-security")] + (ContentSecurityPolicy, ContentSecurityPolicyOwned), + #[cfg(feature = "headers-security")] + (ReferrerPolicy, ReferrerPolicyOwned), + #[cfg(feature = "headers-security")] + (StrictTransportSecurity, StrictTransportSecurityOwned), + #[cfg(feature = "headers-security")] + (XContentTypeOptions, XContentTypeOptionsOwned), + #[cfg(feature = "headers-websocket")] + (SecWebSocketAccept, SecWebSocketAcceptOwned), + #[cfg(feature = "headers-websocket")] + (SecWebSocketExtensions, SecWebSocketExtensionsOwned), + #[cfg(feature = "headers-websocket")] + (SecWebSocketKey, SecWebSocketKeyOwned), + #[cfg(feature = "headers-websocket")] + (SecWebSocketProtocol, SecWebSocketProtocolOwned), + #[cfg(feature = "headers-websocket")] + (SecWebSocketVersion, SecWebSocketVersionOwned), + #[cfg(feature = "headers-authorization")] + (Authorization, AuthorizationOwned), + #[cfg(feature = "headers-authorization")] + (Authorization, AuthorizationOwned), +); + +#[cfg(feature = "headers-cors")] +impl Serialize for AccessControlAllowCredentialsOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(iter::once(FieldValueRef::new(b"true")), serializer) + } +} + +#[cfg(feature = "headers-cors")] +impl Serialize for AccessControlMaxAgeOwned { + fn serialize(&self, serializer: S) -> Result { + let value = FieldValue::from(self.seconds()); + serialize_field_values(iter::once(value.as_field_value_ref()), serializer) + } +} + +#[cfg(feature = "headers-cors")] +impl Serialize for AccessControlAllowOriginOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(iter::once(self.field_value().as_field_value_ref()), serializer) + } +} + +#[cfg(feature = "headers-cors")] +impl Serialize for AccessControlRequestMethodOwned { + fn serialize(&self, serializer: S) -> Result { + serialize_field_values(iter::once(self.field_value()), serializer) + } +} + +#[cfg(feature = "headers-content-length")] +impl Serialize for ContentLengthOwned { + fn serialize(&self, serializer: S) -> Result { + let value = FieldValue::from(self.get()); + serialize_field_values(iter::once(value.as_field_value_ref()), serializer) + } +} + +serde_deserialize_owned!( + #[cfg(feature = "headers-cors")] + (AccessControlAllowCredentials, AccessControlAllowCredentialsOwned), + #[cfg(feature = "headers-cors")] + (AccessControlMaxAge, AccessControlMaxAgeOwned), + #[cfg(feature = "headers-cors")] + (AccessControlAllowOrigin, AccessControlAllowOriginOwned), + #[cfg(feature = "headers-cors")] + (AccessControlRequestMethod, AccessControlRequestMethodOwned), + #[cfg(feature = "headers-content-length")] + (ContentLength, ContentLengthOwned), +); + +#[cfg(feature = "headers-set-cookie")] +impl Serialize for SetCookieOwned { + fn serialize(&self, serializer: S) -> Result { + serializer.collect_seq(self.iter().map(|value| SerializedFieldValue { + value: value.as_field_value_ref(), + sensitivity: FieldSensitivity::Sensitive, + })) + } +} + +#[cfg(feature = "headers-set-cookie")] +impl<'de> Deserialize<'de> for SetCookieOwned { + fn deserialize>(deserializer: D) -> Result { + let source = SerdeSource { + values: deserialize_field_values_for(Some(SetCookie::name()), deserializer)?, + }; + if source.values.is_empty() { + return Ok(Self::new()); + } + SetCookie::owned(&source) + .map_err(D::Error::custom)? + .ok_or_else(|| D::Error::custom("nonempty Set-Cookie values must decode as present")) + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use serde::de::value::Error; + + use super::checked_total_bytes; + + #[test] + fn aggregate_byte_count_preserves_representable_totals() { + for (total, additional, expected) in [ + (0, 0, 0), + (12, 34, 46), + (0, usize::MAX, usize::MAX), + (usize::MAX - 1, 1, usize::MAX), + ] { + assert_eq!(checked_total_bytes::(total, additional).unwrap(), expected); + } + assert_eq!(checked_total_bytes::(usize::MAX, 0).unwrap(), usize::MAX); + } + + #[test] + fn aggregate_byte_count_overflow_reports_source_limit_exceeded() { + for (total, additional) in [(usize::MAX, 1), (1, usize::MAX), (usize::MAX, usize::MAX)] { + assert_eq!( + checked_total_bytes::(total, additional).unwrap_err().to_string(), + "source limit exceeded: aggregate field-value size overflow" + ); + } + } +} diff --git a/crates/http_headers/src/sink/encoded_values.rs b/crates/http_headers/src/sink/encoded_values.rs new file mode 100644 index 000000000..1997f9605 --- /dev/null +++ b/crates/http_headers/src/sink/encoded_values.rs @@ -0,0 +1,479 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Owned encoded field values. + +use std::cmp::Ordering; +use std::collections::TryReserveError; +use std::hash::{Hash, Hasher}; +use std::{fmt, option, slice, vec}; + +use crate::FieldValue; + +/// An iterator over borrowed encoded field values. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::FieldValue; +/// use http_headers::sink::{EncodedValues, EncodedValuesIter}; +/// +/// let values = EncodedValues::from_vec(vec![ +/// FieldValue::from_static("gzip"), +/// FieldValue::from_static("br"), +/// ]); +/// let mut iter: EncodedValuesIter<'_> = values.iter(); +/// assert_eq!(iter.len(), 2); +/// assert_eq!(iter.next().expect("first value"), "gzip"); +/// assert_eq!(iter.next().expect("second value"), "br"); +/// assert!(iter.next().is_none()); +/// ``` +pub struct EncodedValuesIter<'a> { + first: option::Iter<'a, FieldValue>, + rest: slice::Iter<'a, FieldValue>, +} + +/// An iterator over mutably borrowed encoded field values. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::sink::{EncodedValues, EncodedValuesIterMut}; +/// use http_headers::{FieldSensitivity, FieldValue}; +/// +/// let mut values = EncodedValues::from_vec(vec![ +/// FieldValue::from_static("gzip"), +/// FieldValue::from_static("br"), +/// ]); +/// let mut iter: EncodedValuesIterMut<'_> = values.iter_mut(); +/// assert_eq!(iter.len(), 2); +/// iter.next() +/// .expect("first mutable value") +/// .set_sensitivity(FieldSensitivity::Sensitive); +/// drop(iter); +/// assert!(values.iter().next().expect("first value").is_sensitive()); +/// ``` +pub struct EncodedValuesIterMut<'a> { + first: option::IterMut<'a, FieldValue>, + rest: slice::IterMut<'a, FieldValue>, +} + +/// An owning iterator over encoded field values. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::FieldValue; +/// use http_headers::sink::{EncodedValues, EncodedValuesIntoIter}; +/// +/// let values = EncodedValues::from_vec(vec![ +/// FieldValue::from_static("gzip"), +/// FieldValue::from_static("br"), +/// ]); +/// let mut iter: EncodedValuesIntoIter = values.into_iter(); +/// assert_eq!(iter.len(), 2); +/// assert_eq!(iter.next().expect("first value"), "gzip"); +/// assert_eq!(iter.next().expect("second value"), "br"); +/// assert!(iter.next().is_none()); +/// ``` +pub struct EncodedValuesIntoIter { + first: option::IntoIter, + rest: vec::IntoIter, +} + +/// Owned field values ready to be stored by a [`crate::sink::FieldSink`]. +/// +/// The collection holds values only; the field name is supplied separately by +/// [`FieldSink::set_values`]. Each value becomes one field line once stored. +/// +/// [`FieldSink::set_values`]: crate::sink::FieldSink::set_values +/// +/// [`Debug`](std::fmt::Debug) reports only the number of values and never +/// exposes their contents. +/// +/// # Examples +/// +/// ```rust +/// let values = +/// http_headers::sink::EncodedValues::single(http_headers::FieldValue::from_static("gzip")); +/// assert_eq!(values.len(), 1); +/// ``` +#[derive(Clone, Default)] +pub struct EncodedValues { + first: Option, + rest: Vec, +} + +impl fmt::Debug for EncodedValues { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("EncodedValues").field("value_count", &self.len()).finish() + } +} + +impl PartialEq for EncodedValues { + fn eq(&self, other: &Self) -> bool { + self.iter().eq(other.iter()) + } +} + +impl Eq for EncodedValues {} + +impl PartialOrd for EncodedValues { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } +} + +impl Ord for EncodedValues { + fn cmp(&self, other: &Self) -> Ordering { + self.iter().cmp(other.iter()) + } +} + +impl Hash for EncodedValues { + fn hash(&self, state: &mut H) { + self.len().hash(state); + self.iter().for_each(|value| value.hash(state)); + } +} + +impl EncodedValues { + /// Creates an empty collection. + /// + /// # Examples + /// + /// ```rust + /// assert!(http_headers::sink::EncodedValues::new().is_empty()); + /// ``` + #[must_use] + pub const fn new() -> Self { + Self { + first: None, + rest: Vec::new(), + } + } + + /// Creates a collection containing one field value. + /// + /// # Examples + /// + /// ```rust + /// let values = + /// http_headers::sink::EncodedValues::single(http_headers::FieldValue::from_static("gzip")); + /// assert_eq!(values.len(), 1); + /// ``` + #[must_use] + pub fn single(value: FieldValue) -> Self { + Self { + first: Some(value), + rest: Vec::new(), + } + } + + /// Creates a collection from field values in their existing order. + /// + /// # Examples + /// + /// ```rust + /// let values = + /// http_headers::sink::EncodedValues::from_vec(vec![http_headers::FieldValue::from_static( + /// "gzip", + /// )]); + /// assert_eq!(values.len(), 1); + /// ``` + #[must_use] + pub fn from_vec(values: Vec) -> Self { + Self { first: None, rest: values } + } + + /// Appends a field value. + /// + /// # Examples + /// + /// ```rust + /// let mut values = http_headers::sink::EncodedValues::new(); + /// values.push(http_headers::FieldValue::from_static("gzip")); + /// assert_eq!(values.len(), 1); + /// ``` + pub fn push(&mut self, value: FieldValue) { + if self.first.is_some() || !self.rest.is_empty() { + self.rest.push(value); + } else { + self.first = Some(value); + } + } + + pub(crate) fn try_reserve(&mut self, additional: usize) -> Result<(), TryReserveError> { + let rest_additional = additional.saturating_sub(usize::from(self.is_empty())); + self.rest.try_reserve(rest_additional) + } + + /// Returns the number of field values. + /// + /// # Examples + /// + /// ```rust + /// let values = + /// http_headers::sink::EncodedValues::single(http_headers::FieldValue::from_static("gzip")); + /// assert_eq!(values.len(), 1); + /// ``` + #[must_use] + pub fn len(&self) -> usize { + usize::from(self.first.is_some()) + self.rest.len() + } + + /// Returns whether no field values were encoded. + /// + /// # Examples + /// + /// ```rust + /// assert!(http_headers::sink::EncodedValues::new().is_empty()); + /// ``` + #[must_use] + pub fn is_empty(&self) -> bool { + self.first.is_none() && self.rest.is_empty() + } + + /// Iterates the encoded field values. + /// + /// # Examples + /// + /// ```rust + /// let values = + /// http_headers::sink::EncodedValues::single(http_headers::FieldValue::from_static("gzip")); + /// assert!(values.iter().any(|value| value == "gzip")); + /// ``` + pub fn iter(&self) -> EncodedValuesIter<'_> { + EncodedValuesIter { + first: self.first.iter(), + rest: self.rest.iter(), + } + } + + /// Mutably iterates the encoded field values. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldSensitivity; + /// + /// let mut values = + /// http_headers::sink::EncodedValues::single(http_headers::FieldValue::from_static("gzip")); + /// values + /// .iter_mut() + /// .for_each(|value| value.set_sensitivity(FieldSensitivity::Sensitive)); + /// assert!(values.iter().all(|value| value.is_sensitive())); + /// ``` + pub fn iter_mut(&mut self) -> EncodedValuesIterMut<'_> { + EncodedValuesIterMut { + first: self.first.iter_mut(), + rest: self.rest.iter_mut(), + } + } +} + +macro_rules! impl_encoded_iterator { + ($type:ident $(<$lifetime:lifetime>)?, $item:ty) => { + impl$(<$lifetime>)? Iterator for $type$(<$lifetime>)? { + type Item = $item; + + fn next(&mut self) -> Option { + self.first.next().or_else(|| self.rest.next()) + } + + fn size_hint(&self) -> (usize, Option) { + let length = self.first.len() + self.rest.len(); + (length, Some(length)) + } + } + + impl$(<$lifetime>)? DoubleEndedIterator for $type$(<$lifetime>)? { + fn next_back(&mut self) -> Option { + self.rest.next_back().or_else(|| self.first.next_back()) + } + } + + impl$(<$lifetime>)? ExactSizeIterator for $type$(<$lifetime>)? {} + impl$(<$lifetime>)? std::iter::FusedIterator for $type$(<$lifetime>)? {} + }; +} + +impl_encoded_iterator!(EncodedValuesIter<'a>, &'a FieldValue); +impl_encoded_iterator!(EncodedValuesIterMut<'a>, &'a mut FieldValue); +impl_encoded_iterator!(EncodedValuesIntoIter, FieldValue); + +macro_rules! impl_redacted_debug { + ($type:ident $(<$lifetime:lifetime>)?) => { + impl$(<$lifetime>)? fmt::Debug for $type$(<$lifetime>)? { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct(stringify!($type)) + .field("remaining", &self.len()) + .finish() + } + } + }; +} + +impl_redacted_debug!(EncodedValuesIter<'a>); +impl_redacted_debug!(EncodedValuesIterMut<'a>); +impl_redacted_debug!(EncodedValuesIntoIter); + +impl Extend for EncodedValues { + fn extend(&mut self, iter: T) + where + T: IntoIterator, + { + let mut iter = iter.into_iter(); + if self.is_empty() { + self.first = iter.next(); + } + self.rest.reserve(iter.size_hint().0); + self.rest.extend(iter); + } +} + +impl FromIterator for EncodedValues { + fn from_iter>(iter: T) -> Self { + let mut values = Self::new(); + values.extend(iter); + values + } +} + +impl From for EncodedValues { + fn from(value: FieldValue) -> Self { + Self::single(value) + } +} + +impl IntoIterator for EncodedValues { + type Item = FieldValue; + type IntoIter = EncodedValuesIntoIter; + + fn into_iter(self) -> Self::IntoIter { + EncodedValuesIntoIter { + first: self.first.into_iter(), + rest: self.rest.into_iter(), + } + } +} + +impl<'a> IntoIterator for &'a EncodedValues { + type Item = &'a FieldValue; + type IntoIter = EncodedValuesIter<'a>; + + fn into_iter(self) -> Self::IntoIter { + self.iter() + } +} + +impl<'a> IntoIterator for &'a mut EncodedValues { + type Item = &'a mut FieldValue; + type IntoIter = EncodedValuesIterMut<'a>; + + fn into_iter(self) -> Self::IntoIter { + self.iter_mut() + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::cell::Cell; + + use super::EncodedValues; + use crate::FieldValue; + + struct HintingValues<'a> { + remaining: usize, + hint_calls: &'a Cell, + } + + impl Iterator for HintingValues<'_> { + type Item = FieldValue; + + fn next(&mut self) -> Option { + if self.remaining == 0 { + return None; + } + self.remaining -= 1; + Some(FieldValue::from_static("reserved")) + } + + fn size_hint(&self) -> (usize, Option) { + self.hint_calls.set(self.hint_calls.get() + 1); + (self.remaining, Some(self.remaining)) + } + } + + #[test] + fn collection_and_iterators_preserve_values_and_redact_debug_output() { + let mut values = EncodedValues::from_vec(vec![FieldValue::from_static("first"), FieldValue::from_static("second")]); + assert_eq!(format!("{values:?}"), "EncodedValues { value_count: 2 }"); + + let mut iter = values.iter(); + assert_eq!(iter.size_hint(), (2, Some(2))); + assert_eq!(format!("{iter:?}"), "EncodedValuesIter { remaining: 2 }"); + assert_eq!(iter.next().expect("first value"), "first"); + assert_eq!(iter.next().expect("second value"), "second"); + assert!(iter.next().is_none()); + + let mut iter = values.iter_mut(); + assert_eq!(iter.size_hint(), (2, Some(2))); + assert_eq!(format!("{iter:?}"), "EncodedValuesIterMut { remaining: 2 }"); + iter.next().expect("first mutable value").set_sensitive(true); + assert!(values.iter().next().expect("first value").is_sensitive()); + + let mut iter = values.into_iter(); + assert_eq!(iter.size_hint(), (2, Some(2))); + assert_eq!(format!("{iter:?}"), "EncodedValuesIntoIter { remaining: 2 }"); + assert_eq!(iter.next().expect("first owned value"), "first"); + assert_eq!(iter.next().expect("second owned value"), "second"); + assert!(iter.next().is_none()); + } + + #[test] + fn construction_extension_and_borrowed_iteration_cover_all_storage_shapes() { + let mut values = EncodedValues::new(); + assert!(values.is_empty()); + values.push(FieldValue::from_static("first")); + assert!(!values.is_empty()); + values.extend([FieldValue::from_static("second"), FieldValue::from_static("third")]); + assert_eq!(values.len(), 3); + assert_eq!( + (&values).into_iter().map(FieldValue::as_bytes).collect::>(), + [b"first".as_slice(), b"second".as_slice(), b"third".as_slice()] + ); + + let hint_calls = Cell::new(0); + let reserved = HintingValues { + remaining: 16, + hint_calls: &hint_calls, + } + .collect::(); + assert_eq!(reserved.len(), 16); + assert_ne!(hint_calls.get(), 0); + + for value in &mut values { + value.set_sensitive(true); + } + assert!(values.iter().all(FieldValue::is_sensitive)); + + let one = EncodedValues::from(FieldValue::from_static("one")); + assert_eq!(one.len(), 1); + let collected: EncodedValues = [FieldValue::from_static("a"), FieldValue::from_static("b")].into_iter().collect(); + assert_eq!(collected.len(), 2); + } + + #[test] + fn fallible_growth_reports_failure_without_changing_values() { + let mut values = EncodedValues::single(FieldValue::from_static("original")); + assert!(values.try_reserve(usize::MAX).is_err()); + assert_eq!(values.len(), 1); + assert_eq!(values.iter().next().unwrap(), "original"); + + values.try_reserve(1).unwrap(); + values.push(FieldValue::from_static("second")); + assert_eq!(values.len(), 2); + } +} diff --git a/crates/http_headers/src/sink/field_encoder.rs b/crates/http_headers/src/sink/field_encoder.rs new file mode 100644 index 000000000..67584cb1b --- /dev/null +++ b/crates/http_headers/src/sink/field_encoder.rs @@ -0,0 +1,319 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Encoding a typed field into the field lines a sink stores. + +use std::marker::PhantomData; + +use crate::sink::{EncodedValues, InsertError}; +use crate::{FieldValue, FieldValueRef}; + +/// Describes whether an encoded field value contains sensitive data. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::FieldSensitivity; +/// +/// assert!(!FieldSensitivity::NonSensitive.is_sensitive()); +/// assert!(FieldSensitivity::Sensitive.is_sensitive()); +/// ``` +#[cfg_attr(feature = "serde", derive(serde::Deserialize, serde::Serialize))] +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub enum FieldSensitivity { + /// The value may be logged or displayed normally. + NonSensitive, + /// The value should be redacted by supporting containers and formatters. + Sensitive, +} + +impl FieldSensitivity { + pub(crate) const fn from_sensitive(sensitive: bool) -> Self { + if sensitive { Self::Sensitive } else { Self::NonSensitive } + } + + /// Returns whether the value is sensitive. + #[must_use] + pub const fn is_sensitive(self) -> bool { + matches!(self, Self::Sensitive) + } +} + +/// Writes one encoded HTTP field value into a sink-provided destination. +/// +/// The [`FieldEncodeOutput`] example shows a writer implementation that +/// validates the announced length and preserves sensitivity. +pub trait FieldValueWriter { + /// Appends bytes to the current field value. + /// + /// # Errors + /// + /// Returns an error when the destination cannot accept the bytes. + fn write_bytes(&mut self, bytes: &[u8]) -> Result<(), InsertError>; + + /// Completes this field value. + /// + /// # Errors + /// + /// Returns an error when the completed bytes are not a valid field value + /// or the destination cannot retain them. + fn finish(self) -> Result<(), InsertError>; +} + +/// Receives the field values emitted by a [`FieldEncoder`], one per field line. +/// +/// Implement this trait when a [`super::FieldSink`] needs to write encoded values +/// directly into its own container type. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::sink::{ +/// FieldEncodeOutput, FieldEncoder, FieldValueWriter, InsertError, InsertErrorKind, +/// U64Encoder, ValueRefsEncoder, +/// }; +/// use http_headers::{FieldSensitivity, FieldValue, FieldValueRef}; +/// +/// #[derive(Default)] +/// struct Output(Vec); +/// +/// struct Writer<'a> { +/// output: &'a mut Vec, +/// bytes: Vec, +/// expected: usize, +/// sensitivity: FieldSensitivity, +/// } +/// +/// impl FieldValueWriter for Writer<'_> { +/// fn write_bytes(&mut self, bytes: &[u8]) -> Result<(), InsertError> { +/// if self.bytes.len().saturating_add(bytes.len()) > self.expected { +/// return Err(InsertError::new(InsertErrorKind::InvalidEncoding)); +/// } +/// self.bytes.extend_from_slice(bytes); +/// Ok(()) +/// } +/// +/// fn finish(self) -> Result<(), InsertError> { +/// if self.bytes.len() != self.expected { +/// return Err(InsertError::new(InsertErrorKind::InvalidEncoding)); +/// } +/// let value = FieldValue::from_bytes(self.bytes) +/// .map_err(|_| InsertError::new(InsertErrorKind::InvalidValue))? +/// .with_sensitivity(self.sensitivity); +/// self.output.push(value); +/// Ok(()) +/// } +/// } +/// +/// impl FieldEncodeOutput for Output { +/// type Writer<'a> = Writer<'a>; +/// +/// fn begin_value( +/// &mut self, +/// length: usize, +/// sensitivity: FieldSensitivity, +/// ) -> Result, InsertError> { +/// let mut bytes = Vec::new(); +/// bytes +/// .try_reserve_exact(length) +/// .map_err(|_| InsertError::new(InsertErrorKind::AllocationFailed))?; +/// Ok(Writer { +/// output: &mut self.0, +/// bytes, +/// expected: length, +/// sensitivity, +/// }) +/// } +/// +/// fn push_value(&mut self, value: FieldValue) -> Result<(), InsertError> { +/// self.0.push(value); +/// Ok(()) +/// } +/// } +/// +/// let mut output = Output::default(); +/// U64Encoder::new(42).encode(&mut output)?; +/// ValueRefsEncoder::new([FieldValueRef::new(b"a=1"), FieldValueRef::new(b"b=2")]) +/// .with_sensitivity(FieldSensitivity::Sensitive) +/// .encode(&mut output)?; +/// +/// assert_eq!(output.0[0].as_bytes(), b"42"); +/// assert!(output.0[1..].iter().all(FieldValue::is_sensitive)); +/// # Ok::<(), InsertError>(()) +/// ``` +pub trait FieldEncodeOutput { + /// Writer used for one newly encoded field value. + type Writer<'a>: FieldValueWriter + where + Self: 'a; + + /// Starts one field value with its exact encoded length and sensitivity. + /// + /// The encoder must write exactly `length` bytes before calling + /// [`FieldValueWriter::finish`]. Implementations must reject a completed + /// value with a different length and preserve the sensitivity marker. + /// + /// # Errors + /// + /// Returns an error when the destination cannot reserve the requested + /// storage. + fn begin_value(&mut self, length: usize, sensitivity: FieldSensitivity) -> Result, InsertError>; + + /// Receives an already-owned field value. + /// + /// # Errors + /// + /// Returns an error when the destination cannot retain the value. + fn push_value(&mut self, value: FieldValue) -> Result<(), InsertError>; + + /// Writes one unsigned decimal field value. + /// + /// # Errors + /// + /// Returns an error when the destination cannot retain the value. + fn push_u64(&mut self, mut value: u64) -> Result<(), InsertError> { + let mut storage = [0_u8; 20]; + let mut start = storage.len(); + loop { + start -= 1; + storage[start] = b'0' + (value % 10) as u8; + value /= 10; + if value == 0 { + break; + } + } + let digits = &storage[start..]; + let mut writer = self.begin_value(digits.len(), FieldSensitivity::NonSensitive)?; + writer.write_bytes(digits)?; + writer.finish() + } +} + +/// A value that can encode one field's lines. +/// +/// The [`FieldEncodeOutput`] example implements this trait's destination and +/// drives it with reusable encoders. +pub trait FieldEncoder { + /// Emits every field value into `output`, one per field line. + /// + /// # Errors + /// + /// Returns an error when the destination cannot encode or retain the + /// complete field. + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput; +} + +impl FieldEncoder for FieldValueRef<'_> { + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + let mut writer = output.begin_value(self.as_bytes().len(), FieldSensitivity::from_sensitive(self.is_sensitive()))?; + writer.write_bytes(self.as_bytes())?; + writer.finish() + } +} + +/// Encodes borrowed values in iteration order, one field line each. +/// +/// The [`FieldEncodeOutput`] example encodes multiple borrowed values and +/// applies one sensitivity marker to every emitted line. +#[derive(Clone, Debug)] +pub struct ValueRefsEncoder<'a, I> { + values: I, + sensitivity: Option, + marker: PhantomData>, +} + +impl ValueRefsEncoder<'_, I> { + /// Creates an encoder that preserves each borrowed value's sensitivity. + #[must_use] + pub const fn new(values: I) -> Self { + Self { + values, + sensitivity: None, + marker: PhantomData, + } + } + + /// Sets the sensitivity marker on every emitted field line. + /// + /// This replaces the marker carried by each individual + /// [`FieldValueRef`]. + #[must_use] + pub const fn with_sensitivity(mut self, sensitivity: FieldSensitivity) -> Self { + self.sensitivity = Some(sensitivity); + self + } + + #[cfg(any(test, feature = "headers-set-cookie",))] + pub(crate) const fn with_sensitive(mut self, sensitive: bool) -> Self { + self.sensitivity = Some(FieldSensitivity::from_sensitive(sensitive)); + self + } +} + +impl<'a, I> FieldEncoder for ValueRefsEncoder<'a, I> +where + I: IntoIterator>, +{ + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + for value in self.values { + match self.sensitivity { + Some(sensitivity) => value.with_sensitivity(sensitivity).encode(output)?, + None => value.encode(output)?, + } + } + Ok(()) + } +} + +impl FieldEncoder for FieldValue { + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + output.push_value(self) + } +} + +impl FieldEncoder for EncodedValues { + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + for value in self { + output.push_value(value)?; + } + Ok(()) + } +} + +/// Encodes one unsigned integer as a decimal field value. +/// +/// The [`FieldEncodeOutput`] example encodes an integer and inspects its +/// decimal bytes. +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct U64Encoder(u64); + +impl U64Encoder { + /// Creates an unsigned decimal encoder. + #[must_use] + pub const fn new(value: u64) -> Self { + Self(value) + } +} + +impl FieldEncoder for U64Encoder { + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + output.push_u64(self.0) + } +} diff --git a/crates/http_headers/src/sink/field_sink.rs b/crates/http_headers/src/sink/field_sink.rs new file mode 100644 index 000000000..b1c660020 --- /dev/null +++ b/crates/http_headers/src/sink/field_sink.rs @@ -0,0 +1,613 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! The trait a field container implements to store field lines. + +use smallvec::SmallVec; + +use crate::sink::{EncodedValues, FieldEncodeOutput, FieldEncoder, FieldSensitivity, FieldValueWriter, InsertError, InsertErrorKind}; +use crate::source::FieldSource; +use crate::{FieldName, FieldValue}; + +// Match FieldValue's inline representation to avoid temporary heap storage. +const FIELD_VALUE_INLINE_CAPACITY: usize = 64; +type CollectBytes = SmallVec<[u8; FIELD_VALUE_INLINE_CAPACITY]>; + +struct CollectOutput { + values: EncodedValues, +} + +struct CollectWriter<'a> { + values: &'a mut EncodedValues, + bytes: WriterBytes, + expected: usize, + sensitive: bool, +} + +enum WriterBytes { + Inline(CollectBytes), + Spilled(Vec), +} + +impl FieldEncodeOutput for CollectOutput { + type Writer<'a> = CollectWriter<'a>; + + #[expect( + clippy::inline_always, + reason = "keeps the inline-buffer dispatch within the measured instruction gate" + )] + #[inline(always)] + fn begin_value(&mut self, length: usize, sensitivity: FieldSensitivity) -> Result, InsertError> { + if length > isize::MAX as usize { + return Err(InsertError::new(InsertErrorKind::CapacityExceeded)); + } + let bytes = if length <= FIELD_VALUE_INLINE_CAPACITY { + WriterBytes::Inline(CollectBytes::new()) + } else { + WriterBytes::Spilled(reserve_bytes(length)?) + }; + Ok(CollectWriter { + values: &mut self.values, + bytes, + expected: length, + sensitive: sensitivity.is_sensitive(), + }) + } + + fn push_value(&mut self, value: FieldValue) -> Result<(), InsertError> { + try_push_collected_value(&mut self.values, value) + } +} + +fn reserve_collected_values(values: &mut EncodedValues, additional: usize) -> Result<(), InsertError> { + values + .try_reserve(additional) + .map_err(|_error| InsertError::new(InsertErrorKind::AllocationFailed)) +} + +fn try_push_collected_value(values: &mut EncodedValues, value: FieldValue) -> Result<(), InsertError> { + reserve_collected_values(values, 1)?; + values.push(value); + Ok(()) +} + +fn reserve_bytes(length: usize) -> Result, InsertError> { + let mut bytes = Vec::new(); + bytes + .try_reserve_exact(length) + .map_err(|_error| InsertError::new(InsertErrorKind::AllocationFailed))?; + Ok(bytes) +} + +impl FieldValueWriter for CollectWriter<'_> { + #[expect( + clippy::inline_always, + reason = "keeps the inline-buffer dispatch within the measured instruction gate" + )] + #[inline(always)] + fn write_bytes(&mut self, bytes: &[u8]) -> Result<(), InsertError> { + let written = match &self.bytes { + WriterBytes::Inline(buffer) => buffer.len(), + WriterBytes::Spilled(buffer) => buffer.len(), + }; + if bytes.len() > self.expected.saturating_sub(written) { + return Err(InsertError::new(InsertErrorKind::InvalidEncoding)); + } + match &mut self.bytes { + WriterBytes::Inline(buffer) => buffer.extend_from_slice(bytes), + WriterBytes::Spilled(buffer) => buffer.extend_from_slice(bytes), + } + Ok(()) + } + + #[expect( + clippy::inline_always, + reason = "keeps the inline-buffer dispatch within the measured instruction gate" + )] + #[inline(always)] + fn finish(self) -> Result<(), InsertError> { + let length = match &self.bytes { + WriterBytes::Inline(bytes) => bytes.len(), + WriterBytes::Spilled(bytes) => bytes.len(), + }; + if length != self.expected { + return Err(InsertError::new(InsertErrorKind::InvalidEncoding)); + } + let value = match self.bytes { + WriterBytes::Inline(bytes) => FieldValue::from_bytes(bytes), + WriterBytes::Spilled(bytes) => FieldValue::try_from(bytes), + } + .map_err(|_invalid| InsertError::new(InsertErrorKind::InvalidValue))? + .with_sensitive(self.sensitive); + try_push_collected_value(self.values, value) + } +} + +/// A container that can store the field lines of a field. +/// +/// With a header-family feature enabled, most callers create headers with +/// `crate::sink::FieldSinkExt` or [`crate::Field::insert`]. Core-only builds +/// can use this trait directly. Implement it only when integrating a custom +/// field container. +/// +/// All name-taking methods require static descriptors, just like +/// [`FieldSource`]. Custom descriptors can use a `static LazyLock`. +/// Use the container's native API to insert, append, or remove fields whose +/// names are constructed locally at runtime. +/// +/// # Required and provided methods +/// +/// [`set_values`], [`append_values`], and [`remove_values`] are required. +/// Separating replacement from append lets inherited encoded appends buffer +/// only their new field lines before one atomic storage operation. +/// +/// The other two methods take a [`FieldEncoder`] rather than finished values, +/// which lets a typed field write its bytes straight into whatever storage the +/// container prefers. There is deliberately no `remove_encoded`, because +/// removal has nothing to encode. +/// +/// [`set_values`]: FieldSink::set_values +/// [`append_values`]: FieldSink::append_values +/// [`remove_values`]: FieldSink::remove_values +/// [`append_encoded`]: FieldSink::append_encoded +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # fn main() -> Result<(), http_headers::sink::InsertError> { +/// use http::HeaderMap; +/// use http_headers::sink::{EncodedValues, FieldSink}; +/// use http_headers::source::FieldSource; +/// use http_headers::{FieldName, FieldValue}; +/// +/// let mut map = HeaderMap::new(); +/// let values = EncodedValues::single(FieldValue::from_static("client/1")); +/// FieldSink::set_values(&mut map, &FieldName::UserAgent, values)?; +/// let stored = FieldSource::lines(&map, &FieldName::UserAgent).expect("value was stored"); +/// assert_eq!( +/// stored.exactly_one().expect("exactly one value").as_bytes(), +/// b"client/1" +/// ); +/// FieldSink::remove_values(&mut map, &FieldName::UserAgent); +/// assert!(!FieldSource::contains(&map, &FieldName::UserAgent)); +/// # Ok(()) +/// # } +/// # #[cfg(not(feature = "http"))] +/// # fn main() {} +/// ``` +pub trait FieldSink: FieldSource { + /// Replaces every field line stored under `name` from an encoding plan. + /// + /// # Errors + /// + /// Returns an error without changing the sink when encoding or storage + /// fails. + fn set_encoded(&mut self, name: &'static FieldName, encoder: E) -> Result<(), InsertError> + where + E: FieldEncoder, + Self: Sized, + { + let mut output = CollectOutput { + values: EncodedValues::new(), + }; + encoder.encode(&mut output)?; + self.set_values(name, output.values) + } + + /// Appends field lines produced by an encoding plan. + /// + /// # Errors + /// + /// Returns an error without changing the sink when encoding or storage + /// fails. + fn append_encoded(&mut self, name: &'static FieldName, encoder: E) -> Result<(), InsertError> + where + E: FieldEncoder, + Self: Sized, + { + let mut output = CollectOutput { + values: EncodedValues::new(), + }; + encoder.encode(&mut output)?; + self.append_values(name, output.values) + } + + /// Replaces every field line stored under `name`. + /// + /// An empty `values` removes the field. + /// + /// # Errors + /// + /// Returns an error, leaving the container unchanged, when it cannot hold + /// the values. + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError>; + + /// Appends finished field lines under `name`. + /// + /// # Errors + /// + /// Returns an error without changing the sink when storage cannot retain + /// the appended values. + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError>; + + /// Removes every field line stored under `name`. + fn remove_values(&mut self, name: &'static FieldName); +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #[cfg(feature = "http")] + use http::{HeaderMap, HeaderValue, header}; + + use super::{CollectOutput, FIELD_VALUE_INLINE_CAPACITY, FieldSink, reserve_bytes, reserve_collected_values}; + use crate::sink::{ + EncodedValues, FieldEncodeOutput, FieldEncoder, FieldSensitivity, FieldValueWriter, InsertError, InsertErrorKind, U64Encoder, + ValueRefsEncoder, + }; + use crate::source::{FieldLines, FieldSource}; + use crate::{FieldName, FieldValue, FieldValueRef}; + + type BorrowedFieldValue<'a> = FieldValueRef<'a>; + + #[derive(Default)] + struct Sink { + values: Vec, + append_calls: usize, + } + + impl FieldSource for Sink { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_slice(name, &self.values) + } + } + + impl FieldSink for Sink { + fn set_values(&mut self, _name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.values = values.into_iter().collect(); + Ok(()) + } + + fn append_values(&mut self, _name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.append_calls += 1; + self.values.extend(values); + Ok(()) + } + + fn remove_values(&mut self, _name: &'static FieldName) { + self.values.clear(); + } + } + + struct BytesEncoder { + expected: usize, + bytes: &'static [u8], + sensitive: bool, + } + + impl FieldEncoder for BytesEncoder { + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + let mut writer = output.begin_value(self.expected, crate::sink::FieldSensitivity::from_sensitive(self.sensitive))?; + writer.write_bytes(self.bytes)?; + writer.finish() + } + } + + struct RejectBegin; + + struct RejectWrite; + + struct RejectFinish; + + struct RejectOwned; + + struct Writer { + reject_write: bool, + reject_finish: bool, + } + + impl FieldValueWriter for Writer { + fn write_bytes(&mut self, _bytes: &[u8]) -> Result<(), InsertError> { + if self.reject_write { + Err(InsertError::new(InsertErrorKind::InvalidEncoding)) + } else { + Ok(()) + } + } + + fn finish(self) -> Result<(), InsertError> { + if self.reject_finish { + Err(InsertError::new(InsertErrorKind::InvalidValue)) + } else { + Ok(()) + } + } + } + + macro_rules! writer_output { + ($output:ty, $write:literal, $finish:literal) => { + impl FieldEncodeOutput for $output { + type Writer<'a> = Writer; + + fn begin_value(&mut self, _length: usize, _sensitivity: FieldSensitivity) -> Result, InsertError> { + Ok(Writer { + reject_write: $write, + reject_finish: $finish, + }) + } + + fn push_value(&mut self, _value: FieldValue) -> Result<(), InsertError> { + Ok(()) + } + } + }; + } + + impl FieldEncodeOutput for RejectBegin { + type Writer<'a> = Writer; + + fn begin_value(&mut self, _length: usize, _sensitivity: FieldSensitivity) -> Result, InsertError> { + Err(InsertError::new(InsertErrorKind::AllocationFailed)) + } + + fn push_value(&mut self, _value: FieldValue) -> Result<(), InsertError> { + Ok(()) + } + } + + writer_output!(RejectWrite, true, false); + writer_output!(RejectFinish, false, true); + + impl FieldEncodeOutput for RejectOwned { + type Writer<'a> = Writer; + + fn begin_value(&mut self, _length: usize, _sensitivity: FieldSensitivity) -> Result, InsertError> { + Ok(Writer { + reject_write: false, + reject_finish: false, + }) + } + + fn push_value(&mut self, _value: FieldValue) -> Result<(), InsertError> { + Err(InsertError::new(InsertErrorKind::CapacityExceeded)) + } + } + + #[test] + fn default_sink_encoders_cover_owned_borrowed_numeric_and_append_paths() { + let mut sink = Sink::default(); + assert!(!FieldSource::contains(&sink, &FieldName::ContentLength)); + assert!(!FieldSource::contains(&&sink, &FieldName::ContentLength)); + + sink.set_encoded(&FieldName::ContentLength, U64Encoder::new(0)) + .expect("zero encodes"); + assert_eq!(sink.values[0].as_bytes(), b"0"); + + sink.set_encoded(&FieldName::ContentLength, U64Encoder::new(u64::MAX)) + .expect("maximum encodes"); + assert_eq!(sink.values[0].as_bytes(), u64::MAX.to_string().as_bytes()); + + sink.set_encoded( + &FieldName::Authorization, + BorrowedFieldValue::new(b"Bearer token").with_sensitive(true), + ) + .expect("borrowed value encodes"); + assert!(sink.values[0].is_sensitive()); + + let borrowed = [FieldValueRef::new(b"a=1"), FieldValueRef::new(b"b=2")]; + sink.set_encoded(&FieldName::SetCookie, ValueRefsEncoder::new(borrowed).with_sensitive(true)) + .expect("borrowed values encode"); + assert_eq!(sink.values.len(), 2); + assert!(sink.values.iter().all(FieldValue::is_sensitive)); + + let borrowed = [FieldValueRef::new(b"public"), FieldValueRef::new(b"private").with_sensitive(true)]; + sink.set_encoded(&FieldName::Vary, ValueRefsEncoder::new(borrowed)) + .expect("borrowed values preserve markers"); + assert!(!sink.values[0].is_sensitive()); + assert!(sink.values[1].is_sensitive()); + + sink.set_encoded(&FieldName::UserAgent, FieldValue::from_static("client/1")) + .expect("owned value transfers"); + sink.append_encoded(&FieldName::UserAgent, EncodedValues::single(FieldValue::from_static("client/2"))) + .expect("owned values append"); + assert_eq!(sink.values.len(), 2); + assert_eq!(sink.append_calls, 1); + + sink.remove_values(&FieldName::UserAgent); + assert!(sink.values.is_empty()); + } + + #[test] + fn default_sink_rejects_bad_encoders_without_replacing_values() { + let mut sink = Sink { + values: vec![FieldValue::from_static("original")], + append_calls: 0, + }; + for (encoder, expected_kind) in [ + BytesEncoder { + expected: 2, + bytes: b"x", + sensitive: false, + }, + BytesEncoder { + expected: 1, + bytes: b"\n", + sensitive: false, + }, + BytesEncoder { + expected: 0, + bytes: b"x", + sensitive: false, + }, + BytesEncoder { + expected: usize::MAX, + bytes: b"x", + sensitive: false, + }, + ] + .into_iter() + .zip([ + InsertErrorKind::InvalidEncoding, + InsertErrorKind::InvalidValue, + InsertErrorKind::InvalidEncoding, + InsertErrorKind::CapacityExceeded, + ]) { + assert_eq!( + sink.set_encoded(&FieldName::UserAgent, encoder), + Err(InsertError::new(expected_kind)) + ); + assert_eq!(sink.values[0].as_bytes(), b"original"); + } + + assert_eq!( + sink.append_encoded( + &FieldName::UserAgent, + BytesEncoder { + expected: 1, + bytes: b"\n", + sensitive: false, + }, + ), + Err(InsertError::new(InsertErrorKind::InvalidValue)) + ); + assert_eq!(sink.values.len(), 1); + assert_eq!(sink.append_calls, 0); + } + + #[test] + fn impossible_capacities_are_rejected_before_sink_mutation() { + let mut sink = Sink { + values: vec![FieldValue::from_static("original")], + append_calls: 0, + }; + for length in [isize::MAX as usize + 1, usize::MAX] { + let encoder = || BytesEncoder { + expected: length, + bytes: b"x", + sensitive: false, + }; + let expected = Err(InsertError::new(InsertErrorKind::CapacityExceeded)); + assert_eq!(sink.set_encoded(&FieldName::UserAgent, encoder()), expected); + assert_eq!(sink.append_encoded(&FieldName::UserAgent, encoder()), expected); + assert_eq!(sink.values, [FieldValue::from_static("original")]); + assert_eq!(sink.append_calls, 0); + + #[cfg(feature = "http")] + { + let mut map = HeaderMap::new(); + map.insert(header::USER_AGENT, HeaderValue::from_static("original")); + let original = map.clone(); + assert_eq!(map.set_encoded(&FieldName::UserAgent, encoder()), expected); + assert_eq!(map.append_encoded(&FieldName::UserAgent, encoder()), expected); + assert_eq!(map, original); + } + } + } + + #[test] + fn buffer_reservation_reports_failure_without_a_large_allocation() { + // Capacity overflow exercises reservation failure without requesting memory. + assert_eq!(reserve_bytes(usize::MAX), Err(InsertError::new(InsertErrorKind::AllocationFailed))); + let buffer = reserve_bytes(FIELD_VALUE_INLINE_CAPACITY + 1).unwrap(); + assert!(buffer.is_empty()); + assert!(buffer.capacity() > FIELD_VALUE_INLINE_CAPACITY); + + let mut values = EncodedValues::single(FieldValue::from_static("original")); + assert_eq!( + reserve_collected_values(&mut values, usize::MAX), + Err(InsertError::new(InsertErrorKind::AllocationFailed)) + ); + assert_eq!(values.iter().next().unwrap(), "original"); + } + + #[test] + fn encoders_propagate_each_output_failure_without_masking_it() { + assert_eq!( + U64Encoder::new(42).encode(&mut RejectBegin), + Err(InsertError::new(InsertErrorKind::AllocationFailed)) + ); + assert_eq!( + U64Encoder::new(42).encode(&mut RejectWrite), + Err(InsertError::new(InsertErrorKind::InvalidEncoding)) + ); + assert_eq!( + FieldValueRef::new(b"value").encode(&mut RejectBegin), + Err(InsertError::new(InsertErrorKind::AllocationFailed)) + ); + assert_eq!( + FieldValueRef::new(b"value").encode(&mut RejectWrite), + Err(InsertError::new(InsertErrorKind::InvalidEncoding)) + ); + assert_eq!( + FieldValueRef::new(b"value").encode(&mut RejectFinish), + Err(InsertError::new(InsertErrorKind::InvalidValue)) + ); + + let borrowed = [FieldValueRef::new(b"value")]; + assert_eq!( + ValueRefsEncoder::new(borrowed).encode(&mut RejectWrite), + Err(InsertError::new(InsertErrorKind::InvalidEncoding)) + ); + assert_eq!( + EncodedValues::single(FieldValue::from_static("value")).encode(&mut RejectOwned), + Err(InsertError::new(InsertErrorKind::CapacityExceeded)) + ); + assert_eq!(FieldValue::from_static("value").encode(&mut RejectWrite), Ok(())); + assert_eq!(FieldValue::from_static("value").encode(&mut RejectBegin), Ok(())); + assert_eq!(FieldValueRef::new(b"value").encode(&mut RejectOwned), Ok(())); + + let encoder = BytesEncoder { + expected: 5, + bytes: b"value", + sensitive: false, + }; + assert_eq!( + encoder.encode(&mut RejectBegin), + Err(InsertError::new(InsertErrorKind::AllocationFailed)) + ); + let encoder = BytesEncoder { + expected: 5, + bytes: b"value", + sensitive: false, + }; + assert_eq!( + encoder.encode(&mut RejectWrite), + Err(InsertError::new(InsertErrorKind::InvalidEncoding)) + ); + + let sink = Sink { + values: vec![FieldValue::from_static("value")], + append_calls: 0, + }; + assert_eq!( + FieldSource::lines(&&sink, &FieldName::UserAgent) + .expect("forwarded values") + .exactly_one() + .expect("one value"), + "value" + ); + } + + #[test] + fn default_writer_handles_inline_boundary_spill_and_chunking() { + for length in [FIELD_VALUE_INLINE_CAPACITY, FIELD_VALUE_INLINE_CAPACITY + 1] { + let bytes = vec![b'a'; length]; + let mut output = CollectOutput { + values: EncodedValues::new(), + }; + let mut writer = output.begin_value(length, FieldSensitivity::Sensitive).expect("writer starts"); + for chunk in bytes.chunks(7) { + writer.write_bytes(chunk).expect("chunk writes"); + } + writer.finish().expect("complete value finishes"); + + let value = output.values.iter().next().expect("one value"); + assert_eq!(value.as_bytes(), bytes); + assert!(value.is_sensitive()); + } + } +} diff --git a/crates/http_headers/src/sink/field_sink_ext.rs b/crates/http_headers/src/sink/field_sink_ext.rs new file mode 100644 index 000000000..94457337f --- /dev/null +++ b/crates/http_headers/src/sink/field_sink_ext.rs @@ -0,0 +1,975 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Fluent APIs for creating response headers. + +#[cfg(any(test, feature = "headers-cors"))] +use std::time::Duration; + +#[cfg(any( + test, + feature = "headers-authorization", + feature = "headers-conditional", + feature = "headers-location", +))] +use crate::FieldValueRef; +#[cfg(any(test, feature = "headers-negotiation"))] +use crate::headers::{ + Accept, AcceptEncoding, AcceptEncodingOwned, AcceptEncodingView, AcceptLanguage, AcceptLanguageOwned, AcceptLanguageView, AcceptOwned, + AcceptView, Allow, AllowOwned, AllowView, Host, HostOwned, HostView, Server, ServerOwned, ServerView, Vary, VaryOwned, VaryView, +}; +#[cfg(any(test, feature = "headers-range"))] +use crate::headers::{ + AcceptRanges, AcceptRangesOwned, AcceptRangesView, ContentRange, ContentRangeOwned, ContentRangeView, Range, RangeOwned, RangeView, +}; +#[cfg(any(test, feature = "headers-cors"))] +use crate::headers::{ + AccessControlAllowCredentials, AccessControlAllowCredentialsOwned, AccessControlAllowCredentialsView, AccessControlAllowHeaders, + AccessControlAllowHeadersOwned, AccessControlAllowHeadersView, AccessControlAllowMethods, AccessControlAllowMethodsOwned, + AccessControlAllowMethodsView, AccessControlAllowOrigin, AccessControlAllowOriginOwned, AccessControlAllowOriginView, + AccessControlExposeHeaders, AccessControlExposeHeadersOwned, AccessControlExposeHeadersView, AccessControlMaxAge, + AccessControlMaxAgeOwned, AccessControlRequestHeaders, AccessControlRequestHeadersOwned, AccessControlRequestHeadersView, + AccessControlRequestMethod, AccessControlRequestMethodOwned, AccessControlRequestMethodView, +}; +#[cfg(any(test, feature = "headers-authorization"))] +use crate::headers::{Authorization, AuthorizationOwned, AuthorizationView, Basic, Bearer}; +#[cfg(any(test, feature = "headers-cache-control"))] +use crate::headers::{CacheControl, CacheControlBuilder, CacheControlOwned, CacheControlView}; +#[cfg(any(test, feature = "headers-content-length"))] +use crate::headers::{ContentLength, ContentLengthOwned}; +#[cfg(any(test, feature = "headers-security"))] +use crate::headers::{ + ContentSecurityPolicy, ContentSecurityPolicyOwned, ContentSecurityPolicyView, ReferrerPolicy, ReferrerPolicyOwned, ReferrerPolicyView, + StrictTransportSecurity, StrictTransportSecurityOwned, StrictTransportSecurityView, XContentTypeOptions, XContentTypeOptionsOwned, + XContentTypeOptionsView, +}; +#[cfg(any(test, feature = "headers-content-type"))] +use crate::headers::{ContentType, ContentTypeOwned, ContentTypeView}; +#[cfg(any(test, feature = "headers-etag"))] +use crate::headers::{ETag, ETagOwned, ETagView}; +#[cfg(any(test, feature = "headers-conditional"))] +use crate::headers::{ + IfMatch, IfMatchOwned, IfMatchView, IfModifiedSince, IfModifiedSinceOwned, IfModifiedSinceView, IfNoneMatch, IfNoneMatchOwned, + IfNoneMatchView, IfRange, IfRangeOwned, IfRangeView, IfUnmodifiedSince, IfUnmodifiedSinceOwned, IfUnmodifiedSinceView, LastModified, + LastModifiedOwned, LastModifiedView, +}; +#[cfg(any(test, feature = "headers-location"))] +use crate::headers::{Location, LocationOwned, LocationView}; +#[cfg(any(test, feature = "headers-websocket"))] +use crate::headers::{ + SecWebSocketAccept, SecWebSocketAcceptOwned, SecWebSocketAcceptView, SecWebSocketExtensions, SecWebSocketExtensionsOwned, + SecWebSocketExtensionsView, SecWebSocketKey, SecWebSocketKeyOwned, SecWebSocketKeyView, SecWebSocketProtocol, + SecWebSocketProtocolOwned, SecWebSocketProtocolView, SecWebSocketVersion, SecWebSocketVersionOwned, +}; +#[cfg(any(test, feature = "headers-set-cookie"))] +use crate::headers::{SetCookie, SetCookieOwned, SetCookieView}; +#[cfg(any(test, feature = "headers-user-agent"))] +use crate::headers::{UserAgent, UserAgentOwned, UserAgentView}; +#[cfg(any(test, feature = "headers-cors"))] +use crate::sink::InsertErrorKind; +#[cfg(any(test, feature = "headers-content-length", feature = "headers-cors"))] +use crate::sink::U64Encoder; +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-cors", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-websocket", +))] +use crate::sink::ValueRefsEncoder; +use crate::sink::{FieldSink, InsertError}; +use crate::source::FieldSource; +use crate::{DecodeError, Field}; + +macro_rules! built_in_fields { + ($macro:ident) => { + $macro! { + #[cfg(any(test, feature = "headers-cache-control"))] + (CacheControlOwned, CacheControl), + #[cfg(any(test, feature = "headers-negotiation"))] + (AcceptOwned, Accept), + #[cfg(any(test, feature = "headers-negotiation"))] + (AcceptEncodingOwned, AcceptEncoding), + #[cfg(any(test, feature = "headers-negotiation"))] + (AcceptLanguageOwned, AcceptLanguage), + #[cfg(any(test, feature = "headers-negotiation"))] + (AllowOwned, Allow), + #[cfg(any(test, feature = "headers-negotiation"))] + (HostOwned, Host), + #[cfg(any(test, feature = "headers-negotiation"))] + (ServerOwned, Server), + #[cfg(any(test, feature = "headers-negotiation"))] + (VaryOwned, Vary), + #[cfg(any(test, feature = "headers-range"))] + (AcceptRangesOwned, AcceptRanges), + #[cfg(any(test, feature = "headers-range"))] + (ContentRangeOwned, ContentRange), + #[cfg(any(test, feature = "headers-range"))] + (RangeOwned, Range), + #[cfg(any(test, feature = "headers-etag"))] + (ETagOwned, ETag), + #[cfg(any(test, feature = "headers-location"))] + (LocationOwned, Location), + #[cfg(any(test, feature = "headers-user-agent"))] + (UserAgentOwned, UserAgent), + #[cfg(any(test, feature = "headers-content-type"))] + (ContentTypeOwned, ContentType), + #[cfg(any(test, feature = "headers-conditional"))] + (IfMatchOwned, IfMatch), + #[cfg(any(test, feature = "headers-conditional"))] + (IfNoneMatchOwned, IfNoneMatch), + #[cfg(any(test, feature = "headers-conditional"))] + (IfModifiedSinceOwned, IfModifiedSince), + #[cfg(any(test, feature = "headers-conditional"))] + (IfUnmodifiedSinceOwned, IfUnmodifiedSince), + #[cfg(any(test, feature = "headers-conditional"))] + (IfRangeOwned, IfRange), + #[cfg(any(test, feature = "headers-conditional"))] + (LastModifiedOwned, LastModified), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowCredentialsOwned, AccessControlAllowCredentials), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowHeadersOwned, AccessControlAllowHeaders), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowMethodsOwned, AccessControlAllowMethods), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowOriginOwned, AccessControlAllowOrigin), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlExposeHeadersOwned, AccessControlExposeHeaders), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlMaxAgeOwned, AccessControlMaxAge), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlRequestHeadersOwned, AccessControlRequestHeaders), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlRequestMethodOwned, AccessControlRequestMethod), + #[cfg(any(test, feature = "headers-security"))] + (ContentSecurityPolicyOwned, ContentSecurityPolicy), + #[cfg(any(test, feature = "headers-security"))] + (ReferrerPolicyOwned, ReferrerPolicy), + #[cfg(any(test, feature = "headers-security"))] + (StrictTransportSecurityOwned, StrictTransportSecurity), + #[cfg(any(test, feature = "headers-security"))] + (XContentTypeOptionsOwned, XContentTypeOptions), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketAcceptOwned, SecWebSocketAccept), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketExtensionsOwned, SecWebSocketExtensions), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketKeyOwned, SecWebSocketKey), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketProtocolOwned, SecWebSocketProtocol), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketVersionOwned, SecWebSocketVersion), + #[cfg(any(test, feature = "headers-authorization"))] + (AuthorizationOwned, Authorization), + #[cfg(any(test, feature = "headers-authorization"))] + (AuthorizationOwned, Authorization), + #[cfg(any(test, feature = "headers-set-cookie"))] + (SetCookieOwned, SetCookie), + #[cfg(any(test, feature = "headers-content-length"))] + (ContentLengthOwned, ContentLength), + } + }; +} + +macro_rules! insert_owned { + ($($(#[$meta:meta])* ($owned:ty, $header:ty)),+ $(,)?) => { + $( + $(#[$meta])* + impl $owned { + /// Replaces every field line for this value's header. + /// + /// # Errors + /// + /// Returns an error without changing the sink when encoding + /// or storage fails. + pub fn insert_into(self, sink: &mut S) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + <$header as Field>::insert(sink, self) + } + } + )+ + }; +} + +macro_rules! inherent_field_operations { + ($($(#[$meta:meta])* ($owned:ty, $header:ty)),+ $(,)?) => { + $( + $(#[$meta])* + impl $header { + /// Reads a borrowed typed view from `source`. + /// + /// # Errors + /// + /// Returns an error when a present field is malformed. + #[inline] + pub fn view( + source: &S, + ) -> Result::View<'_>>, DecodeError> + where + S: FieldSource + ?Sized, + { + ::view(source) + } + + /// Reads an independently owned field from `source`. + /// + /// # Errors + /// + /// Returns an error when a present field is malformed. + #[inline] + pub fn owned(source: &S) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + ::owned(source) + } + + /// Inserts a field, replacing every existing value with that name. + /// + /// # Errors + /// + /// Returns an error without changing the sink when it cannot + /// hold the encoded values. + #[inline] + pub fn insert(sink: &mut S, value: $owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + ::insert(sink, value) + } + + /// Removes every field value stored for this field. + #[inline] + pub fn remove(sink: &mut S) + where + S: FieldSink + ?Sized, + { + ::remove(sink); + } + } + )+ + }; +} + +built_in_fields!(insert_owned); +built_in_fields!(inherent_field_operations); + +macro_rules! insert_field_view { + ($($(#[$meta:meta])* ($view:ident, $header:ty, $method:ident)),+ $(,)?) => { + $( + $(#[$meta])* + impl<'a> $view<'a> { + /// Replaces every field line for this view's header. + /// + /// # Errors + /// + /// Returns an error without changing the sink when encoding + /// or storage fails. + pub fn insert_into(self, sink: &mut S) -> Result<(), InsertError> + where + S: FieldSink, + { + let value = self.$method(); + sink.set_encoded(<$header as Field>::name(), value) + } + } + )+ + }; +} + +macro_rules! insert_values_view { + ($($(#[$meta:meta])* ($view:ident, $header:ty, $method:ident)),+ $(,)?) => { + $( + $(#[$meta])* + impl<'a> $view<'a> { + /// Replaces every field line for this view's header. + /// + /// # Errors + /// + /// Returns an error without changing the sink when encoding + /// or storage fails. + pub fn insert_into(self, sink: &mut S) -> Result<(), InsertError> + where + S: FieldSink, + { + sink.set_encoded( + <$header as Field>::name(), + ValueRefsEncoder::new(self.$method()), + ) + } + } + )+ + }; +} + +insert_field_view! { + #[cfg(any(test, feature = "headers-content-type"))] + (ContentTypeView, ContentType, as_field_value), + #[cfg(any(test, feature = "headers-conditional"))] + (IfModifiedSinceView, IfModifiedSince, as_field_value), + #[cfg(any(test, feature = "headers-conditional"))] + (IfUnmodifiedSinceView, IfUnmodifiedSince, as_field_value), + #[cfg(any(test, feature = "headers-conditional"))] + (LastModifiedView, LastModified, as_field_value), + #[cfg(any(test, feature = "headers-conditional"))] + (IfRangeView, IfRange, as_field_value), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowCredentialsView, AccessControlAllowCredentials, as_field_value), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowOriginView, AccessControlAllowOrigin, as_field_value), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlRequestMethodView, AccessControlRequestMethod, as_field_value), + #[cfg(any(test, feature = "headers-range"))] + (ContentRangeView, ContentRange, as_field_value), + #[cfg(any(test, feature = "headers-range"))] + (RangeView, Range, as_field_value), + #[cfg(any(test, feature = "headers-security"))] + (StrictTransportSecurityView, StrictTransportSecurity, as_field_value), + #[cfg(any(test, feature = "headers-security"))] + (XContentTypeOptionsView, XContentTypeOptions, as_field_value), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketAcceptView, SecWebSocketAccept, as_field_value), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketKeyView, SecWebSocketKey, as_field_value), + #[cfg(any(test, feature = "headers-etag"))] + (ETagView, ETag, field_value), + #[cfg(any(test, feature = "headers-user-agent"))] + (UserAgentView, UserAgent, field_value), +} + +#[cfg(any(test, feature = "headers-authorization"))] +impl AuthorizationView<'_, Scheme> { + /// Replaces `Authorization` from this borrowed view. + /// + /// # Errors + /// + /// Returns an error without changing the sink when storage fails. + pub fn insert_into(self, sink: &mut S) -> Result<(), InsertError> + where + S: FieldSink, + Authorization: Field, + { + sink.set_encoded( + as Field>::name(), + FieldValueRef::new(self.as_field_value().as_bytes()).with_sensitive(true), + ) + } +} + +#[cfg(any(test, feature = "headers-location"))] +impl LocationView<'_> { + /// Replaces `Location` and marks the inserted value as sensitive. + /// + /// # Errors + /// + /// Returns an error without changing the sink when storage fails. + pub fn insert_into(self, sink: &mut S) -> Result<(), InsertError> + where + S: FieldSink, + { + sink.set_encoded( + ::name(), + FieldValueRef::new(self.as_bytes()).with_sensitive(true), + ) + } +} + +insert_values_view! { + #[cfg(any(test, feature = "headers-negotiation"))] + (AcceptView, Accept, values), + #[cfg(any(test, feature = "headers-negotiation"))] + (AcceptEncodingView, AcceptEncoding, values), + #[cfg(any(test, feature = "headers-negotiation"))] + (AcceptLanguageView, AcceptLanguage, values), + #[cfg(any(test, feature = "headers-negotiation"))] + (AllowView, Allow, values), + #[cfg(any(test, feature = "headers-negotiation"))] + (VaryView, Vary, values), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowHeadersView, AccessControlAllowHeaders, field_values), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlAllowMethodsView, AccessControlAllowMethods, field_values), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlExposeHeadersView, AccessControlExposeHeaders, field_values), + #[cfg(any(test, feature = "headers-cors"))] + (AccessControlRequestHeadersView, AccessControlRequestHeaders, field_values), + #[cfg(any(test, feature = "headers-range"))] + (AcceptRangesView, AcceptRanges, field_values), + #[cfg(any(test, feature = "headers-cache-control"))] + (CacheControlView, CacheControl, field_values), + #[cfg(any(test, feature = "headers-security"))] + (ContentSecurityPolicyView, ContentSecurityPolicy, field_values), + #[cfg(any(test, feature = "headers-security"))] + (ReferrerPolicyView, ReferrerPolicy, field_values), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketExtensionsView, SecWebSocketExtensions, field_values), + #[cfg(any(test, feature = "headers-websocket"))] + (SecWebSocketProtocolView, SecWebSocketProtocol, field_values), +} + +insert_field_view! { + #[cfg(any(test, feature = "headers-negotiation"))] + (HostView, Host, as_field_value), + #[cfg(any(test, feature = "headers-negotiation"))] + (ServerView, Server, as_field_value), +} + +#[cfg(any(test, feature = "headers-set-cookie"))] +impl SetCookieView<'_> { + /// Replaces `Set-Cookie` and marks every inserted field line as sensitive. + /// + /// # Errors + /// + /// Returns an error without changing the sink when storage fails. + pub fn insert_into(self, sink: &mut S) -> Result<(), InsertError> + where + S: FieldSink, + { + sink.set_encoded( + ::name(), + ValueRefsEncoder::new(self.iter()).with_sensitive(true), + ) + } +} + +macro_rules! insert_conditional_tags { + ($($(#[$meta:meta])* ($view:ident, $header:ty)),+ $(,)?) => { + $( + $(#[$meta])* + impl<'a> $view<'a> { + /// Replaces every field line for this conditional view. + /// + /// # Errors + /// + /// Returns an error without changing the sink when storage + /// fails. + pub fn insert_into(self, sink: &mut S) -> Result<(), InsertError> + where + S: FieldSink, + { + if self.is_wildcard() { + sink.set_encoded( + <$header as Field>::name(), + FieldValueRef::new(b"*"), + ) + } else { + sink.set_encoded( + <$header as Field>::name(), + ValueRefsEncoder::new( + self.tags() + .map(|tag| FieldValueRef::new(tag.as_bytes())), + ), + ) + } + } + } + )+ + }; +} + +insert_conditional_tags! { + #[cfg(any(test, feature = "headers-conditional"))] + (IfMatchView, IfMatch), + #[cfg(any(test, feature = "headers-conditional"))] + (IfNoneMatchView, IfNoneMatch), +} + +macro_rules! response_methods { + ($($(#[$meta:meta])* ($method:ident, $header:ty, $owned:ty)),+ $(,)?) => { + $( + $(#[$meta])* + #[doc = concat!("Replaces `", stringify!($header), "` and continues the fluent chain.")] + /// + /// # Errors + /// + /// Returns an error without changing the sink when storage fails. + fn $method(&mut self, value: $owned) -> Result<&mut Self, InsertError> { + <$header as Field>::insert(self, value)?; + Ok(self) + } + )+ + }; +} + +/// Fluent methods for creating headers in any [`FieldSink`]. +/// +/// These write-only conveniences construct response fields. Methods replace +/// the named header unless their name begins with `append_`. +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(all( +/// # feature = "http", +/// # feature = "headers-content-length", +/// # feature = "headers-content-type" +/// # ))] +/// # fn main() -> Result<(), http_headers::sink::InsertError> { +/// use http_headers::headers::ContentType; +/// use http_headers::sink::FieldSinkExt; +/// +/// let mut headers = http::HeaderMap::new(); +/// headers +/// .set_content_type(ContentType::json())? +/// .set_content_length(128)?; +/// # Ok::<(), http_headers::sink::InsertError>(()) +/// # } +/// # #[cfg(not(all( +/// # feature = "http", +/// # feature = "headers-content-length", +/// # feature = "headers-content-type" +/// # )))] +/// # fn main() {} +/// ``` +pub trait FieldSinkExt: FieldSink + Sized { + response_methods! { + #[cfg(any(test, feature = "headers-negotiation"))] + (set_accept, Accept, AcceptOwned), + #[cfg(any(test, feature = "headers-negotiation"))] + (set_accept_encoding, AcceptEncoding, AcceptEncodingOwned), + #[cfg(any(test, feature = "headers-negotiation"))] + (set_accept_language, AcceptLanguage, AcceptLanguageOwned), + #[cfg(any(test, feature = "headers-negotiation"))] + (set_allow, Allow, AllowOwned), + #[cfg(any(test, feature = "headers-negotiation"))] + (set_host, Host, HostOwned), + #[cfg(any(test, feature = "headers-negotiation"))] + (set_server, Server, ServerOwned), + #[cfg(any(test, feature = "headers-negotiation"))] + (set_vary, Vary, VaryOwned), + #[cfg(any(test, feature = "headers-range"))] + (set_accept_ranges, AcceptRanges, AcceptRangesOwned), + #[cfg(any(test, feature = "headers-range"))] + (set_content_range, ContentRange, ContentRangeOwned), + #[cfg(any(test, feature = "headers-range"))] + (set_range, Range, RangeOwned), + #[cfg(any(test, feature = "headers-etag"))] + (set_etag, ETag, ETagOwned), + #[cfg(any(test, feature = "headers-location"))] + (set_location, Location, LocationOwned), + #[cfg(any(test, feature = "headers-user-agent"))] + (set_user_agent, UserAgent, UserAgentOwned), + #[cfg(any(test, feature = "headers-content-type"))] + (set_content_type, ContentType, ContentTypeOwned), + #[cfg(any(test, feature = "headers-conditional"))] + (set_if_match, IfMatch, IfMatchOwned), + #[cfg(any(test, feature = "headers-conditional"))] + (set_if_none_match, IfNoneMatch, IfNoneMatchOwned), + #[cfg(any(test, feature = "headers-conditional"))] + (set_if_modified_since, IfModifiedSince, IfModifiedSinceOwned), + #[cfg(any(test, feature = "headers-conditional"))] + (set_if_unmodified_since, IfUnmodifiedSince, IfUnmodifiedSinceOwned), + #[cfg(any(test, feature = "headers-conditional"))] + (set_if_range, IfRange, IfRangeOwned), + #[cfg(any(test, feature = "headers-conditional"))] + (set_last_modified, LastModified, LastModifiedOwned), + #[cfg(any(test, feature = "headers-cors"))] + (set_access_control_allow_credentials, AccessControlAllowCredentials, AccessControlAllowCredentialsOwned), + #[cfg(any(test, feature = "headers-cors"))] + (set_access_control_allow_headers, AccessControlAllowHeaders, AccessControlAllowHeadersOwned), + #[cfg(any(test, feature = "headers-cors"))] + (set_access_control_allow_methods, AccessControlAllowMethods, AccessControlAllowMethodsOwned), + #[cfg(any(test, feature = "headers-cors"))] + (set_access_control_allow_origin, AccessControlAllowOrigin, AccessControlAllowOriginOwned), + #[cfg(any(test, feature = "headers-cors"))] + (set_access_control_expose_headers, AccessControlExposeHeaders, AccessControlExposeHeadersOwned), + #[cfg(any(test, feature = "headers-cors"))] + (set_access_control_max_age_value, AccessControlMaxAge, AccessControlMaxAgeOwned), + #[cfg(any(test, feature = "headers-cors"))] + (set_access_control_request_headers, AccessControlRequestHeaders, AccessControlRequestHeadersOwned), + #[cfg(any(test, feature = "headers-cors"))] + (set_access_control_request_method, AccessControlRequestMethod, AccessControlRequestMethodOwned), + #[cfg(any(test, feature = "headers-security"))] + (set_content_security_policy, ContentSecurityPolicy, ContentSecurityPolicyOwned), + #[cfg(any(test, feature = "headers-security"))] + (set_referrer_policy, ReferrerPolicy, ReferrerPolicyOwned), + #[cfg(any(test, feature = "headers-security"))] + (set_strict_transport_security, StrictTransportSecurity, StrictTransportSecurityOwned), + #[cfg(any(test, feature = "headers-security"))] + (set_x_content_type_options, XContentTypeOptions, XContentTypeOptionsOwned), + #[cfg(any(test, feature = "headers-websocket"))] + (set_sec_websocket_accept, SecWebSocketAccept, SecWebSocketAcceptOwned), + #[cfg(any(test, feature = "headers-websocket"))] + (set_sec_websocket_extensions, SecWebSocketExtensions, SecWebSocketExtensionsOwned), + #[cfg(any(test, feature = "headers-websocket"))] + (set_sec_websocket_key, SecWebSocketKey, SecWebSocketKeyOwned), + #[cfg(any(test, feature = "headers-websocket"))] + (set_sec_websocket_protocol, SecWebSocketProtocol, SecWebSocketProtocolOwned), + #[cfg(any(test, feature = "headers-websocket"))] + (set_sec_websocket_version, SecWebSocketVersion, SecWebSocketVersionOwned), + #[cfg(any(test, feature = "headers-authorization"))] + (set_basic_authorization, Authorization, AuthorizationOwned), + #[cfg(any(test, feature = "headers-authorization"))] + (set_bearer_authorization, Authorization, AuthorizationOwned), + #[cfg(any(test, feature = "headers-set-cookie"))] + (set_set_cookie, SetCookie, SetCookieOwned), + } + + /// Replaces `Content-Length` from its semantic integer value. + /// + /// # Errors + /// + /// Returns an error without changing the sink when storage fails. + #[cfg(any(test, feature = "headers-content-length"))] + fn set_content_length(&mut self, length: u64) -> Result<&mut Self, InsertError> { + self.set_encoded(::name(), U64Encoder::new(length))?; + Ok(self) + } + + /// Replaces `Cache-Control` from a response cache policy. + /// + /// # Errors + /// + /// Returns an error without changing the sink when the plan is invalid + /// or storage fails. + #[cfg(any(test, feature = "headers-cache-control"))] + fn set_cache_control(&mut self, plan: CacheControlBuilder) -> Result<&mut Self, InsertError> { + self.set_encoded(::name(), plan)?; + Ok(self) + } + + /// Replaces `Access-Control-Max-Age` from seconds. + /// + /// # Errors + /// + /// Returns an error without changing the sink when storage fails. + #[cfg(any(test, feature = "headers-cors"))] + fn set_access_control_max_age(&mut self, seconds: u64) -> Result<&mut Self, InsertError> { + self.set_encoded(::name(), U64Encoder::new(seconds))?; + Ok(self) + } + + /// Replaces `Access-Control-Max-Age` from a whole-second duration. + /// + /// # Errors + /// + /// Returns [`InsertErrorKind::InvalidValue`] without changing the sink + /// if the duration contains fractional seconds, or an error if storage fails. + #[cfg(any(test, feature = "headers-cors"))] + fn set_access_control_max_age_duration(&mut self, duration: Duration) -> Result<&mut Self, InsertError> { + let value = + AccessControlMaxAgeOwned::from_duration(duration).map_err(|_invalid| InsertError::new(InsertErrorKind::InvalidValue))?; + self.set_access_control_max_age(value.seconds()) + } + + /// Appends the supplied `Set-Cookie` field lines. + /// + /// # Errors + /// + /// Returns an error without changing the sink when storage fails. + #[cfg(any(test, feature = "headers-set-cookie"))] + fn append_set_cookie(&mut self, value: SetCookieOwned) -> Result<&mut Self, InsertError> { + self.append_values(::name(), value.into_sensitive_encoded_values()?)?; + Ok(self) + } +} + +impl FieldSinkExt for S where S: FieldSink {} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::*; + use crate::source::FieldLines; + use crate::{FieldName, FieldValue, TestSink}; + + struct Source<'a> { + name: &'static FieldName, + values: &'a [FieldValue], + } + + impl FieldSource for Source<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::from_slice(name, self.values)).flatten() + } + } + + struct RejectSink { + removed: bool, + } + + impl FieldSource for RejectSink { + fn lines(&self, _name: &'static FieldName) -> Option> { + None + } + } + + impl FieldSink for RejectSink { + fn set_values(&mut self, _name: &'static FieldName, _values: crate::sink::EncodedValues) -> Result<(), InsertError> { + Err(InsertError::new(InsertErrorKind::CapacityExceeded)) + } + + fn append_values(&mut self, _name: &'static FieldName, _values: crate::sink::EncodedValues) -> Result<(), InsertError> { + Err(InsertError::new(InsertErrorKind::CapacityExceeded)) + } + + fn remove_values(&mut self, _name: &'static FieldName) { + self.removed = true; + } + } + + macro_rules! exercise_insert { + ($header:ty, [$($value:literal),+ $(,)?]) => {{ + let values = [$(FieldValue::from_static($value)),+]; + let source = Source { + name: <$header as Field>::name(), + values: &values, + }; + let mut sink = TestSink::new(); + <$header as Field>::owned(&source) + .expect("fixture is valid") + .expect("fixture is present") + .insert_into(&mut sink) + .expect("owned insertion succeeds"); + <$header as Field>::view(&source) + .expect("fixture is valid") + .expect("fixture is present") + .insert_into(&mut sink) + .expect("view insertion succeeds"); + sink + }}; + } + + macro_rules! exercise { + ($header:ty, $setter:ident, [$($value:literal),+ $(,)?]) => {{ + let values = [$(FieldValue::from_static($value)),+]; + let source = Source { + name: <$header as Field>::name(), + values: &values, + }; + let mut sink = exercise_insert!($header, [$($value),+]); + sink.$setter( + <$header as Field>::owned(&source) + .expect("fixture is valid") + .expect("fixture is present"), + ) + .expect("fluent insertion succeeds"); + sink + }}; + } + + #[test] + fn every_generated_owned_view_and_fluent_path_inserts() { + drop(exercise!(Accept, set_accept, ["text/html"])); + drop(exercise!(AcceptEncoding, set_accept_encoding, ["gzip"])); + drop(exercise!(AcceptLanguage, set_accept_language, ["en-US"])); + drop(exercise!(Allow, set_allow, ["GET, POST"])); + drop(exercise!(Host, set_host, ["example.com:443"])); + drop(exercise!(Server, set_server, ["example/1.0"])); + drop(exercise!(Vary, set_vary, ["accept-encoding, origin"])); + drop(exercise!(AcceptRanges, set_accept_ranges, ["bytes"])); + drop(exercise!(ContentRange, set_content_range, ["bytes 0-1/2"])); + drop(exercise!(Range, set_range, ["bytes=0-1"])); + drop(exercise!(ETag, set_etag, ["\"a\""])); + drop(exercise!(Location, set_location, ["/docs"])); + drop(exercise!(UserAgent, set_user_agent, ["client/1"])); + drop(exercise!(ContentType, set_content_type, ["application/json"])); + drop(exercise!(IfMatch, set_if_match, ["\"a\""])); + drop(exercise!(IfNoneMatch, set_if_none_match, ["W/\"a\""])); + drop(exercise!(IfModifiedSince, set_if_modified_since, ["Sun, 06 Nov 1994 08:49:37 GMT"])); + drop(exercise!( + IfUnmodifiedSince, + set_if_unmodified_since, + ["Sun, 06 Nov 1994 08:49:37 GMT"] + )); + drop(exercise!(IfRange, set_if_range, ["\"a\""])); + drop(exercise!(LastModified, set_last_modified, ["Sun, 06 Nov 1994 08:49:37 GMT"])); + drop(exercise!( + AccessControlAllowCredentials, + set_access_control_allow_credentials, + ["true"] + )); + drop(exercise!( + AccessControlAllowHeaders, + set_access_control_allow_headers, + ["content-type, x-request-id"] + )); + drop(exercise!( + AccessControlAllowMethods, + set_access_control_allow_methods, + ["GET, POST"] + )); + drop(exercise!( + AccessControlAllowOrigin, + set_access_control_allow_origin, + ["https://example.com"] + )); + drop(exercise!( + AccessControlExposeHeaders, + set_access_control_expose_headers, + ["etag, x-request-id"] + )); + drop(exercise!(AccessControlMaxAge, set_access_control_max_age_value, ["600"])); + drop(exercise!( + AccessControlRequestHeaders, + set_access_control_request_headers, + ["content-type, x-request-id"] + )); + drop(exercise!(AccessControlRequestMethod, set_access_control_request_method, ["POST"])); + drop(exercise!( + ContentSecurityPolicy, + set_content_security_policy, + ["default-src 'self'"] + )); + drop(exercise!(ReferrerPolicy, set_referrer_policy, ["no-referrer"])); + drop(exercise!( + StrictTransportSecurity, + set_strict_transport_security, + ["max-age=60; includeSubDomains"] + )); + drop(exercise!(XContentTypeOptions, set_x_content_type_options, ["nosniff"])); + drop(exercise!( + SecWebSocketAccept, + set_sec_websocket_accept, + ["s3pPLMBiTxaQ9kYGzzhZRbK+xOo="] + )); + drop(exercise!( + SecWebSocketExtensions, + set_sec_websocket_extensions, + ["permessage-deflate"] + )); + drop(exercise!(SecWebSocketKey, set_sec_websocket_key, ["dGhlIHNhbXBsZSBub25jZQ=="])); + drop(exercise!(SecWebSocketProtocol, set_sec_websocket_protocol, ["chat, superchat"])); + drop(exercise!(SecWebSocketVersion, set_sec_websocket_version, ["13"])); + drop(exercise!(Authorization, set_basic_authorization, ["Basic Zm9vOmJhcg=="])); + drop(exercise!(Authorization, set_bearer_authorization, ["Bearer token"])); + drop(exercise!(SetCookie, set_set_cookie, ["a=1"])); + } + + #[test] + fn semantic_and_specialized_response_paths_insert() { + let mut sink = exercise_insert!(CacheControl, ["private, max-age=60"]); + sink.set_cache_control( + CacheControl::public() + .no_store() + .must_revalidate() + .immutable() + .max_age(Duration::from_secs(u64::MAX)) + .extension_value("x-test", "enabled"), + ) + .expect("cache plan inserts"); + sink.set_cache_control(CacheControl::private()) + .expect("private plan inserts") + .set_cache_control(CacheControl::no_cache()) + .expect("no-cache plan inserts") + .set_content_length(u64::MAX) + .expect("content length inserts") + .set_access_control_max_age(600) + .expect("max age inserts") + .set_access_control_max_age_duration(Duration::from_mins(5)) + .expect("duration inserts") + .set_content_type(ContentType::json()) + .expect("JSON content type inserts"); + + let mut cookies = SetCookieOwned::new(); + cookies.push_str("a=1").expect("cookie is valid"); + sink.append_set_cookie(cookies).expect("first cookie append succeeds"); + let mut cookies = SetCookieOwned::new(); + cookies.push_str("b=2").expect("cookie is valid"); + sink.append_set_cookie(cookies).expect("second cookie append succeeds"); + + ContentLengthOwned::new(42) + .insert_into(&mut sink) + .expect("numeric owned insertion succeeds"); + } + + #[test] + fn specialized_views_cover_wildcards_and_general_storage() { + drop(exercise_insert!(IfMatch, ["*"])); + drop(exercise_insert!(IfNoneMatch, ["*"])); + drop(exercise_insert!(AcceptRanges, ["items"])); + drop(exercise_insert!(SecWebSocketVersion, ["7, 8, 13"])); + drop(exercise_insert!(ContentSecurityPolicy, ["default-src 'self'", "img-src https:"])); + drop(exercise_insert!(ReferrerPolicy, ["no-referrer", "same-origin"])); + drop(exercise_insert!(SecWebSocketExtensions, ["permessage-deflate", "x-test"])); + drop(exercise_insert!(SecWebSocketProtocol, ["chat", "superchat"])); + drop(exercise_insert!(SetCookie, ["a=1", "b=2"])); + } + + #[test] + fn borrowed_view_insertion_preserves_source_sensitivity() { + let values = [FieldValue::from_static("client/1").with_sensitive(true)]; + let source = Source { + name: &FieldName::UserAgent, + values: &values, + }; + let mut sink = TestSink::new(); + UserAgent::view(&source) + .expect("fixture is valid") + .expect("fixture is present") + .insert_into(&mut sink) + .expect("view insertion succeeds"); + assert!( + sink.lines(&FieldName::UserAgent) + .expect("inserted value") + .exactly_one() + .expect("one value") + .is_sensitive() + ); + + let values = [ + FieldValue::from_static("text/html"), + FieldValue::from_static("application/json").with_sensitive(true), + ]; + let source = Source { + name: &FieldName::Accept, + values: &values, + }; + Accept::view(&source) + .expect("fixture is valid") + .expect("fixture is present") + .insert_into(&mut sink) + .expect("view insertion succeeds"); + let inserted: Vec<_> = sink + .lines(&FieldName::Accept) + .expect("inserted values") + .repeated() + .map(FieldValueRef::is_sensitive) + .collect(); + assert_eq!(inserted, [false, true]); + } + + #[test] + fn empty_cache_plan_reports_insertion_error() { + let mut sink = TestSink::new(); + assert_eq!( + sink.set_cache_control(CacheControlOwned::builder()).err(), + Some(InsertError::new(InsertErrorKind::InvalidValue)) + ); + } + + #[test] + fn fluent_methods_propagate_sink_failures() { + let mut sink = RejectSink { removed: false }; + assert_eq!( + sink.set_user_agent(UserAgentOwned::try_from_static("client/1").expect("valid")) + .err(), + Some(InsertError::new(InsertErrorKind::CapacityExceeded)) + ); + assert_eq!( + sink.set_content_length(42).err(), + Some(InsertError::new(InsertErrorKind::CapacityExceeded)) + ); + assert_eq!( + sink.set_access_control_max_age(600).err(), + Some(InsertError::new(InsertErrorKind::CapacityExceeded)) + ); + + let mut cookies = SetCookieOwned::new(); + cookies.push_str("a=1").expect("valid cookie"); + assert_eq!( + sink.append_set_cookie(cookies).err(), + Some(InsertError::new(InsertErrorKind::CapacityExceeded)) + ); + + sink.remove_values(&FieldName::UserAgent); + assert!(sink.removed); + } +} diff --git a/crates/http_headers/src/sink/insert_error.rs b/crates/http_headers/src/sink/insert_error.rs new file mode 100644 index 000000000..cea211ac2 --- /dev/null +++ b/crates/http_headers/src/sink/insert_error.rs @@ -0,0 +1,115 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! The error reported when a field cannot be encoded or stored. + +use std::error::Error; +use std::fmt; + +/// The category of an encoding or storage failure. +/// +/// Categories describe the failure, not whether retrying will succeed. +/// +/// # Examples +/// +/// ``` +/// use http_headers::sink::{InsertError, InsertErrorKind}; +/// +/// let error = InsertError::new(InsertErrorKind::CapacityExceeded); +/// assert_eq!(error.kind(), InsertErrorKind::CapacityExceeded); +/// ``` +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +#[non_exhaustive] +pub enum InsertErrorKind { + /// A value does not satisfy the field's grammar or HTTP field-byte rules. + InvalidValue, + /// An encoder wrote a different number of bytes than it announced. + InvalidEncoding, + /// A buffer reservation failed. + /// + /// The built-in sinks check supported size limits before reserving. + /// This category alone does not mean retrying will succeed. + AllocationFailed, + /// A value or container would exceed a supported size limit. + CapacityExceeded, +} + +impl fmt::Display for InsertErrorKind { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(match self { + Self::InvalidValue => "invalid field value", + Self::InvalidEncoding => "encoded field length does not match the announced length", + Self::AllocationFailed => "field buffer capacity could not be reserved", + Self::CapacityExceeded => "field value or container capacity exceeded", + }) + } +} + +/// An error produced when a field cannot be stored. +/// +/// The error retains a compact category, not field contents or underlying +/// errors. This keeps it [`Copy`] and avoids retaining sensitive values. +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(all(feature = "http", feature = "headers-user-agent"))] +/// # fn main() -> Result<(), Box> { +/// use http::HeaderMap; +/// use http_headers::Field; +/// use http_headers::headers::{UserAgent, UserAgentOwned}; +/// use http_headers::sink::InsertError; +/// +/// let mut map = HeaderMap::new(); +/// let stored: Result<(), InsertError> = +/// UserAgent::insert(&mut map, UserAgentOwned::try_from_static("client/1")?); +/// assert!(stored.is_ok()); +/// # Ok::<(), Box>(()) +/// # } +/// # #[cfg(not(all(feature = "http", feature = "headers-user-agent")))] +/// # fn main() {} +/// ``` +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)] +pub struct InsertError { + kind: InsertErrorKind, +} + +impl InsertError { + /// Creates an error with the supplied failure category. + /// + /// # Examples + /// + /// ``` + /// use http_headers::sink::{InsertError, InsertErrorKind}; + /// + /// let error = InsertError::new(InsertErrorKind::InvalidValue); + /// assert_eq!(error.kind(), InsertErrorKind::InvalidValue); + /// ``` + #[must_use] + pub const fn new(kind: InsertErrorKind) -> Self { + Self { kind } + } + + /// Returns the encoding or storage failure category. + /// + /// # Examples + /// + /// ``` + /// use http_headers::sink::{InsertError, InsertErrorKind}; + /// + /// let error = InsertError::new(InsertErrorKind::InvalidEncoding); + /// assert_eq!(error.kind(), InsertErrorKind::InvalidEncoding); + /// ``` + #[must_use] + pub const fn kind(self) -> InsertErrorKind { + self.kind + } +} + +impl fmt::Display for InsertError { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.kind.fmt(f) + } +} + +impl Error for InsertError {} diff --git a/crates/http_headers/src/sink/mod.rs b/crates/http_headers/src/sink/mod.rs new file mode 100644 index 000000000..ea25f3018 --- /dev/null +++ b/crates/http_headers/src/sink/mod.rs @@ -0,0 +1,73 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Definitions involved in writing fields to field containers. +//! +//! A container implements [`FieldSink`] to accept field lines, and +//! `FieldSinkExt` provides the feature-gated fluent setters most callers use. A +//! field reaches a sink through [`FieldEncoder`], which emits one value per +//! field line into a [`FieldEncodeOutput`]; [`EncodedValues`] is the owned +//! form those values collect into. +//! +//! The `http` cargo feature implements [`FieldSink`] for `http::HeaderMap`. +//! Implement it yourself only when integrating a different container. +//! +//! [`crate::FieldSensitivity`] is a crate-root value type, +//! rather than a second export from this module. +//! +//! ```compile_fail +//! use http_headers::sink::FieldSensitivity; +//! ``` + +mod encoded_values; +pub(crate) mod field_encoder; +mod field_sink; +#[cfg(any( + test, + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +mod field_sink_ext; +mod insert_error; + +#[doc(inline)] +pub use encoded_values::{EncodedValues, EncodedValuesIntoIter, EncodedValuesIter, EncodedValuesIterMut}; +#[doc(inline)] +pub(crate) use field_encoder::FieldSensitivity; +#[doc(inline)] +pub use field_encoder::{FieldEncodeOutput, FieldEncoder, FieldValueWriter, U64Encoder, ValueRefsEncoder}; +#[doc(inline)] +pub use field_sink::FieldSink; +#[cfg(any( + test, + feature = "headers-authorization", + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-content-length", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-etag", + feature = "headers-location", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-user-agent", + feature = "headers-websocket", +))] +#[doc(inline)] +pub use field_sink_ext::FieldSinkExt; +#[doc(inline)] +pub use insert_error::{InsertError, InsertErrorKind}; diff --git a/crates/http_headers/src/source/delimited_items.rs b/crates/http_headers/src/source/delimited_items.rs new file mode 100644 index 000000000..0e938c132 --- /dev/null +++ b/crates/http_headers/src/source/delimited_items.rs @@ -0,0 +1,380 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Delimiter-aware iteration over the items within field lines. + +use std::fmt; + +use crate::source::{FieldLines, FieldLinesIter, MAX_CUSTOM_LIST_ITEMS}; +use crate::{DecodeError, DecodeErrorKind, FieldName}; + +/// A borrowing iterator over delimited items across multiple field lines. +/// +/// Each item is yielded as the raw bytes between delimiters, with optional +/// whitespace already trimmed from both ends. Delimiters inside quoted strings +/// are treated as part of the item, and backslash escapes inside quoted +/// strings are honored within each physical field line. A quoted string or +/// escape cannot span field lines because every yielded item borrows one +/// contiguous line. Items are bytes rather than `&str` because a field value +/// may contain obs-text, which need not be UTF-8. +/// +/// Each item is returned as a `Result`; an unterminated quoted string produces +/// [`DecodeErrorKind::UnterminatedQuote`] with the affected field-line index. +/// Custom sources are also rejected before yielding an item when their byte or +/// line budget is exceeded, and after [`MAX_CUSTOM_LIST_ITEMS`] parsed items. +/// These admission failures yield [`DecodeErrorKind::SourceLimitExceeded`] +/// once, then end iteration. Invalid field bytes remain +/// [`DecodeErrorKind::InvalidSyntax`]. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::FieldName; +/// use http_headers::source::{DelimitedItems, FieldLines}; +/// +/// let lines = FieldLines::single(&FieldName::Vary, b" accept , origin "); +/// let mut items: DelimitedItems<'_> = lines.comma_items(); +/// assert_eq!(items.next().transpose()?, Some(b"accept".as_slice())); +/// assert_eq!(items.next().transpose()?, Some(b"origin".as_slice())); +/// assert!(items.next().is_none()); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct DelimitedItems<'a> { + name: &'static FieldName, + values: FieldLinesIter<'a>, + delimiter: u8, + skip_empty: bool, + current: Option<&'a [u8]>, + start: usize, + position: usize, + value_index: usize, + item_count: usize, + pending_error: Option, + limited: bool, + finished: bool, +} + +impl fmt::Debug for DelimitedItems<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("DelimitedItems") + .field("name", &self.name) + .field("delimiter", &self.delimiter) + .field("value_index", &self.value_index) + .finish_non_exhaustive() + } +} + +impl<'a> DelimitedItems<'a> { + pub(super) fn new(values: &FieldLines<'a>, delimiter: u8) -> Self { + Self { + name: values.name(), + values: values.repeated(), + delimiter, + skip_empty: delimiter == b',', + current: None, + start: 0, + position: 0, + value_index: 0, + item_count: 0, + pending_error: values.validate_custom_source().err(), + limited: values.has_custom_source_limits(), + finished: false, + } + } + + #[cfg(any(test, feature = "headers-negotiation"))] + pub(super) fn from_validated_source(values: &FieldLines<'a>, delimiter: u8) -> Self { + Self { + name: values.name(), + values: values.repeated(), + delimiter, + skip_empty: delimiter == b',', + current: None, + start: 0, + position: 0, + value_index: 0, + item_count: 0, + pending_error: None, + limited: values.has_custom_source_limits(), + finished: false, + } + } + + fn item(&mut self, item: &'a [u8]) -> Result<&'a [u8], DecodeError> { + if self.limited && self.item_count == MAX_CUSTOM_LIST_ITEMS { + self.finished = true; + return Err(DecodeError::new(self.name, DecodeErrorKind::SourceLimitExceeded)); + } + if self.limited { + self.item_count += 1; + } + Ok(item) + } +} + +fn trim_ows(bytes: &[u8]) -> &[u8] { + let start = bytes.iter().position(|byte| !matches!(byte, b' ' | b'\t')).unwrap_or(bytes.len()); + let end = bytes + .iter() + .rposition(|byte| !matches!(byte, b' ' | b'\t')) + .map_or(start, |index| index + 1); + &bytes[start..end] +} + +#[inline] +fn find_delimiter_or_quote(bytes: &[u8], delimiter: u8) -> Option { + http_headers_simd::find_either(bytes, delimiter, b'"') +} + +impl<'a> Iterator for DelimitedItems<'a> { + type Item = Result<&'a [u8], DecodeError>; + + fn next(&mut self) -> Option { + if self.finished { + return None; + } + if let Some(error) = self.pending_error.take() { + self.finished = true; + return Some(Err(error)); + } + + loop { + let bytes = if let Some(bytes) = self.current { + bytes + } else { + let value = self.values.next()?; + let bytes = value.as_bytes(); + self.current = Some(bytes); + self.start = 0; + self.position = 0; + bytes + }; + + let mut quoted = false; + let mut escaped = false; + while self.position < bytes.len() { + if !quoted && let Some(skip) = find_delimiter_or_quote(&bytes[self.position..], self.delimiter) { + self.position += skip; + } else if !quoted { + self.position = bytes.len(); + break; + } + + let byte = bytes[self.position]; + if escaped { + escaped = false; + } else if quoted && byte == b'\\' { + escaped = true; + } else if byte == b'"' { + quoted = !quoted; + } else if !quoted && byte == self.delimiter { + let item = trim_ows(&bytes[self.start..self.position]); + self.position += 1; + self.start = self.position; + if self.skip_empty && item.is_empty() { + continue; + } + return Some(self.item(item)); + } + self.position += 1; + } + + let index = self.value_index; + self.value_index += 1; + let item = trim_ows(&bytes[self.start..]); + self.current = None; + if quoted || escaped { + return Some(Err(DecodeError::new(self.name, DecodeErrorKind::UnterminatedQuote).at_value(index))); + } + + if !self.skip_empty || !item.is_empty() { + return Some(self.item(item)); + } + } + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::iter; + + use crate::source::{FieldLines, MAX_CUSTOM_FIELD_BYTES, MAX_CUSTOM_FIELD_LINES, MAX_CUSTOM_LIST_ITEMS}; + use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef}; + + fn assert_validated_iteration_matches_checked(values: &FieldLines<'_>) { + values.validate_custom_source().unwrap(); + for _ in 0..3 { + let mut checked = values.comma_items(); + let mut validated = values.validated_comma_items(); + loop { + let expected = checked.next(); + assert_eq!(validated.next(), expected); + if expected.is_none() { + break; + } + } + assert_eq!(checked.next(), None); + assert_eq!(validated.next(), None); + assert_eq!(validated.next(), None); + } + } + + #[test] + fn validated_iteration_matches_all_retained_source_representations() { + let bytes = [ + b" \t, alpha, \"b,;\\\"\xff\", bare\\, tail,, ".as_slice(), + b"", + b"alpha, beta;q=0.000000000000000000001", + ]; + for line in bytes { + assert_validated_iteration_matches_checked(&FieldLines::single(&FieldName::Accept, line)); + } + let stored: Vec<_> = bytes.into_iter().map(|line| FieldValue::from_bytes(line).unwrap()).collect(); + assert_validated_iteration_matches_checked(&FieldLines::from_slice(&FieldName::Accept, &stored).unwrap()); + let borrowed: Vec<_> = bytes.into_iter().map(FieldValueRef::new).collect(); + assert_validated_iteration_matches_checked(&FieldLines::from_borrowed(&FieldName::Accept, &borrowed).unwrap()); + + #[cfg(feature = "http")] + { + let mut map = http::HeaderMap::new(); + for line in bytes { + map.append(http::header::ACCEPT, http::HeaderValue::from_bytes(line).unwrap()); + } + let values = FieldLines::from_http(&FieldName::Accept, map.get_all(http::header::ACCEPT)).unwrap(); + assert_validated_iteration_matches_checked(&values); + } + } + + #[test] + fn validated_iteration_keeps_quote_errors_and_custom_item_limits() { + let split_quotes = [ + FieldValueRef::new(b"alpha,\"unterminated"), + FieldValueRef::new(b"\"escaped\\"), + FieldValueRef::new(b"omega"), + ]; + let values = FieldLines::from_borrowed(&FieldName::Accept, &split_quotes).unwrap(); + assert_validated_iteration_matches_checked(&values); + let actual = values.validated_comma_items().collect::>(); + assert_eq!( + actual, + [ + Ok(b"alpha".as_slice()), + Err(DecodeError::new(&FieldName::Accept, DecodeErrorKind::UnterminatedQuote).at_value(0)), + Err(DecodeError::new(&FieldName::Accept, DecodeErrorKind::UnterminatedQuote).at_value(1)), + Ok(b"omega".as_slice()), + ] + ); + + for count in [MAX_CUSTOM_LIST_ITEMS - 1, MAX_CUSTOM_LIST_ITEMS, MAX_CUSTOM_LIST_ITEMS + 1] { + let bytes = iter::repeat_n("alpha", count).collect::>().join(","); + let values = FieldLines::single(&FieldName::Accept, bytes.as_bytes()); + assert_validated_iteration_matches_checked(&values); + let mut entries = values.validated_comma_items(); + for _ in 0..count.min(MAX_CUSTOM_LIST_ITEMS) { + assert_eq!(entries.next(), Some(Ok(b"alpha".as_slice()))); + } + if count > MAX_CUSTOM_LIST_ITEMS { + assert_eq!( + entries.next(), + Some(Err(DecodeError::new(&FieldName::Accept, DecodeErrorKind::SourceLimitExceeded))) + ); + } + assert_eq!(entries.next(), None); + } + } + + #[test] + fn public_iteration_keeps_source_preflight_error_precedence() { + fn rejected(values: &FieldLines<'_>, kind: DecodeErrorKind) { + let expected = DecodeError::new(&FieldName::Accept, kind); + assert_eq!(values.validate_custom_source(), Err(expected)); + let mut items = values.comma_items(); + assert_eq!(items.next(), Some(Err(expected))); + assert_eq!(items.next(), None); + assert_eq!(items.next(), None); + } + + let raw_invalid = [FieldValueRef::new(b"\"unterminated"), FieldValueRef::new(b"\r")]; + rejected( + &FieldLines::from_borrowed(&FieldName::Accept, &raw_invalid).unwrap(), + DecodeErrorKind::InvalidSyntax, + ); + let too_many_lines = vec![FieldValueRef::new(b"\"unterminated"); MAX_CUSTOM_FIELD_LINES + 1]; + rejected( + &FieldLines::from_borrowed(&FieldName::Accept, &too_many_lines).unwrap(), + DecodeErrorKind::SourceLimitExceeded, + ); + let mut too_many_bytes = vec![b'a'; MAX_CUSTOM_FIELD_BYTES + 1]; + too_many_bytes[0] = b'"'; + rejected( + &FieldLines::single(&FieldName::Accept, &too_many_bytes), + DecodeErrorKind::SourceLimitExceeded, + ); + } + + #[test] + fn delimited_iteration_handles_ows_empty_items_quotes_escapes_and_errors() { + let stored = [ + FieldValue::from_static(" , alpha, \"bravo,charlie\" ,, "), + FieldValue::from_static("delta"), + ]; + let values = FieldLines::from_slice(&FieldName::Vary, &stored).expect("values"); + let mut items = values.comma_items(); + assert_eq!( + format!("{items:?}"), + "DelimitedItems { name: \"vary\", delimiter: 44, value_index: 0, .. }" + ); + assert_eq!( + items.by_ref().map(|item| item.expect("valid item")).collect::>(), + [b"alpha".as_slice(), b"\"bravo,charlie\"".as_slice(), b"delta".as_slice(),] + ); + assert!(items.next().is_none()); + + let semicolons = FieldLines::single(&FieldName::ContentType, b" alpha ; ; \"b;\\\"c\" ; "); + assert_eq!( + semicolons + .semicolon_items() + .map(|item| item.expect("valid item")) + .collect::>(), + [b"alpha".as_slice(), b"".as_slice(), b"\"b;\\\"c\"".as_slice(), b"".as_slice(),] + ); + + let unterminated = FieldLines::single(&FieldName::Vary, b"alpha,\"bravo"); + let error = unterminated + .comma_items() + .nth(1) + .expect("error item") + .expect_err("unterminated quote fails"); + assert_eq!(error.kind(), DecodeErrorKind::UnterminatedQuote); + assert_eq!(error.value_index(), Some(0)); + + let escaped = FieldLines::single(&FieldName::Vary, b"\"alpha\\"); + assert_eq!( + escaped + .comma_items() + .next() + .expect("error item") + .expect_err("trailing escape fails") + .kind(), + DecodeErrorKind::UnterminatedQuote + ); + + let split_quote = [FieldValue::from_static("\"alpha"), FieldValue::from_static("bravo\"")]; + let values = FieldLines::from_slice(&FieldName::Vary, &split_quote).expect("values"); + let error = values + .comma_items() + .next() + .expect("error item") + .expect_err("quoted strings cannot span physical lines"); + assert_eq!(error.kind(), DecodeErrorKind::UnterminatedQuote); + assert_eq!(error.value_index(), Some(0)); + + let obs_text = [FieldValue::from_bytes(b"\xff").expect("obs-text is a valid field value")]; + let values = FieldLines::from_slice(&FieldName::Vary, &obs_text).expect("values"); + assert_eq!( + values.comma_items().next().expect("one item").expect("obs-text item is yielded"), + b"\xff".as_slice() + ); + } +} diff --git a/crates/http_headers/src/source/field_lines.rs b/crates/http_headers/src/source/field_lines.rs new file mode 100644 index 000000000..5447360e2 --- /dev/null +++ b/crates/http_headers/src/source/field_lines.rs @@ -0,0 +1,1047 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Field lines supplied by a field source for one field name. + +use std::iter::FusedIterator; +use std::{fmt, mem, slice}; + +use crate::source::{DelimitedItems, update_list_item_count}; +use crate::{DecodeError, DecodeErrorKind, FieldName, FieldValue, FieldValueRef}; + +/// Maximum total field-value bytes accepted from a custom [`FieldSource`](crate::source::FieldSource). +/// +/// The limit is enforced before typed decoding accepts custom-source bytes. +/// The validated `http::HeaderMap` adapter is exempt, so callers using it must +/// enforce suitable aggregate header limits at the transport or server layer. +/// Exceeding this budget returns [`DecodeErrorKind::SourceLimitExceeded`]. +pub const MAX_CUSTOM_FIELD_BYTES: usize = 64 * 1024; + +/// Maximum field lines accepted for one name from a custom [`FieldSource`](crate::source::FieldSource). +/// +/// The limit is enforced before typed decoding accepts custom-source lines. +/// The validated `http::HeaderMap` adapter is exempt, so callers using it must +/// enforce suitable aggregate header limits at the transport or server layer. +/// Exceeding this budget returns [`DecodeErrorKind::SourceLimitExceeded`]. +pub const MAX_CUSTOM_FIELD_LINES: usize = 128; + +/// Maximum parsed list items accepted during one custom-source field decode. +/// +/// [`DelimitedItems`] and specialized list parsers enforce this limit before +/// accepting a further item. The validated `http::HeaderMap` adapter is exempt, +/// so callers using it must enforce suitable aggregate header limits at the +/// transport or server layer. +/// Exceeding this budget returns [`DecodeErrorKind::SourceLimitExceeded`]. +pub const MAX_CUSTOM_LIST_ITEMS: usize = 1_024; + +/// The storage a [`FieldLines`] iterates. +/// +/// The set of representations is closed so that every decoder dispatches +/// statically: a field source hands out one of these shapes, and no decoding +/// path ever goes through a trait object. +enum Repr<'a> { + /// Exactly one field line, borrowed from anywhere. + Single(&'a [u8]), + /// One [`FieldValue`] per field line, stored contiguously. + Slice(&'a [FieldValue]), + /// Field lines borrowing arbitrary backing storage. + Borrowed(&'a [FieldValueRef<'a>]), + /// The field lines an `http::HeaderMap` stores for one name. + #[cfg(feature = "http")] + Http(http::header::GetAll<'a, http::HeaderValue>), +} + +/// The raw field lines stored under one field name. +/// +/// A [`crate::source::FieldSource`] returns this type for a present field. An +/// instance always contains at least one field line; absence is represented by +/// `None`. Line order is preserved, a zero-length line remains a present line, +/// and the lines can be iterated repeatedly. +/// +/// The associated name is a static descriptor, independent of the lifetime of +/// the field bytes. Locally constructed runtime names cannot be passed to the +/// constructors; see [`crate::source::FieldSource`]. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::FieldName; +/// use http_headers::source::FieldLines; +/// +/// let lines = FieldLines::single(&FieldName::UserAgent, b"client/1"); +/// assert_eq!(lines.len(), 1); +/// assert_eq!(lines.exactly_one()?.as_bytes(), b"client/1"); +/// # Ok::<(), http_headers::DecodeError>(()) +/// ``` +pub struct FieldLines<'a> { + name: &'static FieldName, + repr: Repr<'a>, +} + +const _: [(); mem::size_of::>()] = [(); mem::size_of::>>()]; + +impl fmt::Debug for FieldLines<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("FieldLines") + .field("name", &self.name) + .field("line_count", &self.len()) + .finish() + } +} + +impl<'a> FieldLines<'a> { + /// Creates field lines holding exactly one line. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// use http_headers::source::FieldLines; + /// + /// assert_eq!( + /// FieldLines::single(&FieldName::Accept, b"text/html").len(), + /// 1 + /// ); + /// ``` + #[must_use] + #[inline] + pub const fn single(name: &'static FieldName, value: &'a [u8]) -> Self { + Self { + name, + repr: Repr::Single(value), + } + } + + /// Creates field lines over a contiguous slice of [`FieldValue`]s. + /// + /// Returns `None` when the slice contains no field lines. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::source::FieldLines; + /// use http_headers::{FieldName, FieldValue}; + /// + /// let stored = [FieldValue::from_static("text/html")]; + /// let lines = FieldLines::from_slice(&FieldName::Accept, &stored); + /// assert_eq!(lines.map(|lines| lines.len()), Some(1)); + /// ``` + #[must_use] + #[inline] + pub const fn from_slice(name: &'static FieldName, values: &'a [FieldValue]) -> Option { + if values.is_empty() { + None + } else { + Some(Self { + name, + repr: Repr::Slice(values), + }) + } + } + + /// Creates field lines from borrowed lines, one per line. + /// + /// Sensitivity markers on the supplied lines are preserved. Returns + /// `None` when the slice contains no field lines. + #[must_use] + #[inline] + pub const fn from_borrowed(name: &'static FieldName, values: &'a [FieldValueRef<'a>]) -> Option { + if values.is_empty() { + None + } else { + Some(Self { + name, + repr: Repr::Borrowed(values), + }) + } + } + + /// Creates field lines over what an `http::HeaderMap` stores. + /// + /// Returns `None` when the map contains no field line for the name. + /// + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # fn main() { + /// use http_headers::FieldName; + /// use http_headers::source::FieldLines; + /// + /// let mut map = http::HeaderMap::new(); + /// map.append( + /// http::header::ACCEPT, + /// http::HeaderValue::from_static("text/html"), + /// ); + /// let lines = FieldLines::from_http(&FieldName::Accept, map.get_all(http::header::ACCEPT)); + /// assert_eq!(lines.map(|lines| lines.len()), Some(1)); + /// # } + /// # #[cfg(not(feature = "http"))] + /// # fn main() {} + /// ``` + #[cfg(feature = "http")] + #[must_use] + #[inline] + pub fn from_http(name: &'static FieldName, values: http::header::GetAll<'a, http::HeaderValue>) -> Option { + if values.iter().next().is_none() { + None + } else { + Some(Self { + name, + repr: Repr::Http(values), + }) + } + } + + /// Returns the associated field name. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// use http_headers::source::FieldLines; + /// + /// let lines = FieldLines::single(&FieldName::Accept, b""); + /// assert_eq!(lines.name().as_str(), "accept"); + /// ``` + #[must_use] + #[inline] + pub const fn name(&self) -> &'static FieldName { + self.name + } + + /// Returns the value of the only field line, or a cardinality error. + /// + /// # Errors + /// + /// Returns an error if there is no field line or more than one. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// use http_headers::source::FieldLines; + /// + /// let lines = FieldLines::single(&FieldName::UserAgent, b"client/1"); + /// assert_eq!(lines.exactly_one()?.as_bytes(), b"client/1"); + /// # Ok::<(), http_headers::DecodeError>(()) + /// ``` + #[expect( + clippy::inline_always, + reason = "typed singleton decoding must eliminate the known source representation and intermediate result" + )] + #[inline(always)] + pub fn exactly_one(self) -> Result, DecodeError> { + // Specialized per representation rather than routed through + // `repeated()`: the cardinality of a slice is already known, and + // `Repr::Single` is by construction exactly one line. + let multiple = || DecodeError::new(self.name, DecodeErrorKind::UnexpectedMultipleValues); + let first = match &self.repr { + Repr::Single(bytes) => return Ok(FieldValueRef::new(bytes)), + Repr::Slice(values) => match values { + [only] => Some(only.as_field_value_ref()), + [] => None, + _ => return Err(multiple()), + }, + Repr::Borrowed(values) => match values { + [only] => Some(*only), + [] => None, + _ => return Err(multiple()), + }, + #[cfg(feature = "http")] + Repr::Http(values) => { + let mut lines = values.iter(); + let first = lines.next().map(FieldValueRef::from); + if first.is_some() && lines.next().is_some() { + return Err(multiple()); + } + first + } + }; + first.ok_or_else(|| DecodeError::new(self.name, DecodeErrorKind::MissingValue)) + } + + /// Returns the value of the only field line as a representation-aware + /// owned value, or a cardinality error. + /// + /// This never builds an intermediate borrowed [`FieldValueRef`]: each + /// [`Repr`] variant is specialized directly to the cheapest owned value + /// it can produce. Short lines are stored inline, so the common case + /// allocates nothing at all; longer ones reuse the source's shared + /// storage where it exists — [`Repr::Slice`] clones the stored + /// [`FieldValue`] — and copy otherwise. [`Repr::Single`] and + /// [`Repr::Borrowed`] are also the only representations that revalidate, + /// because they are the only ones whose bytes did not arrive inside an + /// already-validated value. + /// + /// # Errors + /// + /// Returns an error if there is no field line, if there is more than one, + /// if a source handed out bytes that are not a valid field value, or if a + /// custom source exceeds [`MAX_CUSTOM_FIELD_BYTES`] or + /// [`MAX_CUSTOM_FIELD_LINES`] (reported as [`DecodeErrorKind::SourceLimitExceeded`]). + #[inline] + pub(crate) fn exactly_one_owned(&self) -> Result { + self.validate_custom_source()?; + match &self.repr { + Repr::Single(bytes) => Ok(FieldValueRef::new(bytes).to_validated_field_value()), + Repr::Slice(values) => match values { + [value] => Ok(value.clone()), + _ => Err(DecodeError::new(self.name, DecodeErrorKind::UnexpectedMultipleValues)), + }, + Repr::Borrowed(values) => match values { + [value] => Ok(value.to_validated_field_value()), + _ => Err(DecodeError::new(self.name, DecodeErrorKind::UnexpectedMultipleValues)), + }, + #[cfg(feature = "http")] + Repr::Http(values) => { + let mut values = values.iter(); + let first = values + .next() + .ok_or_else(|| DecodeError::new(self.name, DecodeErrorKind::MissingValue))?; + if values.next().is_some() { + return Err(DecodeError::new(self.name, DecodeErrorKind::UnexpectedMultipleValues)); + } + Ok(FieldValue::from(first)) + } + } + } + + /// Iterates the raw field lines in insertion order. + /// + /// This is equivalent to [`Self::repeated`] and iteration over `&FieldLines`. + /// Line boundaries and sensitivity markers are preserved without parsing + /// or validating the bytes. + /// + /// # Examples + /// + /// ``` + /// use http_headers::FieldName; + /// use http_headers::source::FieldLines; + /// + /// let lines = FieldLines::single(&FieldName::SetCookie, b"a=1"); + /// assert_eq!(lines.iter().next().unwrap().as_bytes(), b"a=1"); + /// assert_eq!(lines.iter().count(), lines.repeated().count()); + /// ``` + #[must_use] + #[inline] + pub fn iter(&self) -> FieldLinesIter<'a> { + self.repeated() + } + + /// Reiterates the field lines in insertion order. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// use http_headers::source::FieldLines; + /// + /// let lines = FieldLines::single(&FieldName::SetCookie, b"a=1"); + /// assert_eq!(lines.repeated().count(), 1); + /// for line in &lines { + /// assert_eq!(line.as_bytes(), b"a=1"); + /// } + /// ``` + #[must_use] + #[inline] + pub fn repeated(&self) -> FieldLinesIter<'a> { + FieldLinesIter { + repr: match &self.repr { + Repr::Single(value) => LinesRepr::Single(Some(value)), + Repr::Slice(values) => LinesRepr::Slice(values.iter()), + Repr::Borrowed(values) => LinesRepr::Borrowed(values.iter()), + #[cfg(feature = "http")] + Repr::Http(values) => LinesRepr::Http(values.iter()), + }, + } + } + + /// Reiterates the field lines in insertion order, pairing each borrowed + /// view with a representation-aware owned clone. + /// + /// This is what an owned decoder should walk instead of calling + /// [`FieldValueRef::try_to_field_value`] on the output of [`Self::repeated`]: + /// it keeps [`Repr::Slice`] from paying for a byte copy that a cheap + /// clone would have avoided. + /// + /// # Errors + /// + /// Returns an invalid-syntax error when raw single or borrowed lines + /// contain bytes that cannot form an HTTP field value. Returns + /// [`DecodeErrorKind::SourceLimitExceeded`] when a custom source exceeds + /// [`MAX_CUSTOM_FIELD_BYTES`] or [`MAX_CUSTOM_FIELD_LINES`]. + #[inline] + #[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-cors", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-websocket", + ))] + pub(crate) fn repeated_owned(&self) -> Result, DecodeError> { + self.validate_custom_source()?; + Ok(FieldLinesOwnedIter { + repr: match &self.repr { + Repr::Single(value) => OwnedLinesRepr::Single(Some(value)), + Repr::Slice(values) => OwnedLinesRepr::Slice(values.iter()), + Repr::Borrowed(values) => OwnedLinesRepr::Borrowed(values.iter()), + #[cfg(feature = "http")] + Repr::Http(values) => OwnedLinesRepr::Http(values.iter()), + }, + }) + } + + /// Returns the number of field lines, which is always at least one. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// use http_headers::source::FieldLines; + /// + /// assert_eq!(FieldLines::single(&FieldName::SetCookie, b"a=1").len(), 1); + /// ``` + #[must_use] + #[inline] + #[expect( + clippy::len_without_is_empty, + reason = "FieldLines is non-empty by construction; absence is represented by Option" + )] + pub fn len(&self) -> usize { + match &self.repr { + Repr::Single(_) => 1, + Repr::Slice(values) => values.len(), + Repr::Borrowed(values) => values.len(), + #[cfg(feature = "http")] + Repr::Http(values) => values.iter().count(), + } + } + + /// Iterates comma-delimited items across all field lines. + /// + /// Items are yielded in order across field-line boundaries without + /// materializing a combined value. Each physical line is scanned + /// independently, so a quoted string or escape cannot span lines. + /// + /// This follows the HTTP list-rule convention: empty items produced by + /// consecutive or trailing commas are silently skipped rather than + /// yielded. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// use http_headers::source::FieldLines; + /// + /// let lines = FieldLines::single(&FieldName::Vary, b"accept, , origin"); + /// assert_eq!(lines.comma_items().count(), 2); + /// ``` + #[must_use] + #[inline] + pub fn comma_items(&self) -> DelimitedItems<'a> { + DelimitedItems::new(self, b',') + } + + /// Iterates comma members after source preflight has already succeeded. + /// + /// The caller must have established that `validate_custom_source()` succeeds + /// for these exact retained lines, not for another call to `FieldSource::lines`. + /// Only source bounds and field-value preflight are omitted: quoting and + /// item limits remain checked. + /// Debug builds recheck the caller's precondition. + #[cfg(any(test, feature = "headers-negotiation"))] + pub(crate) fn validated_comma_items(&self) -> DelimitedItems<'a> { + #[cfg(debug_assertions)] + self.validate_custom_source() + .expect("the caller must retain the exact lines whose source preflight is known to succeed"); + DelimitedItems::from_validated_source(self, b',') + } + + /// Iterates semicolon-delimited items across all field lines. + /// + /// Unlike [`Self::comma_items`], empty items are yielded rather than + /// skipped, allowing the caller to decide whether they are valid for the + /// field being parsed. + /// + /// # Examples + /// + /// ```rust + /// use http_headers::FieldName; + /// use http_headers::source::FieldLines; + /// + /// let lines = FieldLines::single(&FieldName::ContentType, b"text/html; charset=utf-8"); + /// assert_eq!(lines.semicolon_items().count(), 2); + /// ``` + #[must_use] + pub fn semicolon_items(&self) -> DelimitedItems<'a> { + DelimitedItems::new(self, b';') + } + + #[cfg(feature = "http")] + pub(crate) const fn has_custom_source_limits(&self) -> bool { + !matches!(&self.repr, Repr::Http(_)) + } + + #[cfg(not(feature = "http"))] + pub(crate) const fn has_custom_source_limits(&self) -> bool { + matches!(&self.repr, Repr::Single(_) | Repr::Slice(_) | Repr::Borrowed(_)) + } + + #[inline] + pub(crate) fn validate_custom_source_bounds(&self) -> Result<(), DecodeError> { + if !self.has_custom_source_limits() { + return Ok(()); + } + self.validate_bounded_source() + } + + fn validate_bounded_source(&self) -> Result<(), DecodeError> { + let limit = || DecodeError::new(self.name, DecodeErrorKind::SourceLimitExceeded); + if self.len() > MAX_CUSTOM_FIELD_LINES { + return Err(limit()); + } + + let validate_bytes = matches!(&self.repr, Repr::Single(_) | Repr::Borrowed(_)); + let mut total_bytes = 0_usize; + for value in self.repeated() { + let bytes = value.as_bytes(); + total_bytes = total_bytes.checked_add(bytes.len()).ok_or_else(limit)?; + if total_bytes > MAX_CUSTOM_FIELD_BYTES { + return Err(limit()); + } + if validate_bytes && !crate::validate::field_value(bytes) { + return Err(DecodeError::new(self.name, DecodeErrorKind::InvalidSyntax)); + } + } + Ok(()) + } + + #[inline] + pub(crate) fn validate_custom_source(&self) -> Result<(), DecodeError> { + self.validate_custom_source_bounds() + } + + pub(crate) fn validate_list_item_limit(&self, delimiter: u8, skip_empty: bool) -> Result<(), DecodeError> { + self.validate_list_item_limit_with(delimiter, skip_empty, true) + } + + #[cfg(any(test, feature = "headers-conditional"))] + pub(crate) fn validate_entity_tag_item_limit(&self) -> Result<(), DecodeError> { + self.validate_list_item_limit_with(b',', true, false) + } + + #[inline] + fn validate_list_item_limit_with(&self, delimiter: u8, skip_empty: bool, backslash_escapes: bool) -> Result<(), DecodeError> { + if !self.has_custom_source_limits() { + return Ok(()); + } + self.validate_bounded_list(delimiter, skip_empty, backslash_escapes) + } + + fn validate_bounded_list(&self, delimiter: u8, skip_empty: bool, backslash_escapes: bool) -> Result<(), DecodeError> { + let limit = || DecodeError::new(self.name, DecodeErrorKind::SourceLimitExceeded); + if self.len() > MAX_CUSTOM_FIELD_LINES { + return Err(limit()); + } + + let validate_bytes = matches!(&self.repr, Repr::Single(_) | Repr::Borrowed(_)); + let mut total_bytes = 0_usize; + let mut item_count = 0_usize; + for value in self.repeated() { + let bytes = value.as_bytes(); + total_bytes = total_bytes.checked_add(bytes.len()).ok_or_else(limit)?; + if total_bytes > MAX_CUSTOM_FIELD_BYTES { + return Err(limit()); + } + if validate_bytes && !crate::validate::field_value(bytes) { + return Err(DecodeError::new(self.name, DecodeErrorKind::InvalidSyntax)); + } + update_list_item_count(self.name, bytes, delimiter, skip_empty, backslash_escapes, &mut item_count)?; + } + Ok(()) + } +} + +impl<'a> IntoIterator for &FieldLines<'a> { + type Item = FieldValueRef<'a>; + type IntoIter = FieldLinesIter<'a>; + + #[inline] + fn into_iter(self) -> Self::IntoIter { + self.iter() + } +} + +/// The storage a [`FieldLinesIter`] iterator walks. +enum LinesRepr<'a> { + /// At most one remaining field line. + Single(Option<&'a [u8]>), + /// Remaining [`FieldValue`] lines. + Slice(slice::Iter<'a, FieldValue>), + /// Remaining field lines borrowing arbitrary backing storage. + Borrowed(slice::Iter<'a, FieldValueRef<'a>>), + /// Remaining `http` field lines. + #[cfg(feature = "http")] + Http(http::header::ValueIter<'a, http::HeaderValue>), +} + +/// An iterator over the field lines stored under one field name. +/// +/// Yields one [`FieldValueRef`] per field line, in insertion order: the +/// `field-value` each line carries, never the comma-joined combination of +/// them. Field-line boundaries are preserved. +/// +/// # Examples +/// +/// ```rust +/// use http_headers::FieldName; +/// use http_headers::source::{FieldLines, FieldLinesIter}; +/// +/// let lines = FieldLines::single(&FieldName::SetCookie, b"a=1"); +/// let mut iter: FieldLinesIter<'_> = lines.repeated(); +/// assert_eq!(iter.size_hint(), (1, Some(1))); +/// assert_eq!( +/// iter.next().map(|line| line.as_bytes()), +/// Some(b"a=1".as_slice()) +/// ); +/// assert!(iter.next().is_none()); +/// ``` +pub struct FieldLinesIter<'a> { + repr: LinesRepr<'a>, +} + +impl fmt::Debug for FieldLinesIter<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("FieldLinesIter").finish_non_exhaustive() + } +} + +/// Re-borrows a stored value for [`LinesRepr::Slice`]. +/// +/// Kept out of line so that the three remaining arms stay small enough for +/// `next` to inline into a decoder's field-line loop. Inlining the whole +/// `FieldValue` representation here instead costs every other source a call +/// per line, which outweighs the call this leaves on the slice path. +#[inline(never)] +fn slice_line(value: &FieldValue) -> FieldValueRef<'_> { + value.as_field_value_ref() +} + +impl<'a> Iterator for FieldLinesIter<'a> { + type Item = FieldValueRef<'a>; + + #[inline] + fn next(&mut self) -> Option { + match &mut self.repr { + LinesRepr::Single(value) => value.take().map(FieldValueRef::new), + LinesRepr::Slice(values) => values.next().map(slice_line), + LinesRepr::Borrowed(values) => values.next().copied(), + #[cfg(feature = "http")] + LinesRepr::Http(values) => values.next().map(FieldValueRef::from), + } + } + + fn size_hint(&self) -> (usize, Option) { + match &self.repr { + LinesRepr::Single(value) => { + let length = usize::from(value.is_some()); + (length, Some(length)) + } + LinesRepr::Slice(values) => values.size_hint(), + LinesRepr::Borrowed(values) => values.size_hint(), + #[cfg(feature = "http")] + LinesRepr::Http(values) => values.size_hint(), + } + } +} + +impl FusedIterator for FieldLinesIter<'_> {} + +/// The storage a [`FieldLinesOwnedIter`] iterator walks. +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-cors", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-websocket", +))] +enum OwnedLinesRepr<'a> { + /// At most one remaining field line. + Single(Option<&'a [u8]>), + /// Remaining [`FieldValue`] lines. + Slice(slice::Iter<'a, FieldValue>), + /// Remaining field lines borrowing arbitrary backing storage. + Borrowed(slice::Iter<'a, FieldValueRef<'a>>), + /// Remaining `http` field lines. + #[cfg(feature = "http")] + Http(http::header::ValueIter<'a, http::HeaderValue>), +} + +/// An iterator pairing each field line stored under one field name with a +/// representation-aware owned clone. +/// +/// The borrowed half of each item is what a decoder validates or parses; the +/// owned half is what it stores. Cloning the owned half costs no byte copy for +/// [`Repr::Slice`], which already holds shareable storage. Every other +/// representation copies: [`Repr::Single`] and [`Repr::Borrowed`] have only +/// bytes to work from, and with the `http` feature [`Repr::Http`] can reach +/// `HeaderValue` only through `as_bytes`, since `http` keeps the `Bytes` +/// behind it private. +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-cors", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-websocket", +))] +pub(crate) struct FieldLinesOwnedIter<'a> { + repr: OwnedLinesRepr<'a>, +} + +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-cors", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-websocket", +))] +impl fmt::Debug for FieldLinesOwnedIter<'_> { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("FieldLinesOwnedIter").finish_non_exhaustive() + } +} + +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-cors", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-websocket", +))] +impl<'a> Iterator for FieldLinesOwnedIter<'a> { + type Item = (FieldValueRef<'a>, FieldValue); + + #[inline] + fn next(&mut self) -> Option { + match &mut self.repr { + OwnedLinesRepr::Single(value) => value.take().map(|bytes| { + let value = FieldValueRef::new(bytes); + (value, value.to_validated_field_value()) + }), + OwnedLinesRepr::Slice(values) => values.next().map(|value| (value.as_field_value_ref(), value.clone())), + OwnedLinesRepr::Borrowed(values) => values.next().copied().map(|value| (value, value.to_validated_field_value())), + #[cfg(feature = "http")] + OwnedLinesRepr::Http(values) => values.next().map(|value| (FieldValueRef::from(value), FieldValue::from(value))), + } + } + + fn size_hint(&self) -> (usize, Option) { + match &self.repr { + OwnedLinesRepr::Single(value) => { + let length = usize::from(value.is_some()); + (length, Some(length)) + } + OwnedLinesRepr::Slice(values) => values.size_hint(), + OwnedLinesRepr::Borrowed(values) => values.size_hint(), + #[cfg(feature = "http")] + OwnedLinesRepr::Http(values) => values.size_hint(), + } + } +} + +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-conditional", + feature = "headers-cors", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-set-cookie", + feature = "headers-websocket", +))] +impl FusedIterator for FieldLinesOwnedIter<'_> {} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::{FieldLines, Repr, update_list_item_count}; + use crate::headers::SetCookie; + use crate::source::FieldSource; + use crate::{DecodeErrorKind, Field, FieldName, FieldValue, FieldValueRef}; + + #[test] + fn set_cookie_handles_an_internal_empty_field_lines_representation() { + struct EmptySource; + + impl FieldSource for EmptySource { + fn lines(&self, name: &'static FieldName) -> Option> { + Some(FieldLines { + name, + repr: Repr::Slice(&[]), + }) + } + } + + let source = EmptySource; + let view = ::view(&source).unwrap().unwrap(); + assert_eq!(view.len(), 0); + assert_eq!(view.iter().count(), 0); + let owned = ::owned(&source).unwrap().unwrap(); + assert_eq!(owned.len(), 0); + assert!(owned.is_empty()); + assert_eq!(owned.iter().count(), 0); + } + + #[test] + fn exactly_one_specializes_every_representation() { + assert_eq!( + FieldLines::single(&FieldName::Accept, b"text/html") + .exactly_one() + .expect("a single line is exactly one value") + .as_bytes(), + b"text/html" + ); + + let borrowed = [FieldValueRef::new(b"text/html"), FieldValueRef::new(b"application/json")]; + assert_eq!( + FieldLines::from_borrowed(&FieldName::Accept, &borrowed) + .expect("a non-empty slice makes a set") + .exactly_one() + .expect_err("two borrowed lines are not exactly one") + .kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + + let empty = FieldLines { + name: &FieldName::Accept, + repr: Repr::Borrowed(&[]), + }; + assert_eq!( + empty + .exactly_one() + .expect_err("internal empty borrowed representation reports missing") + .kind(), + DecodeErrorKind::MissingValue + ); + } + + #[test] + fn every_storage_representation_preserves_cardinality_order_and_owned_values() { + assert!(FieldLines::from_slice(&FieldName::SetCookie, &[]).is_none()); + assert!(FieldLines::from_borrowed(&FieldName::SetCookie, &[]).is_none()); + + let single = FieldLines::single(&FieldName::SetCookie, b"a=1"); + assert_eq!(single.name(), &FieldName::SetCookie); + assert_eq!(single.len(), 1); + assert_eq!(format!("{single:?}"), "FieldLines { name: \"set-cookie\", line_count: 1 }"); + let mut lines = single.repeated(); + assert_eq!(lines.size_hint(), (1, Some(1))); + assert_eq!(format!("{lines:?}"), "FieldLinesIter { .. }"); + assert_eq!(lines.next().expect("single line"), "a=1"); + assert!(lines.next().is_none()); + assert_eq!(single.exactly_one_owned().expect("single owned value").as_bytes(), b"a=1"); + let mut owned_lines = single.repeated_owned().expect("valid raw line"); + assert_eq!(owned_lines.size_hint(), (1, Some(1))); + assert_eq!(format!("{owned_lines:?}"), "FieldLinesOwnedIter { .. }"); + let (borrowed, owned) = owned_lines.next().expect("single owned line"); + assert_eq!(borrowed, "a=1"); + assert_eq!(owned, "a=1"); + + let stored = [FieldValue::from_static("a=1"), FieldValue::from_static("b=2")]; + let slice = FieldLines::from_slice(&FieldName::SetCookie, &stored).expect("non-empty slice"); + assert_eq!(slice.len(), 2); + assert_eq!( + slice.repeated().map(FieldValueRef::as_bytes).collect::>(), + [b"a=1".as_slice(), b"b=2".as_slice()] + ); + assert_eq!( + slice + .repeated_owned() + .expect("stored values are valid") + .map(|(borrowed, owned)| (borrowed.as_bytes(), owned)) + .collect::>(), + [ + (b"a=1".as_slice(), FieldValue::from_static("a=1")), + (b"b=2".as_slice(), FieldValue::from_static("b=2")), + ] + ); + assert_eq!( + slice.exactly_one_owned().expect_err("multiple owned values fail").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + let slice = FieldLines::from_slice(&FieldName::SetCookie, &stored).expect("non-empty slice"); + assert_eq!( + slice.exactly_one().expect_err("multiple values fail").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + + let refs = [FieldValueRef::new(b"a=1"), FieldValueRef::new(b"b=2")]; + let borrowed = FieldLines::from_borrowed(&FieldName::SetCookie, &refs).expect("non-empty refs"); + assert_eq!(borrowed.len(), 2); + let mut borrowed_lines = borrowed.repeated(); + assert_eq!(borrowed_lines.size_hint(), (2, Some(2))); + assert_eq!(borrowed_lines.by_ref().count(), 2); + let mut borrowed_owned = borrowed.repeated_owned().expect("valid borrowed lines"); + assert_eq!(borrowed_owned.size_hint(), (2, Some(2))); + assert_eq!(borrowed_owned.by_ref().count(), 2); + assert_eq!( + borrowed.exactly_one_owned().expect_err("multiple borrowed values fail").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + let one_ref = [FieldValueRef::new(b"a=1")]; + let one_borrowed = FieldLines::from_borrowed(&FieldName::SetCookie, &one_ref).expect("one ref"); + assert_eq!(one_borrowed.exactly_one_owned().expect("one borrowed owned value"), "a=1"); + + let empty = FieldLines { + name: &FieldName::SetCookie, + repr: Repr::Slice(&[]), + }; + assert_eq!( + empty + .exactly_one() + .expect_err("internal empty representation reports missing") + .kind(), + DecodeErrorKind::MissingValue + ); + } + + #[test] + fn unvalidated_source_bytes_report_a_decode_error_instead_of_panicking() { + let single = FieldLines::single(&FieldName::SetCookie, b"bad\nvalue"); + assert_eq!( + single + .exactly_one_owned() + .expect_err("a source must not hand out invalid bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + let refs = [FieldValueRef::new(b"bad\nvalue")]; + let borrowed = FieldLines::from_borrowed(&FieldName::SetCookie, &refs).expect("one ref"); + assert_eq!( + borrowed + .exactly_one_owned() + .expect_err("a source must not hand out invalid bytes") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + assert!(single.validate_custom_source_bounds().is_err()); + assert!(borrowed.validate_custom_source_bounds().is_err()); + } + + #[test] + fn list_budget_and_counter_overflow_rejections_are_admission_errors() { + for initial in [1_024, usize::MAX] { + for bytes in [b"one,two".as_slice(), b"one"] { + let mut item_count = initial; + assert_eq!( + update_list_item_count(&FieldName::Vary, bytes, b',', true, true, &mut item_count) + .unwrap_err() + .kind(), + DecodeErrorKind::SourceLimitExceeded + ); + } + } + } + + #[test] + fn borrowed_and_owned_iteration_carry_sensitivity_from_stored_values() { + let stored = [FieldValue::from_static("credential").with_sensitive(true)]; + let values = FieldLines::from_slice(&FieldName::Authorization, &stored).expect("non-empty slice"); + + let borrowed = values.repeated().next().expect("one line"); + assert!(borrowed.is_sensitive()); + assert_eq!(format!("{borrowed:?}"), "FieldValueRef(Sensitive)"); + + let (borrowed, owned) = values.repeated_owned().expect("stored value is valid").next().expect("one line"); + assert!(borrowed.is_sensitive()); + assert!(owned.is_sensitive()); + assert!(values.exactly_one_owned().expect("one value").is_sensitive()); + } + + #[cfg(feature = "http")] + #[test] + fn http_storage_representation_handles_empty_single_and_multiple_values() { + let mut map = http::HeaderMap::new(); + assert!(FieldLines::from_http(&FieldName::SetCookie, map.get_all(http::header::SET_COOKIE)).is_none()); + let empty = FieldLines { + name: &FieldName::SetCookie, + repr: Repr::Http(map.get_all(http::header::SET_COOKIE)), + }; + assert_eq!( + empty + .exactly_one_owned() + .expect_err("internal empty http representation reports missing") + .kind(), + DecodeErrorKind::MissingValue + ); + + map.append(http::header::SET_COOKIE, http::HeaderValue::from_static("a=1")); + let one = FieldLines::from_http(&FieldName::SetCookie, map.get_all(http::header::SET_COOKIE)).expect("one http value"); + assert_eq!(one.len(), 1); + let mut lines = one.repeated(); + assert_eq!(lines.size_hint(), (1, Some(1))); + assert_eq!(lines.next().expect("line"), "a=1"); + assert_eq!(one.exactly_one_owned().expect("one owned http value").as_bytes(), b"a=1"); + let mut owned_lines = one.repeated_owned().expect("http values are valid"); + assert_eq!(owned_lines.size_hint(), (1, Some(1))); + assert_eq!(owned_lines.next().expect("owned line").1, "a=1"); + + map.append(http::header::SET_COOKIE, http::HeaderValue::from_static("b=2")); + let multiple = FieldLines::from_http(&FieldName::SetCookie, map.get_all(http::header::SET_COOKIE)).expect("multiple http values"); + assert_eq!(multiple.len(), 2); + assert_eq!(multiple.repeated().count(), 2); + assert_eq!(multiple.repeated_owned().expect("http values are valid").count(), 2); + assert_eq!( + multiple.exactly_one_owned().expect_err("multiple http values fail").kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + } + + #[cfg(feature = "http")] + #[test] + fn http_storage_carries_sensitivity_into_borrowed_and_owned_values() { + let mut map = http::HeaderMap::new(); + let mut value = http::HeaderValue::from_static("Bearer credential"); + value.set_sensitive(true); + map.append(http::header::AUTHORIZATION, value); + + let values = FieldLines::from_http(&FieldName::Authorization, map.get_all(http::header::AUTHORIZATION)).expect("one http value"); + + let borrowed = values.repeated().next().expect("one line"); + assert!(borrowed.is_sensitive()); + let rendered = format!("{borrowed:?}"); + assert!( + !rendered.contains("credential"), + "a borrowed view must never render credential bytes: {rendered}" + ); + + let (borrowed, owned) = values.repeated_owned().expect("http values are valid").next().expect("one line"); + assert!(borrowed.is_sensitive()); + assert!(owned.is_sensitive()); + assert!(values.exactly_one_owned().expect("one value").is_sensitive()); + } +} diff --git a/crates/http_headers/src/source/field_source.rs b/crates/http_headers/src/source/field_source.rs new file mode 100644 index 000000000..b2e385300 --- /dev/null +++ b/crates/http_headers/src/source/field_source.rs @@ -0,0 +1,86 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! The trait a field container implements to supply field lines. + +use crate::FieldName; +use crate::source::FieldLines; + +/// A container that supplies stored HTTP field lines. +/// +/// This trait is implemented for individual providers of HTTP fields and is +/// how this crate acquires fields to parse. +/// +/// You can use the `http` cargo feature to get an implementation of this trait for the common +/// [`HeaderMap`](https://docs.rs/http/latest/http/header/struct.HeaderMap.html) type. +/// +/// Names at this boundary must be static descriptors. Custom descriptors can +/// use a `static LazyLock`; locally constructed runtime names cannot +/// be passed to this trait. Use the container's native API for dynamic lookup. +/// [`FieldName::try_from_bytes`] remains available for validating, comparing, +/// and converting runtime names. +/// +/// # Examples +/// +/// ```rust +/// # #[cfg(feature = "http")] +/// # { +/// use http::HeaderMap; +/// use http_headers::FieldName; +/// use http_headers::source::FieldSource; +/// +/// assert!(FieldSource::lines(&HeaderMap::new(), &FieldName::UserAgent).is_none()); +/// # } +/// ``` +pub trait FieldSource { + /// Returns whether any field line is stored under `name`. + #[inline] + fn contains(&self, name: &'static FieldName) -> bool { + self.lines(name).is_some() + } + + /// Returns the field lines stored under `name`. + /// + /// `None` means the field is absent. `Some` always contains at least one + /// raw field line; a present zero-length field line is therefore distinct + /// from absence. + /// + /// Borrowed and owned decoding from custom sources accept at most + /// [`MAX_CUSTOM_FIELD_BYTES`](crate::source::MAX_CUSTOM_FIELD_BYTES) total + /// bytes and [`MAX_CUSTOM_FIELD_LINES`](crate::source::MAX_CUSTOM_FIELD_LINES) + /// lines for one name. Delimited parsing accepts at most + /// [`MAX_CUSTOM_LIST_ITEMS`](crate::source::MAX_CUSTOM_LIST_ITEMS) items. + /// Exceeding a limit returns + /// [`DecodeErrorKind::SourceLimitExceeded`](crate::DecodeErrorKind::SourceLimitExceeded). + /// The native `http::HeaderMap` representation is exempt; a custom source + /// backed by validated [`FieldValue`](crate::FieldValue) values is not. + /// + /// # Examples + /// + /// ```rust + /// # #[cfg(feature = "http")] + /// # { + /// use http::HeaderMap; + /// use http_headers::FieldName; + /// use http_headers::source::FieldSource; + /// + /// assert!(FieldSource::lines(&HeaderMap::new(), &FieldName::UserAgent).is_none()); + /// # } + /// ``` + fn lines(&self, name: &'static FieldName) -> Option>; +} + +impl FieldSource for &S +where + S: FieldSource + ?Sized, +{ + #[inline] + fn contains(&self, name: &'static FieldName) -> bool { + (**self).contains(name) + } + + #[inline] + fn lines(&self, name: &'static FieldName) -> Option> { + (**self).lines(name) + } +} diff --git a/crates/http_headers/src/source/list_item_count.rs b/crates/http_headers/src/source/list_item_count.rs new file mode 100644 index 000000000..e84948a09 --- /dev/null +++ b/crates/http_headers/src/source/list_item_count.rs @@ -0,0 +1,45 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use super::MAX_CUSTOM_LIST_ITEMS; +use crate::{DecodeError, DecodeErrorKind, FieldName}; + +pub(crate) fn update_list_item_count( + name: &'static FieldName, + bytes: &[u8], + delimiter: u8, + skip_empty: bool, + backslash_escapes: bool, + item_count: &mut usize, +) -> Result<(), DecodeError> { + let limit = || DecodeError::new(name, DecodeErrorKind::SourceLimitExceeded); + let mut start = 0_usize; + let mut quoted = false; + let mut escaped = false; + for (position, byte) in bytes.iter().copied().enumerate() { + if backslash_escapes && escaped { + escaped = false; + } else if backslash_escapes && quoted && byte == b'\\' { + escaped = true; + } else if byte == b'"' { + quoted = !quoted; + } else if !quoted && byte == delimiter { + let item = crate::validate::trim_ows(&bytes[start..position]); + if !skip_empty || !item.is_empty() { + *item_count = item_count.checked_add(1).ok_or_else(limit)?; + if *item_count > MAX_CUSTOM_LIST_ITEMS { + return Err(limit()); + } + } + start = position + 1; + } + } + let item = crate::validate::trim_ows(&bytes[start..]); + if !skip_empty || !item.is_empty() { + *item_count = item_count.checked_add(1).ok_or_else(limit)?; + if *item_count > MAX_CUSTOM_LIST_ITEMS { + return Err(limit()); + } + } + Ok(()) +} diff --git a/crates/http_headers/src/source/mod.rs b/crates/http_headers/src/source/mod.rs new file mode 100644 index 000000000..aefcd26d4 --- /dev/null +++ b/crates/http_headers/src/source/mod.rs @@ -0,0 +1,34 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Definitions involved in reading fields from field containers. +//! +//! A container implements [`FieldSource`] to expose the field lines stored +//! under a field name, and every typed decode in this crate starts there. +//! [`FieldLines`] is what a source hands back: the one-or-more raw field +//! lines for a single name, borrowed from whatever storage the container +//! uses. [`FieldLinesIter`] walks those lines, and [`DelimitedItems`] walks the +//! comma- or semicolon-separated items within them. +//! +//! The `http` cargo feature implements [`FieldSource`] for +//! `http::HeaderMap`. Implement it yourself only when integrating a different +//! container; reading a known header needs nothing from this module. Typed +//! decodes from custom sources accept at most [`MAX_CUSTOM_FIELD_BYTES`] total +//! bytes and [`MAX_CUSTOM_FIELD_LINES`] lines for one name. Delimited parsing +//! accepts at most [`MAX_CUSTOM_LIST_ITEMS`] items per decode. Exceeding a +//! budget returns [`DecodeErrorKind::SourceLimitExceeded`](crate::DecodeErrorKind::SourceLimitExceeded), +//! not an invalid-syntax error. The validated +//! `http::HeaderMap` adapter is exempt from these custom-source budgets. + +mod delimited_items; +mod field_lines; +mod field_source; +mod list_item_count; + +#[doc(inline)] +pub use delimited_items::DelimitedItems; +#[doc(inline)] +pub use field_lines::{FieldLines, FieldLinesIter, MAX_CUSTOM_FIELD_BYTES, MAX_CUSTOM_FIELD_LINES, MAX_CUSTOM_LIST_ITEMS}; +#[doc(inline)] +pub use field_source::FieldSource; +pub(crate) use list_item_count::update_list_item_count; diff --git a/crates/http_headers/src/test_sink.rs b/crates/http_headers/src/test_sink.rs new file mode 100644 index 000000000..ba12d61e9 --- /dev/null +++ b/crates/http_headers/src/test_sink.rs @@ -0,0 +1,67 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::collections::HashMap; + +use crate::sink::{EncodedValues, FieldSink, InsertError}; +use crate::source::{FieldLines, FieldSource}; +use crate::{FieldName, FieldValue}; + +#[derive(Debug, Default)] +pub(crate) struct TestSink { + values: HashMap>, +} + +impl TestSink { + pub(crate) fn new() -> Self { + Self::default() + } +} + +impl FieldSource for TestSink { + fn lines(&self, name: &'static FieldName) -> Option> { + self.values.get(name).and_then(|values| FieldLines::from_slice(name, values)) + } +} + +impl FieldSink for TestSink { + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + if values.is_empty() { + self.values.remove(name); + } else { + self.values.insert(name.clone(), values.into_iter().collect()); + } + Ok(()) + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.values.entry(name.clone()).or_default().extend(values); + Ok(()) + } + + fn remove_values(&mut self, name: &'static FieldName) { + self.values.remove(name); + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::TestSink; + use crate::sink::{EncodedValues, FieldSink}; + use crate::source::FieldSource; + use crate::{FieldName, FieldValue}; + + #[test] + fn setting_no_values_clears_the_entry() { + let mut sink = TestSink::new(); + + sink.set_values(&FieldName::Accept, EncodedValues::single(FieldValue::from_static("text/html"))) + .expect("the test sink always accepts a value"); + assert!(sink.lines(&FieldName::Accept).is_some()); + + sink.set_values(&FieldName::Accept, EncodedValues::new()) + .expect("the test sink always accepts an empty set"); + assert!(sink.lines(&FieldName::Accept).is_none()); + } +} diff --git a/crates/http_headers/src/test_support.rs b/crates/http_headers/src/test_support.rs new file mode 100644 index 000000000..fa52ef883 --- /dev/null +++ b/crates/http_headers/src/test_support.rs @@ -0,0 +1,49 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Native substitutions remain exhaustive; Miri samples grammar classes and ASCII case changes. + +#[cfg(miri)] +const BYTE_CLASSES: &[u8] = b"\0\t\n\x0b\x0c\r\x1f !\"*+,-./0123456789:;=?@AZ[\\]_`az{\x7f\x80\xff"; + +#[cfg(miri)] +pub(crate) fn is_byte_case(byte: u8, original: u8) -> bool { + BYTE_CLASSES.contains(&byte) || byte.eq_ignore_ascii_case(&original) +} + +#[cfg(miri)] +pub(crate) fn byte_cases(original: u8) -> impl Iterator { + let lowercase = original.to_ascii_lowercase(); + let uppercase = original.to_ascii_uppercase(); + BYTE_CLASSES + .iter() + .copied() + .chain((!BYTE_CLASSES.contains(&original)).then_some(original)) + .chain((lowercase != original && !BYTE_CLASSES.contains(&lowercase)).then_some(lowercase)) + .chain((uppercase != original && !BYTE_CLASSES.contains(&uppercase)).then_some(uppercase)) +} + +#[cfg(not(miri))] +pub(crate) fn byte_cases(_original: u8) -> impl Iterator { + u8::MIN..=u8::MAX +} + +#[cfg(miri)] +pub(crate) fn substitution_bytes(original: u8, position: usize, length: usize) -> impl Iterator { + assert!(position < length, "the substitution position must be inside the input"); + byte_cases(original).filter(move |&byte| { + if byte == original { + return position == 0; + } + // Keep structural mutations at every position, and distribute the remaining classes across the literal. + byte.eq_ignore_ascii_case(&original) + || matches!(byte, 0 | b' ' | b'"' | b'\\' | b',' | b';' | 0x7f | 0x80) + || original.is_ascii_digit() && byte.is_ascii_digit() + || usize::from(byte) % length == position + }) +} + +#[cfg(not(miri))] +pub(crate) fn substitution_bytes(original: u8, _position: usize, _length: usize) -> impl Iterator { + byte_cases(original) +} diff --git a/crates/http_headers/src/validate.rs b/crates/http_headers/src/validate.rs new file mode 100644 index 000000000..57018e1ac --- /dev/null +++ b/crates/http_headers/src/validate.rs @@ -0,0 +1,213 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Safe accelerated primitives for built-in and custom header parsers. +//! +//! This module is a private implementation detail: every built-in header +//! parser validates through these functions, but they are not part of the +//! crate's public API. + +/// Returns whether `bytes` is a non-empty RFC 9110 token. +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-content-type", + feature = "headers-cors", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", + feature = "headers-websocket", +))] +#[must_use] +#[inline] +pub(super) fn token(bytes: &[u8]) -> bool { + http_headers_simd::is_token(bytes) +} + +/// Returns whether `bytes` is an RFC 9110 `token68` value. +/// +/// A value contains one or more alphanumeric or `-._~+/` bytes followed only +/// by optional `=` padding. +#[cfg(any(test, feature = "headers-authorization"))] +#[must_use] +#[inline] +pub(super) fn token68(bytes: &[u8]) -> bool { + http_headers_simd::is_token68(bytes) +} + +/// Returns whether one byte is permitted in an RFC 9110 token. +#[must_use] +#[inline] +pub(super) const fn token_byte(byte: u8) -> bool { + matches!( + byte, + b'!' | b'#' + | b'$' + | b'%' + | b'&' + | b'\'' + | b'*' + | b'+' + | b'-' + | b'.' + | b'^' + | b'_' + | b'`' + | b'|' + | b'~' + | b'0'..=b'9' + | b'A'..=b'Z' + | b'a'..=b'z' + ) +} + +/// Returns whether every byte is permitted in an HTTP field value. +/// +/// Empty values are permitted. Carriage return, line feed, other controls, +/// and DEL are rejected. +#[must_use] +#[inline] +pub(super) fn field_value(bytes: &[u8]) -> bool { + http_headers_simd::is_field_value(bytes) +} + +/// Compares byte strings using ASCII case-insensitive equality. +/// +/// Non-ASCII bytes compare exactly. +#[cfg(any( + test, + feature = "headers-content-type", + feature = "headers-negotiation", + feature = "headers-range", + feature = "headers-security", +))] +#[must_use] +#[inline] +pub(super) fn eq_ignore_ascii_case(left: &[u8], right: &[u8]) -> bool { + http_headers_simd::eq_ignore_ascii_case(left, right) +} + +/// Removes optional HTTP whitespace from both ends of a byte string. +#[must_use] +#[inline] +pub(super) fn trim_ows(mut bytes: &[u8]) -> &[u8] { + while bytes.first().is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + bytes = &bytes[1..]; + } + while bytes.last().is_some_and(|byte| matches!(byte, b' ' | b'\t')) { + bytes = &bytes[..bytes.len() - 1]; + } + bytes +} + +/// Parses a nonempty unsigned decimal integer with checked arithmetic. +/// +/// Returns `None` for an empty value, a non-digit, or arithmetic overflow. +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-cors", + feature = "headers-range", + feature = "headers-security", +))] +#[must_use] +#[inline] +pub(super) fn decimal_u64(bytes: &[u8]) -> Option { + // Nineteen digits is the widest decimal `u64` always holds, so a shorter + // input needs no per-digit overflow check and the accumulation collapses + // to a multiply-add. + if bytes.is_empty() || bytes.len() > 19 { + return wide_decimal_u64(bytes); + } + let mut value = 0_u64; + for byte in bytes.iter().copied() { + let digit = byte.wrapping_sub(b'0'); + if digit > 9 { + return None; + } + value = value * 10 + u64::from(digit); + } + Some(value) +} + +/// Parses the decimals long enough to need an overflow check on every digit. +#[cfg(any( + test, + feature = "headers-cache-control", + feature = "headers-cors", + feature = "headers-range", + feature = "headers-security", +))] +#[cold] +#[inline(never)] +fn wide_decimal_u64(bytes: &[u8]) -> Option { + if bytes.is_empty() { + return None; + } + bytes.iter().try_fold(0_u64, |value, byte| { + value + .checked_mul(10)? + .checked_add(u64::from(byte.checked_sub(b'0').filter(|digit| *digit <= 9)?)) + }) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::*; + + #[test] + fn token_accepts_and_rejects() { + assert!(token(b"gzip")); + assert!(!token(b"not a token")); + } + + #[test] + fn token68_accepts_and_rejects() { + assert!(token68(b"YWxpY2U6c2VjcmV0==")); + assert!(!token68(b"data=more")); + } + + #[test] + fn token_byte_accepts_and_rejects() { + assert!(token_byte(b'a')); + assert!(!token_byte(b' ')); + } + + #[test] + fn token_byte_agrees_with_token_for_every_byte() { + for byte in u8::MIN..=u8::MAX { + assert_eq!(token_byte(byte), token(&[byte])); + } + } + + #[test] + fn field_value_accepts_and_rejects() { + assert!(field_value(b"text/plain")); + assert!(!field_value(b"line\r\nbreak")); + } + + #[test] + fn eq_ignore_ascii_case_compares() { + assert!(eq_ignore_ascii_case(b"gzip", b"GZIP")); + assert!(!eq_ignore_ascii_case(b"gzip", b"br")); + } + + #[test] + fn trim_ows_strips_whitespace() { + assert_eq!(trim_ows(b" \tvalue\t "), b"value"); + assert_eq!(trim_ows(b"\rvalue\n"), b"\rvalue\n"); + assert_eq!(trim_ows(b"\t \t"), b""); + } + + #[test] + fn decimal_u64_parses_and_rejects() { + assert_eq!(decimal_u64(b"3600"), Some(3600)); + assert_eq!(decimal_u64(b"-1"), None); + assert_eq!(decimal_u64(b"0"), Some(0)); + assert_eq!(decimal_u64(b"00123"), Some(123)); + assert_eq!(decimal_u64(u64::MAX.to_string().as_bytes()), Some(u64::MAX)); + assert_eq!(decimal_u64(b"12x"), None); + assert_eq!(decimal_u64(b"18446744073709551616"), None); + } +} diff --git a/crates/http_headers/tests/__fuzz__/campaign.toml b/crates/http_headers/tests/__fuzz__/campaign.toml new file mode 100644 index 000000000..45dd8ab18 --- /dev/null +++ b/crates/http_headers/tests/__fuzz__/campaign.toml @@ -0,0 +1,39 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +[campaign] +package = "http_headers" +default_engine = "libfuzzer" +default_time = "300s" +default_max_input_length = 4096 +corpus_root = "crates/http_headers/tests/__fuzz__" + +[[target]] +name = "arbitrary_header_values_do_not_panic" +test = "fuzz" +description = "Exercise every representative parser over arbitrary HeaderValue-compatible bytes." + +[[target]] +name = "expanded_header_values_do_not_panic" +test = "fuzz" +description = "Exercise negotiation, conditional, range, CORS, WebSocket, and security parsers." + +[[target]] +name = "comma_and_quoted_delimiters_match_reference" +test = "fuzz" +description = "Compare comma and semicolon parsing, including quoted delimiters, with an independent reference." + +[[target]] +name = "content_type_cached_and_uncached_agree" +test = "fuzz" +description = "Compare borrowed Content-Type decoding with first and repeated owned metadata outcomes." + +[[target]] +name = "scalar_and_simd_agree" +test = "fuzz" +description = "Compare scalar and runtime-selected byte scanners over arbitrary bytes." + +[[target]] +name = "typed_round_trips" +test = "fuzz" +description = "Construct, encode, decode, and compare representative typed header values." diff --git a/crates/http_headers/tests/access_control_allow_credentials.rs b/crates/http_headers/tests/access_control_allow_credentials.rs new file mode 100644 index 000000000..9cbf4c8d1 --- /dev/null +++ b/crates/http_headers/tests/access_control_allow_credentials.rs @@ -0,0 +1,64 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Credentials decoding preserves singleton and whitespace semantics. + +#![cfg(all(feature = "http", feature = "headers-cors"))] + +use http::{HeaderMap, HeaderValue}; +use http_headers::headers::{AccessControlAllowCredentials, AccessControlAllowCredentialsOwned}; +use http_headers::{DecodeErrorKind, DecodeMode, Field}; + +#[test] +fn owned_and_borrowed_accept_only_true_with_optional_whitespace() { + for wire in ["true", " true", "true\t", " \ttrue\t "] { + let mut map = HeaderMap::new(); + map.insert(http::header::ACCESS_CONTROL_ALLOW_CREDENTIALS, HeaderValue::from_static(wire)); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let view = AccessControlAllowCredentials::view_with(&map, mode).unwrap().unwrap(); + assert_eq!(view.as_str(), "true"); + assert_eq!(view.as_field_value(), "true"); + assert_eq!( + AccessControlAllowCredentials::owned_with(&map, mode).unwrap(), + Some(AccessControlAllowCredentialsOwned::allow()) + ); + } + } + + for wire in ["", " \t ", "True", "tRue", "trUe", "truE", "tru", "truee", " truee ", "true,false"] { + let mut map = HeaderMap::new(); + map.insert(http::header::ACCESS_CONTROL_ALLOW_CREDENTIALS, HeaderValue::from_static(wire)); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!( + AccessControlAllowCredentials::view_with(&map, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + AccessControlAllowCredentials::owned_with(&map, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + } + } +} + +#[test] +fn owned_and_borrowed_distinguish_absence_and_duplicate_values() { + let mut map = HeaderMap::new(); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!(AccessControlAllowCredentials::view_with(&map, mode).unwrap(), None); + assert_eq!(AccessControlAllowCredentials::owned_with(&map, mode).unwrap(), None); + } + + map.append(http::header::ACCESS_CONTROL_ALLOW_CREDENTIALS, HeaderValue::from_static("true")); + map.append(http::header::ACCESS_CONTROL_ALLOW_CREDENTIALS, HeaderValue::from_static("true")); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!( + AccessControlAllowCredentials::view_with(&map, mode).unwrap_err().kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!( + AccessControlAllowCredentials::owned_with(&map, mode).unwrap_err().kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + } +} diff --git a/crates/http_headers/tests/axum_example.rs b/crates/http_headers/tests/axum_example.rs new file mode 100644 index 000000000..59e2dfe21 --- /dev/null +++ b/crates/http_headers/tests/axum_example.rs @@ -0,0 +1,84 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Semantic coverage for the Axum example. + +#![cfg(all(feature = "http", feature = "headers-cache-control", feature = "headers-user-agent"))] +#![expect(clippy::unwrap_used, reason = "test failures provide sufficient context")] + +use axum::body::{Body, to_bytes}; +use axum::http::header::{CACHE_CONTROL, USER_AGENT}; +use axum::http::{HeaderValue, Request, StatusCode}; +use tokio::runtime::Builder; +use tower::ServiceExt; + +#[path = "../examples/axum/app.rs"] +mod app; + +const CACHE_POLICY: &str = "private, max-age=60"; + +fn response(request: Request) -> (StatusCode, Option, String) { + // In-memory requests need no OS I/O driver, which Windows Miri cannot emulate. + Builder::new_current_thread().build().unwrap().block_on(async { + let response = app::router().oneshot(request).await.unwrap(); + let status = response.status(); + let cache_control = response.headers().get(CACHE_CONTROL).cloned(); + let body = to_bytes(response.into_body(), usize::MAX).await.unwrap(); + (status, cache_control, String::from_utf8(body.to_vec()).unwrap()) + }) +} + +#[test] +fn absent_user_agent_returns_unknown_greeting_and_private_cache_policy() { + let request = Request::builder().uri("/").body(Body::empty()).unwrap(); + + let (status, cache_control, body) = response(request); + + assert_eq!(status, StatusCode::OK); + assert_eq!(cache_control, Some(HeaderValue::from_static(CACHE_POLICY))); + assert_eq!(body, r#"{"message":"hello from http_headers","user_agent":"unknown"}"#); +} + +#[test] +fn valid_user_agent_is_returned_with_private_cache_policy() { + let request = Request::builder() + .uri("/") + .header(USER_AGENT, "example-client/1.0") + .body(Body::empty()) + .unwrap(); + + let (status, cache_control, body) = response(request); + + assert_eq!(status, StatusCode::OK); + assert_eq!(cache_control, Some(HeaderValue::from_static(CACHE_POLICY))); + assert_eq!(body, r#"{"message":"hello from http_headers","user_agent":"example-client/1.0"}"#); +} + +#[test] +fn malformed_user_agents_return_bad_request() { + let invalid_utf8 = Request::builder() + .uri("/") + .header(USER_AGENT, HeaderValue::from_bytes(b"\xff").unwrap()) + .body(Body::empty()) + .unwrap(); + + let (status, cache_control, body) = response(invalid_utf8); + + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(cache_control, None); + assert_eq!(body, "invalid user-agent header: invalid UTF-8"); + + let mut multiple_values = Request::builder().uri("/").body(Body::empty()).unwrap(); + multiple_values + .headers_mut() + .append(USER_AGENT, HeaderValue::from_static("client/1")); + multiple_values + .headers_mut() + .append(USER_AGENT, HeaderValue::from_static("client/2")); + + let (status, cache_control, body) = response(multiple_values); + + assert_eq!(status, StatusCode::BAD_REQUEST); + assert_eq!(cache_control, None); + assert_eq!(body, "invalid user-agent header: unexpected multiple values"); +} diff --git a/crates/http_headers/tests/benchmark_ownership.rs b/crates/http_headers/tests/benchmark_ownership.rs new file mode 100644 index 000000000..410782fc3 --- /dev/null +++ b/crates/http_headers/tests/benchmark_ownership.rs @@ -0,0 +1,208 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Ownership regressions for the storage and name-recognition benchmark fixtures. + +#![cfg(all(feature = "http", feature = "headers-content-length", feature = "headers-content-type"))] + +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Arc, LazyLock}; +use std::{iter, ptr}; + +use bytes::Bytes; +use http::header::{CONTENT_LENGTH, CONTENT_TYPE, InvalidHeaderValue, SET_COOKIE}; +use http::{HeaderMap, HeaderName, HeaderValue}; +use http_headers::{FieldName, FieldValue}; + +use self::name_corpus::{crate_names, custom_header_names, custom_names_lowercase, custom_names_mixed_case, http_names, known_names}; +use self::storage_operations::{ + MapOutput, content_length_deferred_http, content_length_materialized, http_append_values, http_writer_borrowed, + http_writer_materialized, http_writer_streamed, http_writer_streamed_sized, +}; + +#[path = "common/http_headers_name_corpus.rs"] +mod name_corpus; +#[path = "common/http_headers_storage_operations.rs"] +mod storage_operations; + +const SENTINEL_NAME: &str = "x-benchmark-owner"; +const SENTINEL_VALUE: &[u8] = b"this allocation must survive until output cleanup"; +const SHORT_VALUE: &str = "application/json; charset=utf-8"; +const LONG_VALUE: &str = "multipart/form-data; boundary=----WebKitFormBoundary7MA4YWxkTrZu0gW; charset=utf-8"; +static LARGE_STREAM: [u8; 65_536] = [b'y'; 65_536]; +static CUSTOM_APPEND_NAME: LazyLock = LazyLock::new(|| FieldName::from_static("x-cookie-batch")); + +#[derive(Debug)] +struct TrackedBytes { + bytes: Vec, + dropped: Arc, +} + +impl AsRef<[u8]> for TrackedBytes { + fn as_ref(&self) -> &[u8] { + &self.bytes + } +} + +impl Drop for TrackedBytes { + fn drop(&mut self) { + self.dropped.fetch_add(1, Ordering::Relaxed); + } +} + +fn assert_batches_reclaim_owners( + mut operation: impl FnMut(HeaderMap) -> MapOutput, + inspect: impl Fn(&HeaderMap), +) -> Result<(), InvalidHeaderValue> { + const BATCH_SIZE: usize = 8; + const SAMPLES: usize = 4; + + let dropped = Arc::new(AtomicUsize::new(0)); + for sample in 0..SAMPLES { + let mut outputs = Vec::with_capacity(BATCH_SIZE); + for _ in 0..BATCH_SIZE { + let sentinel = Bytes::from_owner(TrackedBytes { + bytes: SENTINEL_VALUE.to_vec(), + dropped: Arc::clone(&dropped), + }); + let mut map = HeaderMap::with_capacity(128); + map.insert(SENTINEL_NAME, HeaderValue::from_maybe_shared(sentinel)?); + + let (length, map) = operation(map); + assert_eq!(length, map.len()); + assert_eq!(map[SENTINEL_NAME].as_bytes(), SENTINEL_VALUE); + inspect(&map); + outputs.push((length, map)); + assert_eq!(dropped.load(Ordering::Relaxed), sample * BATCH_SIZE); + } + drop(outputs); + assert_eq!(dropped.load(Ordering::Relaxed), (sample + 1) * BATCH_SIZE); + } + Ok(()) +} + +#[test] +fn writer_outputs_retain_and_release_every_map() -> Result<(), InvalidHeaderValue> { + for value in [SHORT_VALUE, LONG_VALUE] { + assert_batches_reclaim_owners( + |map| http_writer_borrowed((map, value)), + |map| { + assert_eq!(map.len(), 2); + assert_eq!(map[CONTENT_TYPE].as_bytes(), value.as_bytes()); + }, + )?; + assert_batches_reclaim_owners( + |map| http_writer_streamed((map, value)), + |map| { + assert_eq!(map.len(), 2); + assert_eq!(map[CONTENT_TYPE].as_bytes(), value.as_bytes()); + }, + )?; + } + + let streamed: [&'static [u8]; 6] = [&[], &[b's'; 16], &[b'm'; 32], &[b'l'; 65], &[b'x'; 4096], &LARGE_STREAM]; + for value in streamed { + assert_batches_reclaim_owners( + |map| http_writer_streamed_sized((map, value)), + |map| { + assert_eq!(map.len(), 2); + assert_eq!(map[CONTENT_TYPE].as_bytes(), value); + }, + )?; + } + + assert_batches_reclaim_owners( + |map| http_writer_materialized((map, FieldValue::from_static(LONG_VALUE))), + |map| { + assert_eq!(map.len(), 2); + assert_eq!(map[CONTENT_TYPE].as_bytes(), LONG_VALUE.as_bytes()); + }, + )?; + Ok(()) +} + +#[test] +fn append_outputs_retain_existing_values_and_release_every_map() -> Result<(), InvalidHeaderValue> { + let names: [&'static FieldName; 2] = [&FieldName::SetCookie, &CUSTOM_APPEND_NAME]; + for name in names { + let http_name = if name == &FieldName::SetCookie { + SET_COOKIE + } else { + HeaderName::from_static("x-cookie-batch") + }; + for occupied in [false, true] { + for count in [1, 4, 16, 64] { + assert_batches_reclaim_owners( + |mut map| { + if occupied { + map.insert(http_name.clone(), HeaderValue::from_static("existing")); + } + let values = iter::repeat_n(FieldValue::from_static("value"), count).collect(); + http_append_values((map, name, values)) + }, + |map| { + assert_eq!(map.len(), 1 + count + usize::from(occupied)); + for (index, value) in map.get_all(&http_name).iter().enumerate() { + let expected = if occupied && index == 0 { b"existing".as_slice() } else { b"value" }; + assert_eq!(value.as_bytes(), expected); + } + }, + )?; + } + } + } + Ok(()) +} + +#[test] +fn deferred_and_materialized_outputs_release_every_map() -> Result<(), InvalidHeaderValue> { + for operation in [content_length_materialized, content_length_deferred_http] { + assert_batches_reclaim_owners(operation, |map| { + assert_eq!(map.len(), 2); + assert_eq!(map[CONTENT_LENGTH].as_bytes(), b"1024"); + })?; + } + Ok(()) +} + +#[test] +fn name_setups_reuse_immutable_corpora_and_owned_names() { + let known = known_names(); + let lowercase = custom_names_lowercase(); + let mixed_case = custom_names_mixed_case(); + let custom = custom_header_names(); + let http = http_names(); + let crate_names = crate_names(); + + assert_eq!(known.len(), 15); + assert_eq!(lowercase.len(), 8); + assert_eq!(mixed_case.len(), 8); + assert_eq!(custom.len(), 8); + assert_eq!(http.len(), 23); + assert_eq!(crate_names.len(), 23); + for (name, expected) in custom.iter().zip(lowercase) { + assert_eq!(name.as_str().as_bytes(), *expected); + } + for ((http_name, crate_name), expected) in http.iter().zip(crate_names).zip(known.iter().chain(mixed_case)) { + assert!(http_name.as_str().as_bytes().eq_ignore_ascii_case(expected)); + assert_eq!(http_name.as_str(), crate_name.as_str()); + } + + for _ in 0..128 { + let next_custom = custom_header_names(); + let next_http = http_names(); + let next_crate = name_corpus::crate_names(); + assert!(ptr::eq(custom, next_custom)); + assert!(ptr::eq(http, next_http)); + assert!(ptr::eq(crate_names, next_crate)); + for (original, next) in custom.iter().zip(next_custom) { + assert!(ptr::eq(original.as_str(), next.as_str())); + } + for (original, next) in http.iter().zip(next_http) { + assert!(ptr::eq(original.as_str(), next.as_str())); + } + for (original, next) in crate_names.iter().zip(next_crate) { + assert!(ptr::eq(original.as_str(), next.as_str())); + } + } +} diff --git a/crates/http_headers/tests/bolero_fuzz.rs b/crates/http_headers/tests/bolero_fuzz.rs new file mode 100644 index 000000000..a8c28fbf9 --- /dev/null +++ b/crates/http_headers/tests/bolero_fuzz.rs @@ -0,0 +1,1004 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Bounded Bolero properties for header parsing and encoding. + +#![cfg(feature = "headers-all")] + +use std::any::type_name; +use std::sync::LazyLock; +use std::time::Duration; +use std::{iter, slice, str}; + +use http_headers::headers::{ + Accept, AcceptEncoding, AcceptLanguage, AcceptRanges, AccessControlAllowCredentials, AccessControlAllowHeaders, + AccessControlAllowMethods, AccessControlAllowOrigin, AccessControlExposeHeaders, AccessControlMaxAge, AccessControlRequestHeaders, + AccessControlRequestMethod, Allow, Authorization, AuthorizationOwned, Basic, BasicCredentials, Bearer, CacheControl, CacheControlOwned, + ContentLength, ContentLengthOwned, ContentRange, ContentSecurityPolicy, ContentType, ContentTypeOwned, ETag, ETagOwned, Host, IfMatch, + IfModifiedSince, IfNoneMatch, IfRange, IfUnmodifiedSince, LastModified, Location, LocationOwned, MediaTypeParameterView, Range, + ReferrerPolicy, SecWebSocketAccept, SecWebSocketExtensions, SecWebSocketKey, SecWebSocketProtocol, SecWebSocketVersion, Server, + SetCookie, SetCookieOwned, StrictTransportSecurity, UserAgent, UserAgentOwned, Vary, XContentTypeOptions, +}; +use http_headers::sink::{EncodedValues, FieldSink, InsertError}; +use http_headers::source::{FieldSource, MAX_CUSTOM_LIST_ITEMS}; +use http_headers::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef}; + +use self::common::TestMap; + +mod common; + +const BOUNDED_ITERATIONS: usize = 2_048; +const BOUNDED_TEST_TIME: Duration = Duration::from_millis(400); +const MAX_VALUE_LENGTH: usize = 512; +const MAX_FIELD_LINES_U8: u8 = 3; + +static COMMA_LIST: LazyLock = LazyLock::new(|| FieldName::from_static("x-comma-list")); +static SEMICOLON_LIST: LazyLock = LazyLock::new(|| FieldName::from_static("x-semicolon-list")); + +impl TestMap { + fn new() -> Self { + Self::default() + } + + fn append(&mut self, name: FieldName, value: FieldValue) { + self.0.entry(name).or_default().push(value); + } +} + +macro_rules! bounded { + () => { + bolero::check!() + .with_iterations(BOUNDED_ITERATIONS) + .with_test_time(BOUNDED_TEST_TIME) + }; +} + +#[derive(Debug, Eq, PartialEq)] +enum ParseOutcome { + Parsed(Vec>), + Rejected(DecodeErrorKind, Option), +} + +#[derive(Debug, Eq, PartialEq)] +enum DecoderOutcome { + Absent, + Present, + Rejected, +} + +#[derive(Debug, Eq, PartialEq)] +enum ContentTypeOutcome { + Parsed { + type_: String, + subtype: String, + parameters: Vec<(String, Vec)>, + }, + Rejected(DecodeErrorKind, Option), + Absent, +} + +#[derive(Debug, Eq, PartialEq)] +struct CommaItems(Vec>); + +#[derive(Debug, Eq, PartialEq)] +struct SemicolonItems(Vec>); + +type DelimitedOutcome = Result>, (DecodeErrorKind, Option)>; + +// Deliberately disagrees so the generic oracle cannot silently skip a Field. +struct AsymmetricField; + +impl Field for AsymmetricField { + type View<'a> = (); + type Owned = FieldValue; + + fn name() -> &'static FieldName { + &COMMA_LIST + } + + fn view_with(_source: &S, _mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + if BORROWED_ACCEPTS { + Ok(Some(())) + } else { + Err(DecodeError::new(Self::name(), DecodeErrorKind::InvalidSyntax)) + } + } + + fn owned_with(_source: &S, _mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + if BORROWED_ACCEPTS { + Err(DecodeError::new(Self::name(), DecodeErrorKind::InvalidSyntax)) + } else { + Ok(Some(FieldValue::from_static("member"))) + } + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::single(value)) + } +} + +impl Field for CommaItems { + type View<'a> = Self; + type Owned = Self; + + fn name() -> &'static FieldName { + &COMMA_LIST + } + + fn view_with(source: &S, _mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines + .comma_items() + .map(|item| item.map(<[u8]>::to_vec)) + .collect::, _>>() + .map(Self) + .map(Some) + } + + fn owned_with(source: &S, mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + Self::view_with(source, mode) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), encode_items(&value.0)) + } +} + +impl Field for SemicolonItems { + type View<'a> = Self; + type Owned = Self; + + fn name() -> &'static FieldName { + &SEMICOLON_LIST + } + + fn view_with(source: &S, _mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + let Some(lines) = source.lines(Self::name()) else { + return Ok(None); + }; + lines + .semicolon_items() + .map(|item| item.map(<[u8]>::to_vec)) + .collect::, _>>() + .map(Self) + .map(Some) + } + + fn owned_with(source: &S, mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + Self::view_with(source, mode) + } + + fn insert(sink: &mut S, value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), encode_items(&value.0)) + } +} + +fn encode_items(items: &[Vec]) -> EncodedValues { + EncodedValues::from_vec( + items + .iter() + .map(|item| FieldValue::from_bytes(item).expect("delimited items came from valid field-value bytes")) + .collect(), + ) +} + +fn valid_field_byte(byte: u8) -> u8 { + if byte == b'\t' || byte >= b' ' && byte != 0x7f { + byte + } else { + const STRUCTURAL: &[u8] = b",;\"\\ =/%#"; + STRUCTURAL[usize::from(byte) % STRUCTURAL.len()] + } +} + +fn shaped_field_value(input: &[u8]) -> FieldValue { + let bytes: Vec<_> = input.iter().copied().take(MAX_VALUE_LENGTH).map(valid_field_byte).collect(); + FieldValue::from_bytes(&bytes).expect("the generator emits only valid field-value bytes") +} + +fn shaped_field_values(input: &[u8]) -> Vec { + let (selector, input) = input.split_first().map_or((0, &[][..]), |(first, rest)| (*first, rest)); + let count = usize::from(selector % MAX_FIELD_LINES_U8) + 1; + (0..count) + .map(|index| { + let start = input.len() * index / count; + let end = input.len() * (index + 1) / count; + shaped_field_value(&input[start..end]) + }) + .collect() +} + +fn map_with_values(values: &[FieldValue]) -> TestMap { + let mut map = TestMap::new(); + for value in values { + map.append(H::name().clone(), value.clone()); + } + map +} + +fn inserted_bytes(header: H::Owned) -> Vec> { + let mut map = TestMap::new(); + H::insert(&mut map, header).expect("an empty header map has capacity"); + map.get_all(H::name()).iter().map(|value| value.as_bytes().to_vec()).collect() +} + +fn borrowed_outcome(map: &TestMap) -> ParseOutcome { + match H::view(map) { + Ok(Some(_view)) => ParseOutcome::Parsed(Vec::new()), + Ok(None) => ParseOutcome::Parsed(Vec::new()), + Err(error) => ParseOutcome::Rejected(error.kind(), error.value_index()), + } +} + +fn owned_outcome(map: &TestMap) -> ParseOutcome { + match H::owned(map) { + Ok(Some(header)) => ParseOutcome::Parsed(inserted_bytes::(header)), + Ok(None) => ParseOutcome::Parsed(Vec::new()), + Err(error) => ParseOutcome::Rejected(error.kind(), error.value_index()), + } +} + +fn content_type_projection<'a>( + type_: Result<&'a str, DecodeError>, + subtype: Result<&'a str, DecodeError>, + parameters: impl Iterator, DecodeError>>, +) -> ContentTypeOutcome { + let type_ = match type_ { + Ok(type_) => type_, + Err(error) => return ContentTypeOutcome::Rejected(error.kind(), error.value_index()), + }; + let subtype = match subtype { + Ok(subtype) => subtype, + Err(error) => return ContentTypeOutcome::Rejected(error.kind(), error.value_index()), + }; + let mut projected_parameters = Vec::new(); + for parameter in parameters { + match parameter { + Ok(parameter) => projected_parameters.push((parameter.name().to_owned(), parameter.value().to_vec())), + Err(error) => return ContentTypeOutcome::Rejected(error.kind(), error.value_index()), + } + } + ContentTypeOutcome::Parsed { + type_: type_.to_owned(), + subtype: subtype.to_owned(), + parameters: projected_parameters, + } +} + +fn content_type_view_outcome(map: &TestMap) -> ContentTypeOutcome { + match ContentType::view(map) { + Ok(Some(value)) => content_type_projection(value.type_(), value.subtype(), value.parameters()), + Ok(None) => ContentTypeOutcome::Absent, + Err(error) => ContentTypeOutcome::Rejected(error.kind(), error.value_index()), + } +} + +fn content_type_owned_outcome(map: &TestMap) -> ContentTypeOutcome { + match ContentType::owned(map) { + Ok(Some(value)) => content_type_projection(value.type_(), value.subtype(), value.parameters()), + Ok(None) => ContentTypeOutcome::Absent, + Err(error) => ContentTypeOutcome::Rejected(error.kind(), error.value_index()), + } +} + +fn normalized_outcome(result: &Result, DecodeError>) -> DecoderOutcome { + // Field promises the same admission policy, not identical diagnostics: + // empty WebSocket protocol lists report MissingValue from the view and + // InvalidSyntax from the owned decoder. + match result { + Ok(None) => DecoderOutcome::Absent, + Ok(Some(_)) => DecoderOutcome::Present, + Err(_) => DecoderOutcome::Rejected, + } +} + +fn assert_parser_consistency(values: &[FieldValue]) { + let map = map_with_values::(values); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let borrowed = H::view_with(&map, mode); + let owned = H::owned_with(&map, mode); + assert_eq!( + normalized_outcome(&borrowed), + normalized_outcome(&owned), + "{} {mode:?} borrowed and owned decoders disagree: borrowed error {:?}, owned error {:?}, values {values:?}", + type_name::(), + borrowed.as_ref().err(), + owned.as_ref().err() + ); + + if let Ok(header) = owned { + let encoded = header.map_or_else(Vec::new, inserted_bytes::); + let round_trip_values: Vec<_> = encoded + .iter() + .map(|bytes| FieldValue::from_bytes(bytes).expect("a header must encode valid FieldValue bytes")) + .collect(); + let round_trip = map_with_values::(&round_trip_values); + let decoded = H::owned_with(&round_trip, mode).expect("encoding a decoded header must preserve acceptance in the same mode"); + // Empty lists may encode no lines and become absent on decoding. + assert_eq!( + decoded.map_or_else(Vec::new, inserted_bytes::), + encoded, + "{} {mode:?} owned encode/decode round trip changed field bytes", + type_name::() + ); + } + } +} + +fn assert_duplicate_singleton_rejected(value: FieldValue) { + let values = [value.clone(), value]; + assert_eq!( + owned_outcome::(&map_with_values::(&values)), + ParseOutcome::Rejected(DecodeErrorKind::UnexpectedMultipleValues, None) + ); +} + +fn assert_absence_differs_from_empty() { + assert_eq!(owned_outcome::(&TestMap::new()), ParseOutcome::Parsed(Vec::new())); + let empty = [FieldValue::from_static("")]; + assert!( + matches!(owned_outcome::(&map_with_values::(&empty)), ParseOutcome::Rejected(_, _)), + "a present empty {} value must not be treated as absence", + H::name() + ); +} + +fn assert_constructed_round_trip(header: H::Owned) +where + H: Field, +{ + let mut map = TestMap::new(); + H::insert(&mut map, header).expect("an empty header map has capacity"); + let expected: Vec<_> = map.get_all(H::name()).iter().map(|value| value.as_bytes().to_vec()).collect(); + assert_eq!(borrowed_outcome::(&map), ParseOutcome::Parsed(Vec::new())); + assert_eq!(owned_outcome::(&map), ParseOutcome::Parsed(expected.clone())); + + let values: Vec<_> = expected + .iter() + .map(|bytes| FieldValue::from_bytes(bytes).expect("typed encoding must produce FieldValue bytes")) + .collect(); + assert_eq!(owned_outcome::(&map_with_values::(&values)), ParseOutcome::Parsed(expected)); +} + +fn trim_ows(bytes: &[u8]) -> &[u8] { + let start = bytes.iter().position(|byte| !matches!(byte, b' ' | b'\t')).unwrap_or(bytes.len()); + let end = bytes + .iter() + .rposition(|byte| !matches!(byte, b' ' | b'\t')) + .map_or(start, |index| index + 1); + &bytes[start..end] +} + +fn push_reference_item(output: &mut Vec>, item: &[u8]) -> Result<(), (DecodeErrorKind, Option)> { + if output.len() == MAX_CUSTOM_LIST_ITEMS { + return Err((DecodeErrorKind::SourceLimitExceeded, None)); + } + output.push(item.to_vec()); + Ok(()) +} + +fn reference_delimited(values: &[FieldValue], delimiter: u8) -> DelimitedOutcome { + let mut output = Vec::new(); + for (value_index, value) in values.iter().enumerate() { + let bytes = value.as_bytes(); + let mut start = 0; + let mut quoted = false; + let mut escaped = false; + for (position, byte) in bytes.iter().copied().enumerate() { + if escaped { + escaped = false; + } else if quoted && byte == b'\\' { + escaped = true; + } else if byte == b'"' { + quoted = !quoted; + } else if !quoted && byte == delimiter { + let item = trim_ows(&bytes[start..position]); + if delimiter != b',' || !item.is_empty() { + push_reference_item(&mut output, item)?; + } + start = position + 1; + } + } + if quoted || escaped { + return Err((DecodeErrorKind::UnterminatedQuote, Some(value_index))); + } + let item = trim_ows(&bytes[start..]); + if delimiter != b',' || !item.is_empty() { + push_reference_item(&mut output, item)?; + } + } + Ok(output) +} + +fn delimited_outcome(values: &[FieldValue]) -> DelimitedOutcome +where + for<'a> H::View<'a>: IntoDelimitedItems, +{ + let map = map_with_values::(values); + match H::view(&map) { + Ok(Some(items)) => Ok(items.into_items()), + Ok(None) => Ok(Vec::new()), + Err(error) => Err((error.kind(), error.value_index())), + } +} + +trait IntoDelimitedItems { + fn into_items(self) -> Vec>; +} + +impl IntoDelimitedItems for CommaItems { + fn into_items(self) -> Vec> { + self.0 + } +} + +impl IntoDelimitedItems for SemicolonItems { + fn into_items(self) -> Vec> { + self.0 + } +} + +fn token_from(input: &[u8], salt: u8) -> String { + const TOKEN: &[u8] = b"abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789!#$%&'*+-.^_`|~"; + let mut output = String::new(); + for byte in input.iter().copied().take(24) { + output.push(char::from(TOKEN[usize::from(byte.wrapping_add(salt)) % TOKEN.len()])); + } + if output.is_empty() { + output.push(char::from(TOKEN[usize::from(salt) % TOKEN.len()])); + } + output +} + +fn token68_from(input: &[u8]) -> String { + const TOKEN68: &[u8] = b"abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789-._~+/"; + let mut output = String::new(); + for byte in input.iter().copied().take(24) { + output.push(char::from(TOKEN68[usize::from(byte) % TOKEN68.len()])); + } + if output.is_empty() { + output.push('A'); + } + output.extend(iter::repeat_n('=', usize::from(input.first().copied().unwrap_or(0) % 3))); + output +} + +fn opaque_from(input: &[u8]) -> String { + input + .iter() + .copied() + .take(32) + .map(|byte| { + let printable = b'!' + byte % (b'~' - b'!' + 1); + char::from(if printable == b'"' { b'#' } else { printable }) + }) + .collect() +} + +fn quoted_text_from(input: &[u8]) -> String { + const QUOTED: &[u8] = b"abcXYZ019 ,;="; + input + .iter() + .copied() + .take(24) + .map(|byte| char::from(QUOTED[usize::from(byte) % QUOTED.len()])) + .collect() +} + +fn bounded_bytes(input: &[u8], start: usize) -> Vec { + input.iter().copied().skip(start).take(64).collect() +} + +#[test] +#[should_panic(expected = "borrowed and owned decoders disagree")] +fn parser_consistency_rejects_borrowed_only_acceptance() { + assert_parser_consistency::>(&[FieldValue::from_static("member")]); +} + +#[test] +#[should_panic(expected = "borrowed and owned decoders disagree")] +fn parser_consistency_rejects_owned_only_acceptance() { + assert_parser_consistency::>(&[FieldValue::from_static("member")]); +} + +#[test] +fn parser_outcomes_distinguish_absence_presence_and_rejection() { + assert_eq!(normalized_outcome::<()>(&Ok(None)), DecoderOutcome::Absent); + assert_eq!(normalized_outcome(&Ok(Some(()))), DecoderOutcome::Present); + assert_eq!( + normalized_outcome::<()>(&Err(DecodeError::new(&COMMA_LIST, DecodeErrorKind::InvalidSyntax))), + DecoderOutcome::Rejected + ); +} + +#[test] +fn websocket_parity_preserves_distinct_diagnostics_and_empty_line_semantics() { + let empty = [FieldValue::from_static(", ,")]; + let map = map_with_values::(&empty); + assert_eq!( + borrowed_outcome::(&map), + ParseOutcome::Rejected(DecodeErrorKind::MissingValue, None) + ); + assert_eq!( + owned_outcome::(&map), + ParseOutcome::Rejected(DecodeErrorKind::InvalidSyntax, None) + ); + + for raw in [ + &[][..], + &[""], + &[" \t "], + &[", ,"], + &["\"unterminated"], + &["chat"], + &["", "chat"], + &["chat", ""], + &[" ,\t, ", "chat", ""], + ] { + let values: Vec<_> = raw.iter().map(|value| FieldValue::from_static(value)).collect(); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + } +} + +#[test] +#[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] +fn arbitrary_header_values_do_not_panic() { + bounded!().for_each(|input: &[u8]| { + let values = shaped_field_values(input); + assert_parser_consistency::>(&values); + assert_parser_consistency::>(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + + let basic_map = map_with_values::>(&values); + let mut credentials = BasicCredentials::new(); + if let Ok(Some(authorization)) = Authorization::::view(&basic_map) { + let _ = authorization.extract(&mut credentials); + } + }); +} + +#[test] +#[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] +fn expanded_header_values_do_not_panic() { + bounded!().for_each(|input: &[u8]| { + let values = shaped_field_values(input); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + assert_parser_consistency::(&values); + }); +} + +#[test] +#[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] +fn comma_and_quoted_delimiters_match_reference() { + bounded!().for_each(|input: &[u8]| { + let values = shaped_field_values(input); + let comma = reference_delimited(&values, b','); + assert_eq!(delimited_outcome::(&values), comma); + + let semicolon = reference_delimited(&values, b';'); + assert_eq!(delimited_outcome::(&values), semicolon); + }); +} + +#[test] +#[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] +fn content_type_cached_and_uncached_agree() { + bounded!().for_each(|input: &[u8]| { + let value = shaped_field_value(input); + let map = map_with_values::(slice::from_ref(&value)); + let first_view = content_type_view_outcome(&map); + let repeated_view = content_type_view_outcome(&map); + let first_owned = content_type_owned_outcome(&map); + let repeated_owned = content_type_owned_outcome(&map); + + if let ContentTypeOutcome::Parsed { type_, subtype, .. } = &first_view + && let Ok(wire) = str::from_utf8(value.as_bytes()) + && let Ok(reference) = wire.parse::() + { + let reference_subtype = reference.suffix().map_or_else( + || reference.subtype().as_str().to_owned(), + |suffix| format!("{}+{}", reference.subtype(), suffix), + ); + assert!(type_.eq_ignore_ascii_case(reference.type_().as_str())); + assert!(subtype.eq_ignore_ascii_case(&reference_subtype)); + } + assert_ne!(first_view, ContentTypeOutcome::Absent); + assert_eq!(repeated_view, first_view); + assert_eq!(first_owned, first_view); + assert_eq!(repeated_owned, first_view); + }); +} + +#[test] +#[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] +fn scalar_and_simd_agree() { + bounded!().with_type::<(Vec, Vec)>().for_each(|(left, right)| { + let left = &left[..left.len().min(MAX_VALUE_LENGTH)]; + let right = &right[..right.len().min(MAX_VALUE_LENGTH)]; + let token = !left.is_empty() + && left + .iter() + .copied() + .all(|byte| byte.is_ascii_alphanumeric() || b"!#$%&'*+-.^_`|~".contains(&byte)); + let token68_data_end = left.iter().position(|byte| *byte == b'=').unwrap_or(left.len()); + let token68 = token68_data_end != 0 + && left[..token68_data_end] + .iter() + .all(|byte| byte.is_ascii_alphanumeric() || b"-._~+/".contains(byte)) + && left[token68_data_end..].iter().all(|byte| *byte == b'='); + let field_value = left.iter().copied().all(|byte| byte == b'\t' || byte >= b' ' && byte != 0x7f); + let equal = left.len() == right.len() && left.iter().zip(right).all(|(left, right)| left.eq_ignore_ascii_case(right)); + let interesting = left.iter().position(|byte| b",;\"\\ \t".contains(byte)); + + assert_eq!(http_headers_simd::is_token(left), token); + assert_eq!(http_headers_simd::is_token68(left), token68); + assert_eq!(http_headers_simd::is_field_value(left), field_value); + assert_eq!(http_headers_simd::eq_ignore_ascii_case(left, right), equal); + assert_eq!(http_headers_simd::find_interesting(left), interesting); + }); +} + +#[test] +#[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] +fn typed_round_trips() { + bounded!().for_each(|input: &[u8]| { + let number = input + .iter() + .copied() + .take(8) + .fold(0_u64, |value, byte| value.rotate_left(8) ^ u64::from(byte)); + + let content_length = ContentLengthOwned::new(number); + assert_eq!(content_length.get(), number); + assert_constructed_round_trip::(content_length); + + let opaque = opaque_from(input); + let etag = if input.first().is_some_and(|byte| byte & 1 == 1) { + ETagOwned::weak(&opaque).expect("generated opaque tags are valid") + } else { + ETagOwned::strong(&opaque).expect("generated opaque tags are valid") + }; + assert_eq!(etag.opaque_tag(), Ok(opaque.as_bytes())); + assert_constructed_round_trip::(etag); + + let bearer_token = token68_from(input); + let bearer = AuthorizationOwned::::bearer(&bearer_token).expect("generated token68 is valid"); + assert_eq!(bearer.token(), Ok(bearer_token.as_bytes())); + assert_constructed_round_trip::>(bearer); + + let mut username = bounded_bytes(input, 0); + username.retain(|byte| *byte != b':'); + let password = bounded_bytes(input, 64); + let basic = AuthorizationOwned::::basic(&username, &password).expect("bounded generated credentials fit"); + let mut map = TestMap::new(); + Authorization::::insert(&mut map, basic).expect("an empty header map has capacity"); + let authorization = Authorization::::view(&map) + .expect("constructed Basic authorization must decode") + .expect("constructed Basic authorization must be present"); + let mut credentials = BasicCredentials::new(); + let credentials = authorization + .extract(&mut credentials) + .expect("constructed Basic credentials must extract"); + assert_eq!(credentials.username(), username); + assert_eq!(credentials.password(), password); + let basic = Authorization::::owned(&map) + .expect("constructed Basic authorization must decode") + .expect("constructed Basic authorization must be present"); + assert_constructed_round_trip::>(basic); + + let type_ = token_from(input, 17); + let subtype = token_from(input, 93); + let parameter_text = quoted_text_from(input); + let parameter_wire = format!("\"{parameter_text}\""); + let content_type = + ContentTypeOwned::try_from(format!("{type_}/{subtype}; x-note={parameter_wire}")).expect("generated media type is valid"); + assert_eq!(content_type.type_(), Ok(type_.as_str())); + assert_eq!(content_type.subtype(), Ok(subtype.as_str())); + assert_eq!(content_type.parameter("X-NOTE"), Ok(Some(parameter_wire.as_bytes()))); + assert_constructed_round_trip::(content_type); + + let cache_control = CacheControlOwned::builder() + .no_cache() + .max_age(Duration::from_secs(number)) + .extension_value("x-note", ¶meter_wire) + .build() + .expect("the builder contains directives"); + assert!(cache_control.no_cache()); + assert_eq!(cache_control.max_age(), Some(Duration::from_secs(number))); + assert!( + cache_control + .directives() + .any(|directive| { directive.name() == "x-note" && directive.value() == Some(parameter_wire.as_bytes()) }) + ); + assert_constructed_round_trip::(cache_control); + + let low_byte = number.to_le_bytes()[0]; + let location_wire = format!("/resource/{number:016x}?q={low_byte:02x}#section"); + let location = LocationOwned::try_from(location_wire.clone()).expect("generated URI is valid"); + assert_eq!(location.as_str(), Ok(location_wire.as_str())); + assert_constructed_round_trip::(location); + + let first_cookie = format!("a={number:016x}; Path=/"); + let second_cookie = format!("b={low_byte:02x}; HttpOnly"); + let mut cookies = SetCookieOwned::new(); + cookies.push_str(&first_cookie).expect("generated cookie is nonempty"); + cookies.push_str(&second_cookie).expect("generated cookie is nonempty"); + assert_eq!( + cookies.iter().map(FieldValue::as_bytes).collect::>(), + vec![first_cookie.as_bytes(), second_cookie.as_bytes()] + ); + assert_constructed_round_trip::(cookies); + + let user_agent_wire = format!("bolero-agent/{number:016x}"); + let user_agent = UserAgentOwned::try_from(user_agent_wire.clone()).expect("generated user agent is valid"); + assert_eq!(user_agent.as_bytes(), user_agent_wire.as_bytes()); + assert_constructed_round_trip::(user_agent); + }); +} + +#[test] +#[expect( + clippy::too_many_lines, + reason = "one corpus-style test keeps the deterministic protocol seed matrix together" +)] +fn protocol_regression_seeds() { + let bearer_map = map_with_values::>(&[FieldValue::from_static("Bearer mF_9.B5f-4.1JqM")]); + let bearer = Authorization::::view(&bearer_map) + .expect("RFC bearer example is valid") + .expect("authorization is present"); + assert_eq!(bearer.token(), b"mF_9.B5f-4.1JqM"); + + let basic_map = map_with_values::>(&[FieldValue::from_static("Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ==")]); + let authorization = Authorization::::view(&basic_map) + .expect("RFC Basic example is valid") + .expect("authorization is present"); + let mut credentials = BasicCredentials::new(); + let basic = authorization.extract(&mut credentials).expect("RFC Basic credentials extract"); + assert_eq!(basic.username(), b"Aladdin"); + assert_eq!(basic.password(), b"open sesame"); + + let cache_map = map_with_values::(&[FieldValue::from_static("no-cache, private, max-age=60, x-note=\"a,b\"")]); + let cache = CacheControl::view(&cache_map) + .expect("representative Cache-Control is valid") + .expect("cache-control is present"); + assert!(cache.no_cache()); + assert_eq!(cache.max_age(), Some(Duration::from_mins(1))); + assert!( + cache + .directives() + .any(|directive| { directive.name() == "x-note" && directive.value() == Some(b"\"a,b\"".as_slice()) }) + ); + + let length_map = map_with_values::(&[FieldValue::from_static("42, 42")]); + assert_eq!(ContentLength::view(&length_map), Ok(Some(ContentLengthOwned::new(42)))); + + let content_type_map = map_with_values::(&[FieldValue::from_static("text/html; charset=utf-8; boundary=\"a,b\"")]); + let content_type = ContentType::view(&content_type_map) + .expect("representative Content-Type is valid") + .expect("content-type is present"); + assert_eq!(content_type.type_(), Ok("text")); + assert_eq!(content_type.subtype(), Ok("html")); + assert_eq!(content_type.parameter("BOUNDARY"), Ok(Some(b"\"a,b\"".as_slice()))); + + let strong_map = map_with_values::(&[FieldValue::from_static("\"xyzzy\"")]); + let strong = ETag::view(&strong_map) + .expect("strong entity tag is valid") + .expect("etag is present"); + assert!(!strong.is_weak()); + assert_eq!(strong.opaque_tag(), b"xyzzy"); + let weak_map = map_with_values::(&[FieldValue::from_static("W/\"xyzzy\"")]); + let weak = ETag::view(&weak_map).expect("weak entity tag is valid").expect("etag is present"); + assert!(weak.is_weak()); + assert!(strong.weak_eq(weak)); + + let location_wire = "https://example.com/a%20b?q=x#frag"; + let location_map = map_with_values::(&[FieldValue::from_static(location_wire)]); + let location = Location::view(&location_map) + .expect("representative Location is valid") + .expect("location is present"); + assert_eq!(location.as_str(), Ok(location_wire)); + + let cookie_values = [ + FieldValue::from_static("session=abc; Path=/; Secure; HttpOnly"), + FieldValue::from_static("theme=light; SameSite=Lax"), + ]; + let cookie_map = map_with_values::(&cookie_values); + let cookies = SetCookie::view(&cookie_map) + .expect("representative cookies are valid") + .expect("set-cookie is present"); + assert_eq!(cookies.len(), 2); + assert_eq!( + cookies.iter().map(FieldValueRef::as_bytes).collect::>(), + cookie_values.iter().map(FieldValue::as_bytes).collect::>() + ); + + let user_agent_wire = "ExampleBrowser/1.0 (compatible; TestBot/2.0)"; + let user_agent_map = map_with_values::(&[FieldValue::from_static(user_agent_wire)]); + let user_agent = UserAgent::view(&user_agent_map) + .expect("representative User-Agent is valid") + .expect("user-agent is present"); + assert_eq!(user_agent.as_str(), Ok(user_agent_wire)); + + let values = [ + FieldValue::from_static("a, \"b,c\", d"), + FieldValue::from_static("x=\"escaped\\\" comma, stays\"; y=z"), + ]; + assert_eq!( + reference_delimited(&values, b','), + Ok(vec![ + b"a".to_vec(), + b"\"b,c\"".to_vec(), + b"d".to_vec(), + b"x=\"escaped\\\" comma, stays\"; y=z".to_vec(), + ]) + ); + assert_eq!(delimited_outcome::(&values), reference_delimited(&values, b',')); + + let malformed = [FieldValue::from_static("a, \"unterminated")]; + assert_eq!( + delimited_outcome::(&malformed), + Err((DecodeErrorKind::UnterminatedQuote, Some(0))) + ); + + let item_limit_before_quote = [ + FieldValue::from_bytes(vec![b';'; MAX_CUSTOM_LIST_ITEMS - 1]).unwrap(), + FieldValue::from_static(";\""), + ]; + assert_eq!( + reference_delimited(&item_limit_before_quote, b';'), + Err((DecodeErrorKind::SourceLimitExceeded, None)) + ); + assert_eq!( + delimited_outcome::(&item_limit_before_quote), + Err((DecodeErrorKind::SourceLimitExceeded, None)) + ); + + let all_empty = [FieldValue::from_static(",, ,\t,,")]; + assert_eq!(delimited_outcome::(&all_empty), Ok(Vec::new())); + assert_eq!( + owned_outcome::(&map_with_values::(&all_empty)), + ParseOutcome::Parsed(Vec::new()) + ); + + let obs_text = FieldValue::from_bytes([0x80]).expect("obs-text is a valid field value"); + let obs_user_agent_map = map_with_values::(slice::from_ref(&obs_text)); + let obs_user_agent = UserAgent::view(&obs_user_agent_map) + .expect("obs-text User-Agent bytes are a valid field value") + .expect("user-agent is present"); + assert_eq!( + obs_user_agent.as_str().map_err(|error| error.kind()), + Err(DecodeErrorKind::InvalidUtf8) + ); + assert_eq!( + owned_outcome::(&map_with_values::(slice::from_ref(&obs_text))), + ParseOutcome::Rejected(DecodeErrorKind::InvalidSyntax, None) + ); + + let obs_etag = FieldValue::from_bytes([b'"', 0x80, b'"']).expect("obs-text entity tag is valid"); + let obs_etag_map = map_with_values::(&[obs_etag]); + assert_eq!( + ETag::view(&obs_etag_map) + .expect("obs-text entity tag is valid") + .expect("etag is present") + .opaque_tag(), + &[0x80] + ); + + let obs_content_type = FieldValue::from_bytes(b"text/plain; note=\"\x80\"").expect("obs-text quoted parameters are valid field values"); + let obs_content_type_map = map_with_values::(&[obs_content_type]); + let content_type = ContentType::view(&obs_content_type_map) + .expect("obs-text quoted parameter is valid") + .expect("content-type is present"); + assert_eq!(content_type.parameter("note"), Ok(Some(b"\"\x80\"".as_slice()))); + + let obs_cache_control = FieldValue::from_bytes(b"x-note=\"\x80\"").expect("obs-text quoted directive is a valid field value"); + assert_eq!( + owned_outcome::(&map_with_values::(&[obs_cache_control])), + ParseOutcome::Parsed(vec![b"x-note=\"\x80\"".to_vec()]) + ); + + assert_duplicate_singleton_rejected::>(FieldValue::from_static("Bearer abc")); + assert_duplicate_singleton_rejected::>(FieldValue::from_static("Basic dXNlcjpwYXNz")); + assert_duplicate_singleton_rejected::(FieldValue::from_static("text/plain")); + assert_duplicate_singleton_rejected::(FieldValue::from_static("\"tag\"")); + assert_duplicate_singleton_rejected::(FieldValue::from_static("/target")); + assert_duplicate_singleton_rejected::(FieldValue::from_static("agent/1")); + + let conflicting_lengths = [FieldValue::from_static("31"), FieldValue::from_static("32")]; + assert_eq!( + owned_outcome::(&map_with_values::(&conflicting_lengths)), + ParseOutcome::Rejected(DecodeErrorKind::InvalidSyntax, Some(1)) + ); + + assert_absence_differs_from_empty::>(); + assert_absence_differs_from_empty::>(); + assert_absence_differs_from_empty::(); + assert_absence_differs_from_empty::(); + assert_absence_differs_from_empty::(); + assert_absence_differs_from_empty::(); + assert_absence_differs_from_empty::(); + + assert_eq!( + owned_outcome::(&map_with_values::(&[FieldValue::from_static(""),])), + ParseOutcome::Parsed(Vec::new()) + ); + assert_eq!( + owned_outcome::(&map_with_values::(&[FieldValue::from_static("",)])), + ParseOutcome::Parsed(vec![Vec::new()]) + ); + + for length in [31, 32, 33, MAX_VALUE_LENGTH] { + let value = FieldValue::from_bytes(vec![b'a'; length]).expect("ASCII field value is valid"); + assert_parser_consistency::(slice::from_ref(&value)); + assert_parser_consistency::(slice::from_ref(&value)); + assert_eq!(http_headers_simd::find_interesting(value.as_bytes()), None); + } +} diff --git a/crates/http_headers/tests/collection_traits.rs b/crates/http_headers/tests/collection_traits.rs new file mode 100644 index 000000000..0d4a63354 --- /dev/null +++ b/crates/http_headers/tests/collection_traits.rs @@ -0,0 +1,179 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Collection-trait and hashing integration coverage. + +#![cfg(feature = "headers-all")] + +use std::collections::{HashMap, HashSet}; + +#[cfg(feature = "http")] +use http::{HeaderMap, HeaderValue, header}; +use http_headers::headers::{ + self, CacheControlOwned, ContentTypeOwned, ETagOwned, LocationOwned, SetCookie, SetCookieOwned, UserAgentOwned, +}; +use http_headers::sink::{EncodedValues, FieldSink}; +use http_headers::{DecodeError, Field, FieldSensitivity, FieldValue}; + +use self::common::TestMap; + +mod common; + +#[test] +fn encoded_values_support_collection_iteration() { + let mut encoded: EncodedValues = [FieldValue::from_static("first"), FieldValue::from_static("second")] + .into_iter() + .collect(); + + assert_eq!((&encoded).into_iter().len(), 2); + for value in &mut encoded { + value.set_sensitivity(FieldSensitivity::Sensitive); + } + assert!(encoded.iter().all(FieldValue::is_sensitive)); + assert_eq!(encoded.into_iter().len(), 2); +} + +#[test] +fn set_cookie_supports_safe_collection_iteration() { + let mut cookies = SetCookieOwned::new(); + for value in ["a=1", "b=2"] { + cookies.push_str(value).expect("valid cookie"); + } + + assert_eq!((&cookies).into_iter().len(), 2); + for value in &mut cookies { + value.set_sensitivity(FieldSensitivity::Sensitive); + } + let rebuilt = cookies + .into_iter() + .try_fold(SetCookieOwned::new(), |mut rebuilt, value| { + rebuilt.push(value)?; + Ok::<_, DecodeError>(rebuilt) + }) + .expect("stored values preserve the Set-Cookie invariant"); + assert_eq!(rebuilt.len(), 2); +} + +#[test] +fn set_cookie_public_helpers_cover_empty_invalid_and_view_states() { + let mut cookies = SetCookieOwned::default(); + assert!(cookies.is_empty()); + assert!(cookies.iter_mut().next().is_none()); + assert!(cookies.push_str("\n").is_err()); + + let parsed = "a=1".parse::().expect("valid cookie"); + assert!(!parsed.is_empty()); + assert!(format!("{parsed:?}").contains("value_count")); + + let mut table = TestMap::default(); + SetCookie::insert(&mut table, parsed).expect("table accepts cookie"); + let view = SetCookie::view(&table).expect("valid cookie view").expect("cookie is present"); + assert!(!view.is_empty()); + + table + .set_values(SetCookie::name(), EncodedValues::single(FieldValue::from_static(""))) + .expect("table accepts raw empty field value"); + SetCookie::view(&table).expect_err("empty cookie is invalid"); + SetCookie::owned(&table).expect_err("empty cookie is invalid"); + + #[cfg(feature = "http")] + { + let mut map = HeaderMap::new(); + SetCookie::insert(&mut map, "b=2".parse::().expect("valid cookie")).expect("HTTP map accepts cookie"); + assert!(SetCookie::view(&map).expect("valid HTTP view").is_some()); + assert!(SetCookie::owned(&map).expect("valid HTTP value").is_some()); + map.insert(header::SET_COOKIE, HeaderValue::from_static("")); + SetCookie::view(&map).expect_err("empty cookie is invalid"); + SetCookie::owned(&map).expect_err("empty cookie is invalid"); + } +} + +#[test] +fn shared_from_str_error_mapping_covers_foundational_headers() { + "\n".parse::().expect_err("invalid field"); + "\n".parse::().expect_err("invalid field"); + "\n".parse::().expect_err("invalid field"); + "\n".parse::().expect_err("invalid field"); + "\n".parse::().expect_err("invalid field"); +} + +#[test] +fn shared_from_str_error_mapping_covers_every_generated_impl() { + macro_rules! assert_invalid { + ($($owned:path),+ $(,)?) => { + $("\n".parse::<$owned>().expect_err("invalid field");)+ + }; + } + + assert_invalid!( + headers::AcceptOwned, + headers::AcceptEncodingOwned, + headers::AcceptLanguageOwned, + headers::AcceptRangesOwned, + headers::AccessControlAllowCredentialsOwned, + headers::AccessControlAllowHeadersOwned, + headers::AccessControlAllowMethodsOwned, + headers::AccessControlAllowOriginOwned, + headers::AccessControlExposeHeadersOwned, + headers::AccessControlMaxAgeOwned, + headers::AccessControlRequestHeadersOwned, + headers::AccessControlRequestMethodOwned, + headers::AllowOwned, + headers::CacheControlOwned, + headers::ContentRangeOwned, + headers::ContentSecurityPolicyOwned, + headers::ContentTypeOwned, + headers::ETagOwned, + headers::HostOwned, + headers::IfMatchOwned, + headers::IfModifiedSinceOwned, + headers::IfNoneMatchOwned, + headers::IfRangeOwned, + headers::IfUnmodifiedSinceOwned, + headers::LastModifiedOwned, + headers::LocationOwned, + headers::RangeOwned, + headers::ReferrerPolicyOwned, + headers::SecWebSocketAcceptOwned, + headers::SecWebSocketExtensionsOwned, + headers::SecWebSocketKeyOwned, + headers::SecWebSocketProtocolOwned, + headers::SecWebSocketVersionOwned, + headers::ServerOwned, + headers::StrictTransportSecurityOwned, + headers::UserAgentOwned, + headers::VaryOwned, + headers::XContentTypeOptionsOwned, + ); +} + +#[test] +fn shared_ascii_display_covers_every_generated_impl() { + macro_rules! assert_display { + ($owned:path, $wire:literal) => {{ + let value = $wire.parse::<$owned>().expect("valid display value"); + assert_eq!(value.to_string(), $wire); + }}; + } + + assert_display!(headers::ContentRangeOwned, "bytes 0-1/2"); + assert_display!(headers::HostOwned, "example.com"); + assert_display!(headers::IfModifiedSinceOwned, "Sun, 06 Nov 1994 08:49:37 GMT"); + assert_display!(headers::IfUnmodifiedSinceOwned, "Sun, 06 Nov 1994 08:49:37 GMT"); + assert_display!(headers::LastModifiedOwned, "Sun, 06 Nov 1994 08:49:37 GMT"); + assert_display!(headers::RangeOwned, "bytes=0-1"); + assert_display!(headers::SecWebSocketAcceptOwned, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + assert_display!(headers::SecWebSocketKeyOwned, "dGhlIHNhbXBsZSBub25jZQ=="); + assert_display!(headers::XContentTypeOptionsOwned, "nosniff"); +} + +#[test] +fn public_equal_types_are_hashable() { + let mut tags = HashSet::new(); + tags.insert(ETagOwned::strong("revision").expect("valid entity tag")); + assert_eq!(tags.len(), 1); + + let mut media = HashMap::new(); + media.insert(ContentTypeOwned::try_from("application/json").expect("valid media type"), "json"); + assert_eq!(media.len(), 1); +} diff --git a/crates/http_headers/tests/common/http_headers_name_corpus.rs b/crates/http_headers/tests/common/http_headers_name_corpus.rs new file mode 100644 index 000000000..d43702da1 --- /dev/null +++ b/crates/http_headers/tests/common/http_headers_name_corpus.rs @@ -0,0 +1,106 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Immutable name corpora shared by recognition benchmarks and ownership tests. +//! +//! Each parsed corpus is initialized once, so repeated setup reuses its owned names. + +use std::sync::LazyLock; + +use http_headers::FieldName; + +/// The field names of an ordinary browser request, in wire case. +/// +/// The set spans the length range of the well-known table, from the shortest +/// entry to one of the longest, so no case is confined to one length bucket. +const KNOWN_NAMES: &[&[u8]] = &[ + b"Host", + b"User-Agent", + b"Accept", + b"Accept-Language", + b"Accept-Encoding", + b"Connection", + b"Cookie", + b"Referer", + b"Cache-Control", + b"Content-Type", + b"Content-Length", + b"Authorization", + b"TE", + b"Sec-WebSocket-Key", + b"Access-Control-Request-Headers", +]; + +/// Vendor and tracing names no well-known entry matches. +const CUSTOM_NAMES_LOWERCASE: &[&[u8]] = &[ + b"x-request-id", + b"x-forwarded-for", + b"x-forwarded-proto", + b"x-trace-id", + b"cf-ray", + b"x-amzn-trace-id", + b"x-correlation-id", + b"x-real-ip", +]; + +const CUSTOM_NAMES_MIXED_CASE: &[&[u8]] = &[ + b"X-Request-Id", + b"X-Forwarded-For", + b"X-Forwarded-Proto", + b"X-Trace-Id", + b"CF-Ray", + b"X-Amzn-Trace-Id", + b"X-Correlation-Id", + b"X-Real-Ip", +]; + +pub(super) fn known_names() -> &'static [&'static [u8]] { + KNOWN_NAMES +} + +pub(super) fn custom_names_lowercase() -> &'static [&'static [u8]] { + assert_eq!(CUSTOM_NAMES_LOWERCASE.len(), CUSTOM_NAMES_MIXED_CASE.len()); + assert!( + CUSTOM_NAMES_LOWERCASE + .iter() + .zip(CUSTOM_NAMES_MIXED_CASE) + .all(|(lowercase, mixed_case)| lowercase.eq_ignore_ascii_case(mixed_case)) + ); + CUSTOM_NAMES_LOWERCASE +} + +pub(super) fn custom_names_mixed_case() -> &'static [&'static [u8]] { + CUSTOM_NAMES_MIXED_CASE +} + +pub(super) fn custom_header_names() -> &'static [FieldName] { + static NAMES: LazyLock> = LazyLock::new(|| { + CUSTOM_NAMES_MIXED_CASE + .iter() + .map(|name| FieldName::try_from_bytes(name).expect("valid field name")) + .collect() + }); + NAMES.as_slice() +} + +pub(super) fn http_names() -> &'static [http::HeaderName] { + static NAMES: LazyLock> = LazyLock::new(|| { + KNOWN_NAMES + .iter() + .chain(CUSTOM_NAMES_MIXED_CASE) + .map(|name| http::HeaderName::from_bytes(name).expect("valid field name")) + .collect() + }); + NAMES.as_slice() +} + +pub(super) fn crate_names() -> &'static [FieldName] { + static NAMES: LazyLock> = LazyLock::new(|| { + KNOWN_NAMES + .iter() + .chain(CUSTOM_NAMES_MIXED_CASE) + .map(|name| FieldName::try_from_bytes(name).expect("valid field name")) + .collect() + }); + NAMES.as_slice() +} diff --git a/crates/http_headers/tests/common/http_headers_storage_operations.rs b/crates/http_headers/tests/common/http_headers_storage_operations.rs new file mode 100644 index 000000000..c02bce7ba --- /dev/null +++ b/crates/http_headers/tests/common/http_headers_storage_operations.rs @@ -0,0 +1,94 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Storage benchmark operations shared with ownership regression tests. + +use std::hint::black_box; + +use http::HeaderMap; +use http_headers::headers::{ContentLength, ContentType}; +use http_headers::sink::{EncodedValues, FieldEncodeOutput, FieldEncoder, FieldSink, FieldSinkExt, FieldValueWriter, InsertError}; +use http_headers::{Field, FieldName, FieldValue, FieldValueRef}; + +pub(super) type AppendInput = (HeaderMap, &'static FieldName, EncodedValues); + +/// Keep the map alive until the benchmark harness ends measurement and drops its output. +pub(super) type MapOutput = (usize, HeaderMap); + +struct ChunkedEncoder { + bytes: &'static [u8], + chunks: usize, +} + +impl FieldEncoder for ChunkedEncoder { + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + let mut writer = output.begin_value(self.bytes.len(), http_headers::FieldSensitivity::NonSensitive)?; + if !self.bytes.is_empty() { + for piece in self.bytes.chunks(self.bytes.len().div_ceil(self.chunks)) { + writer.write_bytes(piece)?; + } + } + writer.finish() + } +} + +pub(super) fn http_writer_borrowed(state: (HeaderMap, &'static str)) -> MapOutput { + let (mut map, value) = state; + map.set_encoded(::name(), FieldValueRef::new(value.as_bytes())) + .expect("preallocated map has capacity"); + let length = black_box(map.len()); + (length, map) +} + +pub(super) fn http_writer_streamed(state: (HeaderMap, &'static str)) -> MapOutput { + let (mut map, value) = state; + map.set_encoded( + ::name(), + ChunkedEncoder { + bytes: value.as_bytes(), + chunks: 4, + }, + ) + .expect("preallocated map has capacity"); + let length = black_box(map.len()); + (length, map) +} + +pub(super) fn http_writer_streamed_sized(state: (HeaderMap, &'static [u8])) -> MapOutput { + let (mut map, value) = state; + map.set_encoded(::name(), ChunkedEncoder { bytes: value, chunks: 4 }) + .expect("preallocated map has capacity"); + let length = black_box(map.len()); + (length, map) +} + +pub(super) fn http_writer_materialized(state: (HeaderMap, FieldValue)) -> MapOutput { + let (mut map, value) = state; + map.set_values(::name(), EncodedValues::single(value)) + .expect("preallocated map has capacity"); + let length = black_box(map.len()); + (length, map) +} + +pub(super) fn http_append_values(input: AppendInput) -> MapOutput { + let (mut map, name, values) = input; + map.append_values(name, values).expect("preallocated map has capacity"); + let length = black_box(map.len()); + (length, map) +} + +pub(super) fn content_length_materialized(mut map: HeaderMap) -> MapOutput { + map.set_values(::name(), EncodedValues::single(FieldValue::from(1_024_u64))) + .expect("preallocated map has capacity"); + let length = black_box(map.len()); + (length, map) +} + +pub(super) fn content_length_deferred_http(mut map: HeaderMap) -> MapOutput { + map.set_content_length(1_024).expect("preallocated map has capacity"); + let length = black_box(map.len()); + (length, map) +} diff --git a/crates/http_headers/tests/common/mod.rs b/crates/http_headers/tests/common/mod.rs new file mode 100644 index 000000000..d8dde860e --- /dev/null +++ b/crates/http_headers/tests/common/mod.rs @@ -0,0 +1,6 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +mod test_map; + +pub(crate) use test_map::TestMap; diff --git a/crates/http_headers/tests/common/test_map.rs b/crates/http_headers/tests/common/test_map.rs new file mode 100644 index 000000000..c7d1a67c6 --- /dev/null +++ b/crates/http_headers/tests/common/test_map.rs @@ -0,0 +1,43 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +use std::collections::HashMap; + +use http_headers::sink::{EncodedValues, FieldSink, InsertError}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{FieldName, FieldValue}; + +#[derive(Default)] +pub(crate) struct TestMap(pub(crate) HashMap>); + +impl TestMap { + pub(crate) fn get_all(&self, name: &FieldName) -> &[FieldValue] { + self.0.get(name).map_or(&[], Vec::as_slice) + } +} + +impl FieldSource for TestMap { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_slice(name, self.get_all(name)) + } +} + +impl FieldSink for TestMap { + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + if values.is_empty() { + self.0.remove(name); + } else { + self.0.insert(name.clone(), values.into_iter().collect()); + } + Ok(()) + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.0.entry(name.clone()).or_default().extend(values); + Ok(()) + } + + fn remove_values(&mut self, name: &'static FieldName) { + self.0.remove(name); + } +} diff --git a/crates/http_headers/tests/conformance_api.rs b/crates/http_headers/tests/conformance_api.rs new file mode 100644 index 000000000..00f45755c --- /dev/null +++ b/crates/http_headers/tests/conformance_api.rs @@ -0,0 +1,424 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Coverage for public API conformance guarantees. + +#![cfg(feature = "headers-all")] + +use std::fmt::{self, Write as _}; +use std::time::Duration; + +use http_headers::headers::*; +use http_headers::sink::{EncodedValues, FieldSink, InsertError, InsertErrorKind, ValueRefsEncoder}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeError, FieldName, FieldSensitivity, FieldValue, FieldValueRef}; + +macro_rules! assert_field_value_conversion { + ($owned:ty, $wire:literal) => {{ + let from = FieldValue::from(<$owned>::try_from($wire).expect("valid fixture")); + let into: FieldValue = <$owned>::try_from($wire).expect("valid fixture").into(); + let shim = <$owned>::try_from($wire).expect("valid fixture").into_field_value(); + assert_eq!(from, $wire); + assert_eq!(into, $wire); + assert_eq!(shim, $wire); + }}; +} + +struct RejectWriter; + +impl fmt::Write for RejectWriter { + fn write_str(&mut self, _value: &str) -> fmt::Result { + Err(fmt::Error) + } +} + +#[test] +fn all_21_owned_single_value_types_support_standard_conversion() { + assert_field_value_conversion!(StrictTransportSecurityOwned, "max-age=60"); + assert_field_value_conversion!(IfRangeOwned, "\"revision\""); + assert_field_value_conversion!(XContentTypeOptionsOwned, "nosniff"); + assert_field_value_conversion!(IfModifiedSinceOwned, "Sat, 29 Oct 1994 19:43:31 GMT"); + assert_field_value_conversion!(IfUnmodifiedSinceOwned, "Sat, 29 Oct 1994 19:43:31 GMT"); + assert_field_value_conversion!(LastModifiedOwned, "Sat, 29 Oct 1994 19:43:31 GMT"); + assert_field_value_conversion!(RangeOwned, "bytes=0-9"); + assert_field_value_conversion!(ContentRangeOwned, "bytes 0-9/10"); + assert_field_value_conversion!(LocationOwned, "https://example.com/people"); + assert_field_value_conversion!(ContentTypeOwned, "application/json"); + let bearer = AuthorizationOwned::::bearer("abc.def").expect("valid bearer credentials"); + let bearer_from = FieldValue::from(bearer.clone()); + let bearer_into: FieldValue = bearer.clone().into(); + let bearer_shim = bearer.into_field_value(); + for value in [&bearer_from, &bearer_into, &bearer_shim] { + assert_eq!(value.as_bytes(), b"Bearer abc.def"); + assert!(value.is_sensitive()); + } + + let basic = AuthorizationOwned::::basic(b"user", b"pass").expect("valid basic credentials"); + let basic_from = FieldValue::from(basic.clone()); + let basic_into: FieldValue = basic.clone().into(); + let basic_shim = basic.into_field_value(); + for value in [&basic_from, &basic_into, &basic_shim] { + assert_eq!(value.as_bytes(), b"Basic dXNlcjpwYXNz"); + assert!(value.is_sensitive()); + } + assert_field_value_conversion!(ETagOwned, "\"revision\""); + assert_field_value_conversion!(UserAgentOwned, "client/1"); + assert_field_value_conversion!(HostOwned, "example.com"); + assert_field_value_conversion!(ServerOwned, "example/1"); + assert_field_value_conversion!(SecWebSocketKeyOwned, "dGhlIHNhbXBsZSBub25jZQ=="); + assert_field_value_conversion!(SecWebSocketAcceptOwned, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + assert_field_value_conversion!(AccessControlAllowOriginOwned, "https://example.com"); + assert_field_value_conversion!(AccessControlAllowCredentialsOwned, "true"); + assert_field_value_conversion!(AccessControlMaxAgeOwned, "600"); + assert_field_value_conversion!(AccessControlRequestMethodOwned, "DELETE"); +} + +#[test] +fn collection_displays_propagate_formatter_failures() { + let ranges = AcceptRangesOwned::from_units(["bytes"]).expect("valid range unit"); + assert!(write!(&mut RejectWriter, "{ranges}").is_err()); + + let methods = AccessControlAllowMethodsOwned::from_methods(["GET"]).expect("valid method"); + assert!(write!(&mut RejectWriter, "{methods}").is_err()); +} + +#[test] +fn all_11_ascii_owned_types_display_canonical_values() { + assert_eq!(AllowOwned::try_from("GET, HEAD").expect("valid methods").to_string(), "GET, HEAD"); + assert_eq!( + VaryOwned::try_from("accept, origin").expect("valid field names").to_string(), + "accept, origin" + ); + assert_eq!( + AcceptRangesOwned::from_units(["bytes", "items"]) + .expect("valid range units") + .to_string(), + "bytes, items" + ); + assert_eq!( + AccessControlAllowHeadersOwned::from_header_names(["content-type", "x-trace-id"]) + .expect("valid header names") + .to_string(), + "content-type, x-trace-id" + ); + assert_eq!( + AccessControlAllowMethodsOwned::from_methods(["GET", "POST"]) + .expect("valid methods") + .to_string(), + "GET, POST" + ); + assert_eq!( + AccessControlExposeHeadersOwned::from_header_names(["etag", "x-trace-id"]) + .expect("valid header names") + .to_string(), + "etag, x-trace-id" + ); + assert_eq!( + AccessControlRequestHeadersOwned::from_header_names(["content-type", "authorization"]) + .expect("valid header names") + .to_string(), + "content-type, authorization" + ); + assert_eq!( + ReferrerPolicyOwned::new(ReferrerPolicyValue::NoReferrer) + .with_fallback(ReferrerPolicyValue::StrictOrigin) + .to_string(), + "no-referrer, strict-origin" + ); + assert_eq!( + SecWebSocketExtensionsOwned::builder() + .extension("permessage-deflate") + .parameter_flag("client_max_window_bits") + .build() + .expect("nonempty extension list") + .to_string(), + "permessage-deflate; client_max_window_bits" + ); + assert_eq!( + SecWebSocketProtocolOwned::new("chat") + .expect("valid protocol") + .with_protocol("superchat") + .expect("valid protocol") + .to_string(), + "chat, superchat" + ); + assert_eq!(SecWebSocketVersionOwned::new(13).with_version(8).to_string(), "8, 13"); +} + +struct EmptySource; + +impl FieldSource for EmptySource { + fn lines(&self, _name: &'static FieldName) -> Option> { + None + } +} + +#[derive(Default)] +struct Sink; + +impl FieldSource for Sink { + fn lines(&self, _name: &'static FieldName) -> Option> { + None + } +} + +impl FieldSink for Sink { + fn set_values(&mut self, _name: &'static FieldName, _values: EncodedValues) -> Result<(), InsertError> { + Ok(()) + } + + fn append_values(&mut self, _name: &'static FieldName, _values: EncodedValues) -> Result<(), InsertError> { + Ok(()) + } + + fn remove_values(&mut self, _name: &'static FieldName) {} +} + +fn assert_debug() {} + +macro_rules! assert_inherent_operations { + ($(($header:ty, $owned:ty)),+ $(,)?) => { + $( + assert_debug::<$header>(); + assert!( + <$header>::view(&EmptySource) + .expect("absence is valid") + .is_none() + ); + assert!( + <$header>::owned(&EmptySource) + .expect("absence is valid") + .is_none() + ); + let _insert: fn(&mut Sink, $owned) -> Result<(), InsertError> = + <$header>::insert::; + <$header>::remove(&mut Sink); + )+ + }; +} + +#[test] +fn all_42_built_in_instantiations_expose_inherent_operations() { + assert_inherent_operations!( + (CacheControl, CacheControlOwned), + (Accept, AcceptOwned), + (AcceptEncoding, AcceptEncodingOwned), + (AcceptLanguage, AcceptLanguageOwned), + (Allow, AllowOwned), + (Host, HostOwned), + (Server, ServerOwned), + (Vary, VaryOwned), + (AcceptRanges, AcceptRangesOwned), + (ContentRange, ContentRangeOwned), + (Range, RangeOwned), + (ETag, ETagOwned), + (Location, LocationOwned), + (UserAgent, UserAgentOwned), + (ContentType, ContentTypeOwned), + (IfMatch, IfMatchOwned), + (IfNoneMatch, IfNoneMatchOwned), + (IfModifiedSince, IfModifiedSinceOwned), + (IfUnmodifiedSince, IfUnmodifiedSinceOwned), + (IfRange, IfRangeOwned), + (LastModified, LastModifiedOwned), + (AccessControlAllowCredentials, AccessControlAllowCredentialsOwned), + (AccessControlAllowHeaders, AccessControlAllowHeadersOwned), + (AccessControlAllowMethods, AccessControlAllowMethodsOwned), + (AccessControlAllowOrigin, AccessControlAllowOriginOwned), + (AccessControlExposeHeaders, AccessControlExposeHeadersOwned), + (AccessControlMaxAge, AccessControlMaxAgeOwned), + (AccessControlRequestHeaders, AccessControlRequestHeadersOwned), + (AccessControlRequestMethod, AccessControlRequestMethodOwned), + (ContentSecurityPolicy, ContentSecurityPolicyOwned), + (ReferrerPolicy, ReferrerPolicyOwned), + (StrictTransportSecurity, StrictTransportSecurityOwned), + (XContentTypeOptions, XContentTypeOptionsOwned), + (SecWebSocketAccept, SecWebSocketAcceptOwned), + (SecWebSocketExtensions, SecWebSocketExtensionsOwned), + (SecWebSocketKey, SecWebSocketKeyOwned), + (SecWebSocketProtocol, SecWebSocketProtocolOwned), + (SecWebSocketVersion, SecWebSocketVersionOwned), + (Authorization, AuthorizationOwned), + (Authorization, AuthorizationOwned), + (SetCookie, SetCookieOwned), + (ContentLength, ContentLengthOwned), + ); +} + +fn generic_owned(source: &EmptySource) -> Result, DecodeError> { + H::owned(source) +} + +#[test] +fn generic_field_behavior_is_retained() { + assert!(generic_owned::(&EmptySource).expect("absence is valid").is_none()); +} + +#[test] +fn borrowed_byte_input_constructors_accept_owned_values() { + assert_eq!( + FieldName::try_from_bytes(Vec::from(&b"Accept"[..])).expect("valid name"), + FieldName::Accept + ); + assert_eq!( + FieldValue::from_bytes(Vec::from(&b"gzip"[..])).expect("valid bytes").as_bytes(), + b"gzip" + ); + assert_eq!( + AuthorizationOwned::::basic(Vec::from(&b"user"[..]), Vec::from(&b"pass"[..])) + .expect("valid credentials") + .encoded_credentials(), + Ok(b"dXNlcjpwYXNz".as_slice()) + ); + assert_eq!( + ContentSecurityPolicyOwned::from_bytes(Vec::from(&b"img-src *"[..])) + .expect("valid policy") + .policies() + .next(), + Some(b"img-src *".as_slice()) + ); +} + +#[test] +fn borrowed_text_input_constructors_accept_owned_values() { + assert_eq!( + FieldValue::from_str(String::from("gzip")).expect("valid string").as_bytes(), + b"gzip" + ); + assert_eq!( + FieldValue::from(ContentTypeOwned::new(String::from("application"), String::from("json")).expect("valid media type"),).as_bytes(), + b"application/json" + ); + assert_eq!( + AuthorizationOwned::::bearer(String::from("abc.def")) + .expect("valid token") + .token(), + Ok(b"abc.def".as_slice()) + ); + assert_eq!( + ContentSecurityPolicyOwned::new(String::from("default-src 'self'")) + .expect("valid policy") + .policies() + .next(), + Some(b"default-src 'self'".as_slice()) + ); + assert_eq!( + RangeOwned::extension(String::from("items"), String::from("1-5")) + .expect("valid extension") + .as_field_value() + .as_bytes(), + b"items=1-5" + ); + assert_eq!( + ContentRangeOwned::extension(String::from("items"), String::from("0-9/100")) + .expect("valid extension") + .as_field_value() + .as_bytes(), + b"items 0-9/100" + ); + assert_eq!( + AccessControlAllowOriginOwned::from_origin(String::from("https://example.com")) + .expect("valid origin") + .origin(), + Ok(Some("https://example.com")) + ); + assert!(!ETagOwned::strong(String::from("revision")).expect("valid tag").is_weak()); + assert!(ETagOwned::weak(String::from("revision")).expect("valid tag").is_weak()); + assert!( + ETagOwned::try_from_wire(String::from("W/\"revision\"")) + .expect("valid wire tag") + .is_weak() + ); + assert_eq!( + SecWebSocketProtocolOwned::new(String::from("chat")) + .expect("valid protocol") + .selected(), + Ok("chat") + ); + assert_eq!( + HostOwned::new(String::from("example.com")).expect("valid host").as_str(), + Ok("example.com") + ); + assert_eq!( + HostOwned::with_port(String::from("example.com"), 443).expect("valid host").as_str(), + Ok("example.com:443") + ); +} + +#[test] +fn semantic_arguments_and_build_time_validation_preserve_output() { + let known = ContentRangeOwned::bytes(0, 9, CompleteLength::from(10)).expect("valid byte range"); + assert_eq!(known.as_field_value().as_bytes(), b"bytes 0-9/10"); + let unknown = ContentRangeOwned::bytes(0, 9, CompleteLength::Unknown).expect("valid byte range"); + assert_eq!(unknown.as_field_value().as_bytes(), b"bytes 0-9/*"); + + let cache = CacheControlOwned::builder() + .extension("stale-if-error", ExtensionValue::Value("30")) + .extension_flag("immutable-extension") + .build() + .expect("valid cache extensions"); + let cache_directives = cache + .directives() + .map(|directive| directive.as_bytes().to_vec()) + .collect::>(); + assert_eq!( + cache_directives, + [b"stale-if-error=30".as_slice(), b"immutable-extension"].map(<[u8]>::to_vec) + ); + CacheControlOwned::builder() + .extension_value("bad name", "bad value") + .build() + .expect_err("malformed cache extension must fail at build"); + assert_eq!( + Sink.set_encoded( + &FieldName::CacheControl, + CacheControlOwned::builder().extension_value("bad name", "bad value"), + ), + Err(InsertError::new(InsertErrorKind::InvalidValue)) + ); + + Sink.set_encoded( + &FieldName::SetCookie, + ValueRefsEncoder::new([FieldValueRef::new(b"a=1")]).with_sensitivity(FieldSensitivity::Sensitive), + ) + .expect("sensitive value references encode"); + SetCookie::insert(&mut Sink, SetCookieOwned::new()).expect("an empty cookie collection removes the field"); + + let hsts = StrictTransportSecurityOwned::builder(Duration::from_mins(1)) + .extension("report", ExtensionValue::Value("audit")) + .extension_flag("future") + .build() + .expect("valid HSTS extensions"); + assert_eq!(hsts.as_field_value().as_bytes(), b"max-age=60; report=audit; future"); + StrictTransportSecurityOwned::builder(Duration::from_mins(1)) + .extension_flag("preload") + .build() + .expect_err("reserved HSTS extension must fail at build"); + + let websocket = SecWebSocketExtensionsOwned::builder() + .extension("permessage-deflate") + .parameter("server_max_window_bits", ExtensionValue::Value("15")) + .quoted_parameter("mode", "fast") + .build() + .expect("valid WebSocket extension"); + assert_eq!( + websocket.to_string(), + "permessage-deflate; server_max_window_bits=15; mode=\"fast\"" + ); + SecWebSocketExtensionsOwned::builder() + .parameter_flag("orphan") + .build() + .expect_err("orphan parameter must fail at build"); + + CacheControlOwned::builder() + .extension_flag("x=y") + .build() + .expect_err("flag names containing equals must not become valued extensions"); + StrictTransportSecurityOwned::builder(Duration::from_mins(1)) + .extension_flag("x=y") + .build() + .expect_err("flag names containing equals must not become valued extensions"); + + assert!(!FieldSensitivity::NonSensitive.is_sensitive()); + assert!(FieldSensitivity::Sensitive.is_sensitive()); +} diff --git a/crates/http_headers/tests/content_type_identity.rs b/crates/http_headers/tests/content_type_identity.rs new file mode 100644 index 000000000..8040b9c59 --- /dev/null +++ b/crates/http_headers/tests/content_type_identity.rs @@ -0,0 +1,128 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Content-Type equality and hashing independent of construction and metadata. + +#![cfg(feature = "headers-content-type")] + +use std::collections::HashSet; +use std::hash::{DefaultHasher, Hash, Hasher}; + +use http_headers::headers::{ContentType, ContentTypeOwned}; +use http_headers::{FieldSensitivity, FieldValue}; + +use self::common::TestMap; + +mod common; + +const COMMON: [(&str, &str); 6] = [ + ("application", "json"), + ("application", "octet-stream"), + ("text", "html"), + ("text", "plain"), + ("text", "css"), + ("application", "javascript"), +]; + +fn fingerprint(value: &impl Hash) -> u64 { + let mut hasher = DefaultHasher::new(); + value.hash(&mut hasher); + hasher.finish() +} + +#[test] +fn every_common_constructor_matches_raw_parsing_reconstruction_and_hash_set_lookup() { + for (type_, subtype) in COMMON { + let wire = format!("{type_}/{subtype}"); + let field = FieldValue::from_str(&wire).unwrap(); + let constructed = ContentTypeOwned::new(type_, subtype).unwrap(); + let mut map = TestMap::default(); + ContentType::insert(&mut map, constructed.clone()).unwrap(); + let mut values = HashSet::from([constructed.clone()]); + for candidate in [ + ContentTypeOwned::try_from(wire.as_str()).unwrap(), + ContentTypeOwned::try_from(wire.clone()).unwrap(), + ContentTypeOwned::try_from(field.clone()).unwrap(), + ContentTypeOwned::try_from(constructed.clone().into_field_value()).unwrap(), + ContentType::owned(&map).unwrap().unwrap(), + ] { + assert_eq!(candidate, constructed); + assert_eq!(candidate.type_().unwrap(), type_); + assert_eq!(candidate.subtype().unwrap(), subtype); + assert_eq!(candidate.parameters().count(), 0); + assert_eq!(candidate.clone().into_field_value(), field); + assert_eq!(fingerprint(&candidate), fingerprint(&constructed)); + assert_eq!(fingerprint(&candidate), fingerprint(&field)); + assert!(values.contains(&candidate)); + assert!(!values.insert(candidate)); + } + assert_eq!(values.len(), 1); + #[cfg(feature = "serde")] + { + let json = serde_json::to_string(&constructed).unwrap(); + let decoded: ContentTypeOwned = serde_json::from_str(&json).unwrap(); + assert_eq!(decoded, constructed); + assert_eq!(fingerprint(&decoded), fingerprint(&constructed)); + assert!(values.contains(&decoded)); + assert_eq!(decoded.into_field_value(), field); + } + } +} + +#[test] +fn json_convenience_constructors_have_the_same_wire_identity() { + let constructed = ContentTypeOwned::new("application", "json").unwrap(); + for value in [ContentTypeOwned::json(), ContentType::json()] { + assert_eq!(value, constructed); + assert_eq!(fingerprint(&value), fingerprint(&constructed)); + assert_eq!(value.into_field_value().as_bytes(), b"application/json"); + } +} + +#[test] +fn content_type_identity_ignores_sensitivity_without_discarding_the_marker() { + for (type_, subtype) in COMMON { + let constructed = ContentTypeOwned::new(type_, subtype).unwrap(); + let sensitive_field = constructed.clone().into_field_value().with_sensitivity(FieldSensitivity::Sensitive); + let sensitive = ContentTypeOwned::try_from(sensitive_field.clone()).unwrap(); + assert_eq!(sensitive, constructed); + assert_eq!(fingerprint(&sensitive), fingerprint(&constructed)); + assert_eq!(fingerprint(&sensitive), fingerprint(&sensitive_field)); + assert!(sensitive.into_field_value().is_sensitive()); + assert!(!constructed.into_field_value().is_sensitive()); + } +} + +#[test] +fn content_type_identity_keeps_case_parameters_order_and_quoting_distinct() { + for (left, right) in [ + ("application/json", "Application/json"), + ("application/json", "application/JSON"), + ("application/json", "application/javascript"), + ("text/plain; charset=utf-8", "text/plain;charset=utf-8"), + ("text/plain; charset=utf-8", "text/plain; charset=UTF-8"), + ("text/plain;a=1;b=2", "text/plain;b=2;a=1"), + ("text/plain;x=\"a\"", "text/plain;x=a"), + ("application/json", "application/json; charset=utf-8"), + ] { + let left_field = FieldValue::from_str(left).unwrap(); + let right_field = FieldValue::from_str(right).unwrap(); + let left = ContentTypeOwned::try_from(left_field.clone()).unwrap(); + let right = ContentTypeOwned::try_from(right_field.clone()).unwrap(); + assert_ne!(left, right); + assert_eq!(fingerprint(&left), fingerprint(&left_field)); + assert_eq!(fingerprint(&right), fingerprint(&right_field)); + let values = HashSet::from([left.clone(), right.clone()]); + assert_eq!(values.len(), 2); + assert!(values.contains(&left)); + assert!(values.contains(&right)); + #[cfg(feature = "serde")] + for original in [&left, &right] { + let json = serde_json::to_string(original).unwrap(); + let decoded: ContentTypeOwned = serde_json::from_str(&json).unwrap(); + assert_eq!(&decoded, original); + assert_eq!(fingerprint(&decoded), fingerprint(original)); + assert!(values.contains(&decoded)); + } + } +} diff --git a/crates/http_headers/tests/core_api.rs b/crates/http_headers/tests/core_api.rs new file mode 100644 index 000000000..23a572882 --- /dev/null +++ b/crates/http_headers/tests/core_api.rs @@ -0,0 +1,621 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Public core API integration tests using downstream-defined headers. + +#![cfg(feature = "headers-all")] + +use std::sync::LazyLock; +use std::{panic, str}; + +#[cfg(feature = "http")] +use http::HeaderName; +use http_headers::headers::{ + AcceptRanges, AcceptRangesOwned, Authorization, AuthorizationOwned, Basic, BasicCredentials, ContentLength, ContentLengthOwned, + ContentType, ETag, ETagOwned, IfMatch, IfMatchOwned, IfNoneMatch, IfNoneMatchOwned, IfRange, IfRangeOwned, ReferrerPolicy, + ReferrerPolicyOwned, ReferrerPolicyValue, SecWebSocketVersion, SecWebSocketVersionOwned, SetCookie, SetCookieOwned, UserAgent, +}; +use http_headers::sink::{EncodedValues, FieldSink, FieldSinkExt, InsertError}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef, SingleValueField}; + +use self::common::TestMap; + +mod common; +#[cfg(all(miri, feature = "http"))] +#[path = "../src/miri_http_map.rs"] +mod miri_http_map; + +static REQUEST_ID: LazyLock = LazyLock::new(|| FieldName::from_static("x-request-id")); +static EMPTY: LazyLock = LazyLock::new(|| FieldName::from_static("x-empty")); + +impl TestMap { + fn new() -> Self { + Self::default() + } + + fn insert(&mut self, name: FieldName, value: FieldValue) { + self.0.insert(name, vec![value]); + } + + fn append(&mut self, name: FieldName, value: FieldValue) { + self.0.entry(name).or_default().push(value); + } + + fn contains_name(&self, name: &FieldName) -> bool { + self.0.contains_key(name) + } + + fn names(&self) -> impl Iterator { + self.0.keys() + } +} + +fn content_length_source(values: &[&'static str]) -> TestMap { + let mut source = TestMap::new(); + for (index, value) in values.iter().copied().enumerate() { + let value = FieldValue::from_static(value); + if index == 0 { + source.insert(FieldName::ContentLength, value); + } else { + source.append(FieldName::ContentLength, value); + } + } + source +} + +struct BorrowedSource<'a> { + user_agent: [FieldValueRef<'a>; 1], +} + +impl FieldSource for BorrowedSource<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + if name == &FieldName::UserAgent { + FieldLines::from_borrowed(name, &self.user_agent) + } else { + None + } + } +} + +#[test] +fn deferred_and_fluent_writes_work_with_downstream_sinks() { + let mut compatibility = TestMap::new(); + compatibility + .set_content_length(42) + .expect("in-memory insertion succeeds") + .set_content_type(ContentType::json()) + .expect("in-memory insertion succeeds"); + let dynamic_sink: &mut dyn FieldSink = &mut compatibility; + dynamic_sink.remove_values(&FieldName::ContentLength); + + ContentType::json() + .insert_into(&mut compatibility) + .expect("in-memory insertion succeeds"); + compatibility.set_content_length(42).expect("in-memory insertion succeeds"); + + let source = BorrowedSource { + user_agent: [FieldValueRef::new(b"client/1")], + }; + let view = UserAgent::view(&source).expect("valid user agent").expect("user agent is present"); + view.insert_into(&mut compatibility).expect("borrowed insertion succeeds"); +} + +#[derive(Clone, Debug, Eq, PartialEq)] +struct RequestId(FieldValue); + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +struct RequestIdView<'a>(FieldValueRef<'a>); + +impl SingleValueField for RequestId { + type View<'a> = RequestIdView<'a>; + type Owned = Self; + + fn name() -> &'static FieldName { + &REQUEST_ID + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + if value.as_bytes().is_empty() { + Err(DecodeError::new(&REQUEST_ID, DecodeErrorKind::InvalidToken)) + } else { + Ok(RequestIdView(value)) + } + } + + fn decode_owned(value: FieldValue) -> Result { + if value.as_bytes().is_empty() { + Err(DecodeError::new(&REQUEST_ID, DecodeErrorKind::InvalidToken)) + } else { + Ok(Self(value)) + } + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.0 + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.0 + } +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +struct EmptyHeader; + +impl Field for EmptyHeader { + type View<'a> = Self; + type Owned = Self; + + fn name() -> &'static FieldName { + &EMPTY + } + + fn view_with(source: &S, _mode: DecodeMode) -> Result>, DecodeError> + where + S: FieldSource + ?Sized, + { + Ok(source.contains(Self::name()).then_some(Self)) + } + + fn owned_with(source: &S, mode: DecodeMode) -> Result, DecodeError> + where + S: FieldSource + ?Sized, + { + Self::view_with(source, mode) + } + + fn insert(sink: &mut S, _value: Self::Owned) -> Result<(), InsertError> + where + S: FieldSink + ?Sized, + { + sink.set_values(Self::name(), EncodedValues::new()) + } +} + +#[test] +fn constructed_and_parsed_custom_names_are_one_value() { + let constructed = FieldName::from_static("x-trace-id"); + let parsed = FieldName::try_from_bytes(b"X-Trace-Id").expect("valid name"); + + assert_eq!(constructed, parsed); + assert_eq!(parsed.as_bytes(), b"x-trace-id"); + assert_eq!(format!("{parsed}"), "x-trace-id"); + assert_eq!(parsed.index(), None); +} + +#[test] +fn every_known_name_indexes_its_own_slot() { + for (index, name) in FieldName::ALL_KNOWN.iter().enumerate() { + assert_eq!(name.index(), Some(index)); + assert_eq!(FieldName::try_from_bytes(name.as_str().as_bytes()).expect("a valid name"), *name); + } +} + +#[test] +fn delimited_items_cross_lines_and_preserve_quoted_commas() { + let stored = [FieldValue::from_static("a, \"b,c\""), FieldValue::from_static("d")]; + let values = FieldLines::from_slice(&REQUEST_ID, &stored).expect("stored lines are present"); + let items: Result>, _> = values.comma_items().map(|item| item.map(<[u8]>::to_vec)).collect(); + assert_eq!(items, Ok(vec![b"a".to_vec(), b"\"b,c\"".to_vec(), b"d".to_vec()])); +} + +#[test] +fn values_can_be_reiterated_without_collecting() { + let stored = [FieldValue::from_static("a"), FieldValue::from_static("b")]; + let values = FieldLines::from_slice(&REQUEST_ID, &stored).expect("stored lines are present"); + assert_eq!(values.repeated().count(), 2); + assert_eq!(values.repeated().count(), 2); +} + +#[test] +fn comma_items_ignore_empty_members() { + let stored = [FieldValue::from_static(",, a,"), FieldValue::from_static(" , b ,,")]; + let values = FieldLines::from_slice(&REQUEST_ID, &stored).expect("stored lines are present"); + let items: Result>, _> = values.comma_items().map(|item| item.map(<[u8]>::to_vec)).collect(); + assert_eq!(items, Ok(vec![b"a".to_vec(), b"b".to_vec()])); +} + +#[test] +fn constructors_distinguish_absence_from_an_empty_field_line() { + assert!(FieldLines::from_slice(&REQUEST_ID, &[]).is_none()); + let single = FieldLines::single(&REQUEST_ID, b"a"); + assert_eq!(single.len(), 1); + assert_eq!(single.repeated().count(), 1); + + let empty_line = FieldLines::single(&REQUEST_ID, b""); + assert_eq!(empty_line.len(), 1); + assert_eq!(empty_line.repeated().count(), 1); +} + +#[test] +fn debug_does_not_expose_field_values() { + let stored = [FieldValue::from_static("secret")]; + let values = FieldLines::from_slice(&REQUEST_ID, &stored).expect("stored lines are present"); + let debug = format!("{values:?}"); + assert!(!debug.contains("secret")); + assert!(debug.contains("line_count")); + + let debug = format!("{:?}", values.comma_items()); + assert!(!debug.contains("secret")); +} + +#[test] +fn custom_single_value_header_uses_all_map_operations() { + let mut map = TestMap::new(); + RequestId::insert(&mut map, RequestId(FieldValue::from_static("abc-123"))).expect("header map has capacity"); + + assert!(map.contains(&REQUEST_ID)); + assert_eq!( + RequestId::view(&map).expect("valid request ID").expect("request ID present").0, + "abc-123" + ); + assert_eq!(RequestId::owned(&map), Ok(Some(RequestId(FieldValue::from_static("abc-123"))))); + + RequestId::remove(&mut map); + assert!(!map.contains(&REQUEST_ID)); +} + +#[test] +fn typed_custom_name_finds_an_owned_runtime_name() { + let mut map = TestMap::new(); + map.insert( + FieldName::try_from_bytes(b"X-Request-Id").expect("valid runtime name"), + FieldValue::from_static("from-wire"), + ); + + assert_eq!( + RequestId::view(&map).expect("valid request ID").expect("request ID present").0, + "from-wire" + ); + assert_eq!(map.names().next().expect("stored runtime name").as_str(), "x-request-id"); +} + +#[test] +fn source_presence_distinguishes_an_empty_field_line_from_absence() { + let mut map = TestMap::new(); + assert!(map.lines(&FieldName::UserAgent).is_none()); + + map.insert(FieldName::UserAgent, FieldValue::from_static("")); + let values = map.lines(&FieldName::UserAgent).expect("zero-length field line is present"); + assert_eq!(values.len(), 1); + assert_eq!(values.exactly_one().expect("one field line").as_bytes(), b""); +} + +#[test] +fn custom_source_decodes_references_into_external_storage() { + let wire = b"client/1"; + let source = BorrowedSource { + user_agent: [FieldValueRef::new(wire)], + }; + + assert_eq!( + UserAgent::view(&source) + .expect("valid user agent") + .expect("user agent present") + .as_bytes(), + b"client/1" + ); +} + +#[test] +fn duplicate_single_value_is_an_error() { + let mut map = TestMap::new(); + map.append((*REQUEST_ID).clone(), FieldValue::from_static("first")); + map.append((*REQUEST_ID).clone(), FieldValue::from_static("second")); + + let error = RequestId::view(&map).expect_err("duplicates must fail"); + assert_eq!(error.kind(), DecodeErrorKind::UnexpectedMultipleValues); +} + +#[test] +fn empty_encoding_removes_a_stale_value() { + let mut map = TestMap::new(); + map.insert(EMPTY.clone(), FieldValue::from_static("stale")); + EmptyHeader::insert(&mut map, EmptyHeader).expect("header map has capacity"); + assert!(!map.contains_name(&EMPTY)); +} + +#[test] +fn insert_replaces_existing_single_and_repeated_values() { + let mut map = TestMap::new(); + map.insert((*REQUEST_ID).clone(), FieldValue::from_static("old")); + RequestId::insert(&mut map, RequestId(FieldValue::from_static("new"))).expect("header map has capacity"); + assert_eq!(map.get_all(&REQUEST_ID), [FieldValue::from_static("new")]); + + map.append(FieldName::SetCookie, FieldValue::from_static("stale=first")); + map.append(FieldName::SetCookie, FieldValue::from_static("stale=second")); + let mut replacement = SetCookieOwned::new(); + replacement.push_str("fresh=first").expect("valid Set-Cookie value"); + replacement.push_str("fresh=second").expect("valid Set-Cookie value"); + SetCookie::insert(&mut map, replacement).expect("header map has capacity"); + assert_eq!( + map.get_all(&FieldName::SetCookie) + .iter() + .map(FieldValue::as_bytes) + .collect::>(), + [b"fresh=first".as_slice(), b"fresh=second".as_slice()] + ); +} + +#[test] +fn built_in_headers_expose_their_well_known_name() { + fn assert_known(expected: &'static FieldName) { + assert_eq!(H::name(), expected); + assert_eq!(H::name().as_str(), expected.as_str()); + assert!( + H::name().index().expect("a built-in header is well known") < FieldName::COUNT, + "a well-known name addresses a slot directly" + ); + } + + assert_known::(&FieldName::AcceptRanges); + assert_known::>(&FieldName::Authorization); + assert_known::(&FieldName::Etag); + assert_known::(&FieldName::IfMatch); + assert_known::(&FieldName::IfNoneMatch); + assert_known::(&FieldName::IfRange); + assert_known::(&FieldName::ReferrerPolicy); + assert_known::(&FieldName::SecWebSocketVersion); + assert_known::(&FieldName::SetCookie); + + // A downstream header has no dense index, so it is resolved by name. + assert_eq!(::name().index(), None); + assert_eq!(::name().as_str(), "x-request-id"); +} + +#[test] +fn public_header_constructors_cover_success_and_rejection_paths() { + let basic = AuthorizationOwned::::basic(b"user", b"password").expect("valid credentials"); + assert_eq!(basic.encoded_credentials().expect("stored value is valid"), b"dXNlcjpwYXNzd29yZA=="); + + ETagOwned::try_from_wire("\"revision\"").expect("quoted wire tag is valid"); + assert_eq!( + ETagOwned::try_from_wire("revision").expect_err("wire tags require quotes").kind(), + DecodeErrorKind::InvalidSyntax + ); + + let tag = ETagOwned::strong("revision").expect("valid tag"); + IfMatchOwned::from_tags([tag.clone()]).expect("nonempty tag list is valid"); + assert_eq!( + IfMatchOwned::from_tags(Vec::::new()) + .expect_err("an empty tag list is invalid") + .kind(), + DecodeErrorKind::MissingValue + ); + IfNoneMatchOwned::from_tags([tag.clone()]).expect("nonempty tag list is valid"); + assert_eq!( + IfNoneMatchOwned::from_tags(Vec::::new()) + .expect_err("an empty tag list is invalid") + .kind(), + DecodeErrorKind::MissingValue + ); + + IfRangeOwned::entity_tag(tag).expect("strong entity tag is valid for If-Range"); + assert_eq!( + IfRangeOwned::entity_tag(ETagOwned::weak("revision").expect("valid weak tag")) + .expect_err("If-Range requires a strong tag") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + + AcceptRangesOwned::from_units(["bytes", "items"]).expect("token range units are valid"); + assert_eq!( + AcceptRangesOwned::from_units(["none", "bytes"]) + .expect_err("none cannot be combined with another unit") + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + AcceptRangesOwned::from_units(Vec::::new()) + .expect_err("the unit list cannot be empty") + .kind(), + DecodeErrorKind::MissingValue + ); + + let policies = ReferrerPolicyOwned::new(ReferrerPolicyValue::NoReferrer).with_fallback(ReferrerPolicyValue::StrictOrigin); + assert_eq!( + policies.preferred().expect("fallback list is valid"), + ReferrerPolicyValue::StrictOrigin + ); + + let versions = SecWebSocketVersionOwned::new(13).with_version(8); + assert_eq!(versions.versions().collect::>(), [8, 13]); +} + +#[test] +fn credentials_release_excess_capacity() { + let mut map = TestMap::new(); + Authorization::::insert( + &mut map, + AuthorizationOwned::::basic([b'a'; 64], b"").expect("valid credentials"), + ) + .expect("header map has capacity"); + let authorization = Authorization::::view(&map) + .expect("valid credentials") + .expect("authorization present"); + let mut credentials = BasicCredentials::with_retain_limit(8); + authorization.extract(&mut credentials).expect("valid credentials"); + assert!(credentials.capacity() >= 64); + credentials.clear(); + assert!(credentials.capacity() <= 8); + assert!(credentials.username().is_empty()); + assert!(credentials.password().is_empty()); +} + +#[test] +fn credentials_can_release_or_retain_their_allocation() { + let mut map = TestMap::new(); + Authorization::::insert( + &mut map, + AuthorizationOwned::::basic(b"user", b"password").expect("valid credentials"), + ) + .expect("header map has capacity"); + let authorization = Authorization::::view(&map) + .expect("valid credentials") + .expect("authorization present"); + + let mut released = BasicCredentials::with_retain_limit(0); + authorization.extract(&mut released).expect("valid credentials"); + released.clear(); + assert_eq!(released.capacity(), 0); + + let mut retained = BasicCredentials::with_retain_limit(usize::MAX); + authorization.extract(&mut retained).expect("valid credentials"); + let capacity = retained.capacity(); + retained.clear(); + assert_eq!(retained.capacity(), capacity); +} + +/// Covers the one storage limit only the adapted `http::HeaderMap` imposes. +#[cfg(feature = "http")] +#[test] +fn insert_reports_full_map_without_replacing_values() { + use http::{HeaderMap, HeaderValue}; + + let old = HeaderValue::from_static("old"); + #[cfg(not(miri))] + let mut map = HeaderMap::new(); + #[cfg(miri)] + let mut map = HeaderMap::with_capacity(miri_http_map::CAPACITY); + map.insert(REQUEST_ID.as_str(), old.clone()); + for index in 0_u64.. { + #[cfg(not(miri))] + let name = HeaderName::try_from(format!("x-fill-{index}")).expect("generated header name is valid"); + #[cfg(miri)] + let name = HeaderName::from_static(miri_http_map::name(usize::try_from(index).unwrap())); + if map.try_insert(name, HeaderValue::from_static("fill")).is_err() { + break; + } + } + #[cfg(miri)] + assert_eq!(map.len(), miri_http_map::CAPACITY); + + let result = RequestId::insert(&mut map, RequestId(FieldValue::from_static("new"))); + assert!(result.is_err()); + assert_eq!(map.get(REQUEST_ID.as_str()), Some(&old)); + + let append_result = map.append_encoded(&REQUEST_ID, FieldValue::from_static("additional")); + assert!(append_result.is_err()); + assert_eq!(map.get(REQUEST_ID.as_str()), Some(&old)); +} + +#[test] +fn content_length_formats_parses_and_inserts() { + let value = ContentLengthOwned::new(42); + assert_eq!(value.get(), 42); + assert_eq!(value.to_string(), "42"); + assert_eq!(" \t42 ".parse::(), Ok(value)); + assert_eq!( + "42x".parse::().expect_err("trailing data must fail").kind(), + DecodeErrorKind::InvalidNumber + ); + + let mut source = TestMap::new(); + ContentLength::insert(&mut source, value).expect("in-memory insertion succeeds"); + assert_eq!(ContentLength::view(&source), Ok(Some(value))); + assert_eq!(ContentLength::owned(&source), Ok(Some(value))); + source.remove_values(&FieldName::ContentLength); + assert_eq!(ContentLength::view(&source), Ok(None)); +} + +#[test] +fn content_length_decodes_single_repeated_and_missing_lines() { + assert_eq!( + ContentLength::view(&content_length_source(&[" 7\t"])), + Ok(Some(ContentLengthOwned::new(7))) + ); + assert_eq!( + ContentLength::view(&content_length_source(&["7, 7", " 7 ", "7,7", "7"])), + Ok(Some(ContentLengthOwned::new(7))) + ); + + for values in [&["7,8"][..], &["7", "8"], &[","]] { + ContentLength::view(&content_length_source(values)).expect_err("inconsistent or empty content lengths are invalid"); + } + + assert_eq!( + ContentLength::view(&content_length_source(&[""])) + .expect_err("empty value is invalid") + .kind(), + DecodeErrorKind::InvalidNumber + ); +} + +#[cfg(feature = "http")] +#[test] +fn publicly_constructible_names_convert_without_panicking() { + let names = FieldName::ALL_KNOWN + .iter() + .cloned() + .chain([ + FieldName::from_static("x-trace-id"), + FieldName::try_from_bytes(b"X-Vendor-Field").expect("valid name"), + FieldName::try_from_bytes(vec![b'a'; (1 << 16) - 1]).expect("valid name"), + FieldName::from(&HeaderName::from_lowercase(b"x\"y").expect("http admits `\"`")), + ]) + .collect::>(); + + for name in names { + let converted = name.try_to_http_header_name().expect("every name converts"); + assert_eq!(converted.as_str(), name.as_str()); + assert_eq!(HeaderName::from(&name), converted); + let round_trip = FieldName::from(&converted); + assert_eq!(round_trip, name); + assert_eq!(round_trip.index(), name.index()); + } +} + +#[cfg(feature = "http")] +#[test] +fn http_names_holding_a_quote_convert_without_panicking() { + let quoted = HeaderName::from_lowercase(b"x\"y").expect("http admits `\"`"); + let converted = FieldName::from("ed); + assert_eq!(converted.as_str(), "x\"y"); + assert_eq!(converted.index(), None); + assert_eq!(FieldName::from(quoted.clone()), converted); + + assert_eq!(HeaderName::from(&converted), quoted); + assert_eq!(converted.try_to_http_header_name().expect("round trips"), quoted); + + FieldName::try_from_bytes(b"x\"y").expect_err("this crate's own parser rejects `\"`"); +} + +#[test] +fn field_names_respect_the_http_length_limit() { + const MAX: usize = (1 << 16) - 1; + static OVER_LONG: [u8; MAX + 1] = [b'a'; MAX + 1]; + + FieldName::try_from_bytes(vec![b'a'; MAX + 1]).expect_err("an over-long name is rejected"); + let over_long = str::from_utf8(&OVER_LONG).unwrap(); + panic::catch_unwind(|| FieldName::from_static(over_long)).unwrap_err(); + + let longest = FieldName::try_from_bytes(vec![b'a'; MAX]).expect("the longest name is valid"); + assert_eq!(longest.as_str().len(), MAX); + + #[cfg(feature = "http")] + { + let converted = HeaderName::from(&longest); + assert_eq!(converted.as_str().len(), MAX); + assert_eq!(FieldName::from(&converted), longest); + } +} + +#[test] +fn field_name_recognition_only_uses_same_length_candidates() { + for name in FieldName::ALL_KNOWN { + let mut shorter = name.as_str().as_bytes().to_vec(); + let _removed = shorter.pop(); + let mut longer = name.as_str().as_bytes().to_vec(); + longer.push(b'x'); + + for altered in [shorter, longer] { + let parsed = FieldName::try_from_bytes(&altered).expect("still a token"); + assert_ne!(parsed, *name); + } + } + + assert_eq!(FieldName::try_from_bytes(b"a").expect("valid").index(), None); + let over_long = vec![b'a'; 64]; + assert_eq!(FieldName::try_from_bytes(&over_long).expect("valid").index(), None); +} diff --git a/crates/http_headers/tests/cors_iteration.rs b/crates/http_headers/tests/cors_iteration.rs new file mode 100644 index 000000000..91b716f5e --- /dev/null +++ b/crates/http_headers/tests/cors_iteration.rs @@ -0,0 +1,101 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Borrowed iteration over owned CORS collections preserves their token lists. + +#![cfg(feature = "headers-cors")] + +use std::fmt::{self, Write as _}; + +use http_headers::FieldValue; +use http_headers::headers::{ + AccessControlAllowHeadersOwned, AccessControlAllowMethodsOwned, AccessControlExposeHeadersOwned, AccessControlRequestHeadersOwned, +}; + +struct RejectWriter; + +impl fmt::Write for RejectWriter { + fn write_str(&mut self, _text: &str) -> fmt::Result { + Err(fmt::Error) + } +} + +#[test] +fn borrowed_owned_collection_iteration_keeps_order_case_duplicates_and_wildcards() { + macro_rules! check { + ($owned:ty, $first:literal, $last:literal, $expected:expr, $debug_name:literal) => {{ + let value = <$owned>::from_field_values(vec![ + FieldValue::from_static($first), + FieldValue::from_static(""), + FieldValue::from_static($last), + ]) + .unwrap(); + let expected = $expected; + assert_eq!(value.iter().map(|item| item.as_str()).collect::>(), expected); + assert_eq!((&value).into_iter().map(|item| item.as_str()).collect::>(), expected); + let mut seen = Vec::new(); + for item in &value { + seen.push(item.as_str()); + } + assert_eq!(seen, expected); + assert_eq!(value.field_values().count(), 3); + assert_eq!(value.len(), expected.len()); + let mut iter = (&value).into_iter(); + for expected in expected { + assert_eq!(iter.next().unwrap().as_str(), expected); + } + assert!(iter.next().is_none()); + assert!(iter.next().is_none()); + assert_eq!(format!("{iter:?}"), concat!($debug_name, " { .. }")); + assert_eq!(write!(&mut RejectWriter, "{:?}", value.iter()), Err(fmt::Error)); + }}; + } + check!( + AccessControlAllowHeadersOwned, + " \t, X-Trace,, content-type ,", + "x-trace,*,", + ["X-Trace", "content-type", "x-trace", "*"], + "CorsHeaderNames" + ); + check!( + AccessControlExposeHeadersOwned, + " \t, X-Trace,, content-type ,", + "x-trace,*,", + ["X-Trace", "content-type", "x-trace", "*"], + "CorsHeaderNames" + ); + check!( + AccessControlRequestHeadersOwned, + " \t, X-Trace,, content-type ,", + "x-trace,*,", + ["X-Trace", "content-type", "x-trace", "*"], + "CorsHeaderNames" + ); + check!( + AccessControlAllowMethodsOwned, + " \t, GET,, get ,", + "X-CUSTOM,*,GET", + ["GET", "get", "X-CUSTOM", "*", "GET"], + "CorsMethods" + ); +} + +#[test] +fn permitted_empty_collections_are_empty_borrowed_iterators() { + macro_rules! check { + ($owned:ty) => {{ + for value in [ + <$owned>::empty(), + <$owned>::from_field_values(vec![FieldValue::from_static(" ,,\t"), FieldValue::from_static("")]).unwrap(), + ] { + let mut iter = (&value).into_iter(); + assert!(iter.next().is_none()); + assert!(iter.next().is_none()); + assert_eq!(value.iter().count(), 0); + } + }}; + } + check!(AccessControlAllowHeadersOwned); + check!(AccessControlExposeHeadersOwned); + check!(AccessControlAllowMethodsOwned); +} diff --git a/crates/http_headers/tests/cors_tokens.rs b/crates/http_headers/tests/cors_tokens.rs new file mode 100644 index 000000000..8d15c3dce --- /dev/null +++ b/crates/http_headers/tests/cors_tokens.rs @@ -0,0 +1,144 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Shared field-name and method tokens remain available with only CORS enabled. + +#![cfg(feature = "headers-cors")] + +use std::collections::HashSet; +use std::collections::hash_map::DefaultHasher; +use std::hash::{Hash, Hasher}; + +use http_headers::headers::{ + AccessControlAllowHeaders, AccessControlAllowMethods, AccessControlExposeHeaders, AccessControlRequestHeaders, + AccessControlRequestMethod, FieldNameView, InvalidMethod, MethodView, +}; +#[cfg(feature = "http")] +use http_headers::headers::{ + AccessControlAllowHeadersOwned, AccessControlAllowMethodsOwned, AccessControlExposeHeadersOwned, AccessControlRequestHeadersOwned, + AccessControlRequestMethodOwned, +}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{FieldName, FieldValue}; + +struct Source { + name: &'static FieldName, + values: Vec, +} + +impl FieldSource for Source { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::from_slice(name, &self.values)).flatten() + } +} + +fn hash(value: &impl Hash) -> u64 { + let mut hasher = DefaultHasher::new(); + value.hash(&mut hasher); + hasher.finish() +} + +#[test] +fn cors_names_share_case_insensitive_identity_without_changing_wire_spelling() { + macro_rules! check { + ($header:ty, $name:expr) => {{ + let source = Source { + name: $name, + values: vec![FieldValue::from_static("X-Trace, content-type"), FieldValue::from_static("x-trace")], + }; + let owned = <$header>::owned(&source).unwrap().unwrap(); + let borrowed = <$header>::view(&source).unwrap().unwrap(); + let owned_names: Vec> = owned.iter().collect(); + let borrowed_names: Vec> = borrowed.iter().collect(); + assert_eq!(owned_names, borrowed_names); + assert_eq!( + borrowed_names.iter().map(|name| name.as_str()).collect::>(), + ["X-Trace", "content-type", "x-trace"] + ); + assert_eq!( + owned_names.iter().map(|name| name.as_str()).collect::>(), + ["X-Trace", "content-type", "x-trace"] + ); + assert_eq!(borrowed_names[0], borrowed_names[2]); + assert_eq!(borrowed_names[0], FieldNameView::new("x-trace").unwrap()); + assert_eq!(hash(&borrowed_names[0]), hash(&borrowed_names[2])); + assert_eq!(borrowed_names.iter().copied().collect::>().len(), 2); + assert_eq!(borrowed_names[0].as_bytes(), b"X-Trace"); + assert_eq!(borrowed_names[0].try_to_field_name().unwrap().as_str(), "x-trace"); + assert_eq!(borrowed_names[1].try_to_field_name().unwrap(), FieldName::ContentType); + assert_eq!( + owned.field_values().map(|value| value.as_bytes()).collect::>(), + [b"X-Trace, content-type".as_slice(), b"x-trace"] + ); + }}; + } + check!(AccessControlAllowHeaders, &FieldName::AccessControlAllowHeaders); + check!(AccessControlExposeHeaders, &FieldName::AccessControlExposeHeaders); + check!(AccessControlRequestHeaders, &FieldName::AccessControlRequestHeaders); +} + +#[test] +fn cors_methods_share_case_sensitive_identity_and_extension_tokens() { + let source = Source { + name: &FieldName::AccessControlAllowMethods, + values: vec![FieldValue::from_static("GET, get, X-CUSTOM, *"), FieldValue::from_static("GET")], + }; + let borrowed = AccessControlAllowMethods::view(&source).unwrap().unwrap(); + let owned = AccessControlAllowMethods::owned(&source).unwrap().unwrap(); + let methods: Vec> = borrowed.iter().collect(); + assert_eq!(methods, owned.iter().collect::>()); + assert_eq!( + methods.iter().map(|method| method.as_str()).collect::>(), + ["GET", "get", "X-CUSTOM", "*", "GET"] + ); + assert_eq!( + owned.iter().map(MethodView::as_str).collect::>(), + ["GET", "get", "X-CUSTOM", "*", "GET"] + ); + assert_eq!(methods[0], MethodView::GET); + assert_ne!(methods[0], methods[1]); + assert_eq!(methods[2], MethodView::new("X-CUSTOM").unwrap()); + assert_ne!(methods[2], MethodView::new("x-custom").unwrap()); + assert_eq!(methods.iter().copied().collect::>().len(), 4); + assert_eq!(hash(&methods[0]), hash(&methods[4])); + assert_eq!(MethodView::new("bad method").unwrap_err(), InvalidMethod); + + for wire in ["GET", "get", "X-CUSTOM", "*"] { + let source = Source { + name: &FieldName::AccessControlRequestMethod, + values: vec![FieldValue::from_str(wire).unwrap()], + }; + let borrowed = AccessControlRequestMethod::view(&source).unwrap().unwrap(); + let owned = AccessControlRequestMethod::owned(&source).unwrap().unwrap(); + let method = borrowed.method(); + assert_eq!(method, owned.method().unwrap()); + assert_eq!(method, MethodView::new(wire).unwrap()); + assert_eq!(method.as_bytes(), wire.as_bytes()); + } +} + +#[cfg(feature = "http")] +#[test] +fn shared_cors_tokens_convert_to_http_names_and_methods() { + let allowed = AccessControlAllowHeadersOwned::from_header_names(["Content-Type", "X-Trace"]).unwrap(); + let exposed = AccessControlExposeHeadersOwned::from_header_names(["Content-Type", "X-Trace"]).unwrap(); + let requested = AccessControlRequestHeadersOwned::from_header_names(["Content-Type", "X-Trace"]).unwrap(); + for names in [ + allowed.iter().collect::>(), + exposed.iter().collect(), + requested.iter().collect(), + ] { + assert_eq!(names[0].try_to_http_header_name().unwrap(), http::header::CONTENT_TYPE); + assert_eq!(names[1].try_to_http_header_name().unwrap().as_str(), "x-trace"); + assert_eq!(names[1].as_str(), "X-Trace"); + } + let allowed = AccessControlAllowMethodsOwned::from_methods(["GET", "get", "X-CUSTOM", "*"]).unwrap(); + for method in &allowed { + assert_eq!(method.try_to_method().unwrap().as_str(), method.as_str()); + let requested = AccessControlRequestMethodOwned::from_method(method.as_str()).unwrap(); + assert_eq!( + requested.method().unwrap().try_to_method().unwrap(), + method.try_to_method().unwrap() + ); + } +} diff --git a/crates/http_headers/tests/decode_mode_parity.rs b/crates/http_headers/tests/decode_mode_parity.rs new file mode 100644 index 000000000..3c3360e58 --- /dev/null +++ b/crates/http_headers/tests/decode_mode_parity.rs @@ -0,0 +1,387 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Strict/relaxed parity checks across every documented interoperability class. + +#![cfg(feature = "headers-all")] +#![expect( + clippy::too_many_lines, + reason = "one exhaustive matrix keeps every relaxed-decoding class together" +)] +#![expect(clippy::unwrap_used, reason = "test failures provide sufficient context")] + +use http_headers::headers::{ + Accept, AcceptEncoding, AcceptLanguage, ContentRange, ContentType, ETag, Host, IfRange, IfRangeValueView, LastModified, Location, Range, +}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeMode, Field, FieldName}; + +#[derive(Debug, Eq, PartialEq)] +struct Projection { + wire: Vec>, + semantics: String, +} + +struct Source { + name: &'static FieldName, + value: &'static [u8], +} + +impl FieldSource for Source { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::single(name, self.value)) + } +} + +type Project = fn(&Source) -> Projection; +type Reject = fn(&Source) -> bool; + +struct Case { + label: &'static str, + name: &'static FieldName, + relaxed: &'static [u8], + invalid_neighbor: &'static [u8], + strict_rejects: Reject, + relaxed_rejects: Reject, + borrowed: Project, + owned: Project, +} + +fn rejects_strict(source: &Source) -> bool { + H::view(source).is_err() && H::owned(source).is_err() +} + +fn rejects_relaxed(source: &Source) -> bool { + H::view_with(source, DecodeMode::Relaxed).is_err() && H::owned_with(source, DecodeMode::Relaxed).is_err() +} + +macro_rules! list_projection { + ($borrowed:ident, $owned:ident, $header:ty) => { + fn $borrowed(source: &Source) -> Projection { + let value = <$header>::view_with(source, DecodeMode::Relaxed) + .expect("relaxed borrowed decode") + .expect("field is present"); + Projection { + wire: value.values().map(|line| line.as_bytes().to_vec()).collect(), + semantics: format!("{:?}", value.items().map(<[u8]>::to_vec).collect::>()), + } + } + + fn $owned(source: &Source) -> Projection { + let value = <$header>::owned_with(source, DecodeMode::Relaxed) + .expect("relaxed owned decode") + .expect("field is present"); + Projection { + wire: value.values().map(|line| line.as_bytes().to_vec()).collect(), + semantics: format!("{:?}", value.items().map(<[u8]>::to_vec).collect::>()), + } + } + }; +} + +list_projection!(accept_view, accept_owned, Accept); +list_projection!(accept_encoding_view, accept_encoding_owned, AcceptEncoding); +list_projection!(accept_language_view, accept_language_owned, AcceptLanguage); + +fn etag_view(source: &Source) -> Projection { + let value = ETag::view_with(source, DecodeMode::Relaxed) + .expect("relaxed borrowed decode") + .expect("field is present"); + Projection { + wire: vec![source.value.to_vec()], + semantics: format!("{}:{:?}", value.is_weak(), value.opaque_tag()), + } +} + +fn etag_owned(source: &Source) -> Projection { + let value = ETag::owned_with(source, DecodeMode::Relaxed) + .expect("relaxed owned decode") + .expect("field is present"); + let semantics = format!("{}:{:?}", value.is_weak(), value.opaque_tag().expect("validated metadata")); + Projection { + wire: vec![value.into_field_value().as_bytes().to_vec()], + semantics, + } +} + +fn content_type_view(source: &Source) -> Projection { + let value = ContentType::view_with(source, DecodeMode::Relaxed) + .expect("relaxed borrowed decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: format!("{}:{}", value.type_().unwrap(), value.subtype().unwrap()), + } +} + +fn content_type_owned(source: &Source) -> Projection { + let value = ContentType::owned_with(source, DecodeMode::Relaxed) + .expect("relaxed owned decode") + .expect("field is present"); + let semantics = format!("{}:{}", value.type_().unwrap(), value.subtype().unwrap()); + Projection { + wire: vec![value.into_field_value().as_bytes().to_vec()], + semantics, + } +} + +fn range_view(source: &Source) -> Projection { + let value = Range::view_with(source, DecodeMode::Relaxed) + .expect("relaxed borrowed decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: format!("{:?}", value.byte_ranges().expect("byte ranges").collect::>()), + } +} + +fn range_owned(source: &Source) -> Projection { + let value = Range::owned_with(source, DecodeMode::Relaxed) + .expect("relaxed owned decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: format!("{:?}", value.byte_ranges().expect("byte ranges").collect::>()), + } +} + +fn content_range_view(source: &Source) -> Projection { + let value = ContentRange::view_with(source, DecodeMode::Relaxed) + .expect("relaxed borrowed decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: format!("{:?}", value.byte_range()), + } +} + +fn content_range_owned(source: &Source) -> Projection { + let value = ContentRange::owned_with(source, DecodeMode::Relaxed) + .expect("relaxed owned decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: format!("{:?}", value.byte_range()), + } +} + +fn last_modified_view(source: &Source) -> Projection { + let value = LastModified::view_with(source, DecodeMode::Relaxed) + .expect("relaxed borrowed decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: format!("{:?}", value.date()), + } +} + +fn last_modified_owned(source: &Source) -> Projection { + let value = LastModified::owned_with(source, DecodeMode::Relaxed) + .expect("relaxed owned decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: format!("{:?}", value.date()), + } +} + +fn if_range_view(source: &Source) -> Projection { + let value = IfRange::view_with(source, DecodeMode::Relaxed) + .expect("relaxed borrowed decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: if_range_semantics(value.value()), + } +} + +fn if_range_owned(source: &Source) -> Projection { + let value = IfRange::owned_with(source, DecodeMode::Relaxed) + .expect("relaxed owned decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: if_range_semantics(value.value().expect("validated metadata")), + } +} + +fn if_range_semantics(value: IfRangeValueView<'_>) -> String { + match value { + IfRangeValueView::EntityTag(tag) => format!("tag:{:?}", tag.opaque_tag()), + IfRangeValueView::Date(date) => format!("date:{date:?}"), + } +} + +fn host_view(source: &Source) -> Projection { + let value = Host::view_with(source, DecodeMode::Relaxed) + .expect("relaxed borrowed decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: format!("{}:{:?}", value.host(), value.port()), + } +} + +fn host_owned(source: &Source) -> Projection { + let value = Host::owned_with(source, DecodeMode::Relaxed) + .expect("relaxed owned decode") + .expect("field is present"); + Projection { + wire: vec![value.as_field_value().as_bytes().to_vec()], + semantics: format!("{}:{:?}", value.host().unwrap(), value.port().unwrap()), + } +} + +fn location_view(source: &Source) -> Projection { + let value = Location::view_with(source, DecodeMode::Relaxed) + .expect("relaxed borrowed decode") + .expect("field is present"); + Projection { + wire: vec![value.as_bytes().to_vec()], + semantics: value.as_str().unwrap().to_owned(), + } +} + +fn location_owned(source: &Source) -> Projection { + let value = Location::owned_with(source, DecodeMode::Relaxed) + .expect("relaxed owned decode") + .expect("field is present"); + Projection { + wire: vec![value.as_bytes().to_vec()], + semantics: value.as_str().unwrap().to_owned(), + } +} + +#[test] +fn strict_and_relaxed_borrowed_owned_semantics_stay_in_parity() { + let cases = [ + Case { + label: "Accept qvalue precision", + name: &FieldName::Accept, + relaxed: b"text/html; q = .12345", + invalid_neighbor: b"text/html;q=-.5", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: accept_view, + owned: accept_owned, + }, + Case { + label: "Accept-Encoding qvalue whitespace", + name: &FieldName::AcceptEncoding, + relaxed: b"gzip; q=.5", + invalid_neighbor: b"gzip;q=1.1", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: accept_encoding_view, + owned: accept_encoding_owned, + }, + Case { + label: "Accept-Language qvalue precision", + name: &FieldName::AcceptLanguage, + relaxed: b"en-US; Q = 1.0000", + invalid_neighbor: b"en_US;q=.5", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: accept_language_view, + owned: accept_language_owned, + }, + Case { + label: "lowercase weak ETag", + name: &FieldName::Etag, + relaxed: b"w/\"revision\"", + invalid_neighbor: b"w/\"bad tag\"", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: etag_view, + owned: etag_owned, + }, + Case { + label: "Content-Type delimiter whitespace", + name: &FieldName::ContentType, + relaxed: b"text / html", + invalid_neighbor: b"text // html", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: content_type_view, + owned: content_type_owned, + }, + Case { + label: "Range delimiter whitespace", + name: &FieldName::Range, + relaxed: b"bytes = 0 - 499", + invalid_neighbor: b"bytes = +1 - 2", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: range_view, + owned: range_owned, + }, + Case { + label: "Content-Range delimiter whitespace", + name: &FieldName::ContentRange, + relaxed: b"bytes 0 - 499 / 1234", + invalid_neighbor: b"bytes 0 - 1 / 2", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: content_range_view, + owned: content_range_owned, + }, + Case { + label: "relaxed IMF-style date", + name: &FieldName::LastModified, + relaxed: b" Sun, 6 Nov 1994 8:49:37 UTC ", + invalid_neighbor: b"Sun,\t6 Nov 1994 8:49:37 UTC", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: last_modified_view, + owned: last_modified_owned, + }, + Case { + label: "relaxed If-Range date", + name: &FieldName::IfRange, + relaxed: b" Tue, 8 Nov 1994 8:49:37 UTC ", + invalid_neighbor: b"Tue,\t8 Nov 1994 8:49:37 UTC", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: if_range_view, + owned: if_range_owned, + }, + Case { + label: "internationalized Host", + name: &FieldName::Host, + relaxed: "münich.example:443".as_bytes(), + invalid_neighbor: "münich@example".as_bytes(), + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: host_view, + owned: host_owned, + }, + Case { + label: "Location backslashes", + name: &FieldName::Location, + relaxed: br"/a\b\c", + invalid_neighbor: br"bad\%zz", + strict_rejects: rejects_strict::, + relaxed_rejects: rejects_relaxed::, + borrowed: location_view, + owned: location_owned, + }, + ]; + + for case in cases { + let source = Source { + name: case.name, + value: case.relaxed, + }; + assert!((case.strict_rejects)(&source), "{} must remain strict by default", case.label); + let borrowed = (case.borrowed)(&source); + let owned = (case.owned)(&source); + assert_eq!(borrowed, owned, "{} borrowed/owned parity", case.label); + assert_eq!(borrowed.wire, [case.relaxed.to_vec()], "{} preserves wire bytes", case.label); + + let invalid = Source { + name: case.name, + value: case.invalid_neighbor, + }; + assert!((case.relaxed_rejects)(&invalid), "{} invalid neighbor", case.label); + } +} diff --git a/crates/http_headers/tests/documentation_examples.rs b/crates/http_headers/tests/documentation_examples.rs new file mode 100644 index 000000000..08a45000b --- /dev/null +++ b/crates/http_headers/tests/documentation_examples.rs @@ -0,0 +1,59 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Compile coverage for the concise examples shared across public API docs. + +#![cfg(feature = "headers-all")] + +use std::time::Duration; + +use http_headers::headers::*; +use http_headers::*; + +#[test] +fn header_type_examples_compile() { + let authorization = AuthorizationOwned::::bearer("abc.def").unwrap(); + assert_eq!(authorization.token().unwrap(), b"abc.def"); + + let cache = CacheControlOwned::try_from("max-age=60").unwrap(); + assert_eq!(cache.max_age(), Some(Duration::from_mins(1))); + + let conditional = IfRangeOwned::try_from("\"revision\"").unwrap(); + assert!(matches!(conditional.value().unwrap(), IfRangeValueView::EntityTag(_))); + + assert_eq!(ContentLengthOwned::new(42).get(), 42); + + let content_type = ContentTypeOwned::try_from("text/html; charset=utf-8").unwrap(); + assert_eq!(content_type.parameter("charset").unwrap(), Some(b"utf-8".as_slice())); + + let origin = AccessControlAllowOriginOwned::try_from("https://example.com").unwrap(); + assert_eq!(origin.origin().unwrap(), Some("https://example.com")); + + assert_eq!(ETagOwned::strong("revision").unwrap().opaque_tag().unwrap(), b"revision"); + assert_eq!(LocationOwned::try_from("/next").unwrap().as_str().unwrap(), "/next"); + assert_eq!(HostOwned::try_from("example.com:443").unwrap().host().unwrap(), "example.com"); + assert!(RangeOwned::try_from("bytes=0-99").unwrap().is_bytes()); + + let policy = ReferrerPolicyOwned::new(ReferrerPolicyValue::NoReferrer); + assert_eq!(policy.preferred().unwrap(), ReferrerPolicyValue::NoReferrer); + + let mut cookies = SetCookieOwned::new(); + cookies.push_str("a=1").unwrap(); + assert_eq!(cookies.len(), 1); + + assert_eq!(UserAgentOwned::try_from("client/1").unwrap().as_bytes(), b"client/1"); + assert_eq!(SecWebSocketVersionOwned::new(13).versions().next().unwrap(), 13); +} + +#[test] +fn core_api_examples_compile() { + let error = DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax); + assert_eq!(error.header().as_str(), "content-type"); + + #[cfg(feature = "http")] + { + let mut map = http::HeaderMap::new(); + map.insert(http::header::USER_AGENT, http::HeaderValue::from_static("client/1")); + assert!(UserAgent::view(&map).unwrap().is_some()); + } +} diff --git a/crates/http_headers/tests/duration_precision.rs b/crates/http_headers/tests/duration_precision.rs new file mode 100644 index 000000000..438efd3a1 --- /dev/null +++ b/crates/http_headers/tests/duration_precision.rs @@ -0,0 +1,100 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Whole-second duration contracts for max-age construction and insertion. + +#![cfg(all(feature = "headers-cors", feature = "headers-cache-control", feature = "headers-security"))] + +use std::time::Duration; + +use http_headers::headers::{AccessControlMaxAgeOwned, CacheControlOwned, StrictTransportSecurityOwned}; +use http_headers::{DecodeError, DecodeErrorKind, FieldName}; + +const FRACTIONAL: [Duration; 4] = [ + Duration::from_nanos(1), + Duration::from_nanos(999_999_999), + Duration::from_millis(2_500), + Duration::new(u64::MAX, 1), +]; + +fn assert_invalid_duration(error: DecodeError, name: &FieldName) { + assert_eq!(error.header(), name); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber); + assert_eq!(error.value_index(), None); +} + +#[test] +fn fractional_max_ages_are_rejected_without_flooring() { + for duration in FRACTIONAL { + let error = AccessControlMaxAgeOwned::from_duration(duration).unwrap_err(); + assert_invalid_duration(error, &FieldName::AccessControlMaxAge); + assert_eq!(AccessControlMaxAgeOwned::try_from(duration).unwrap_err(), error); + + assert_invalid_duration( + CacheControlOwned::builder().public().max_age(duration).build().unwrap_err(), + &FieldName::CacheControl, + ); + + let error = StrictTransportSecurityOwned::new(duration).unwrap_err(); + assert_invalid_duration(error, &FieldName::StrictTransportSecurity); + assert_eq!(StrictTransportSecurityOwned::builder(duration).build().unwrap_err(), error); + } +} + +#[test] +fn whole_second_max_ages_round_trip_including_numeric_bounds() { + for seconds in [0, 1, 2, 600, u64::MAX] { + let duration = Duration::from_secs(seconds); + let cors = AccessControlMaxAgeOwned::from_duration(duration).unwrap(); + assert_eq!(cors, AccessControlMaxAgeOwned::try_from(duration).unwrap()); + assert_eq!(cors.seconds(), seconds); + assert_eq!(Duration::from(cors), duration); + assert_eq!(cors.into_field_value().as_bytes(), seconds.to_string().as_bytes()); + + let cache = CacheControlOwned::builder().max_age(duration).build().unwrap(); + assert_eq!(cache.max_age(), Some(duration)); + let hsts = StrictTransportSecurityOwned::builder(duration).build().unwrap(); + assert_eq!(hsts.max_age(), duration); + assert_eq!(hsts.as_field_value().as_bytes(), format!("max-age={seconds}").as_bytes()); + } +} + +#[cfg(feature = "http")] +#[test] +fn invalid_duration_insertion_and_append_preserve_existing_headers() { + use http::{HeaderMap, HeaderValue, header}; + use http_headers::sink::{FieldSink, FieldSinkExt, InsertErrorKind}; + + let mut map = HeaderMap::new(); + map.append(header::CACHE_CONTROL, HeaderValue::from_static("public")); + map.append(header::CACHE_CONTROL, HeaderValue::from_static("max-age=60")); + map.insert(header::ACCESS_CONTROL_MAX_AGE, HeaderValue::from_static("600")); + let original = map.clone(); + + for duration in FRACTIONAL { + let builder = CacheControlOwned::builder().public().max_age(duration); + assert_eq!( + map.set_cache_control(builder.clone()).unwrap_err().kind(), + InsertErrorKind::InvalidValue + ); + assert_eq!(map, original); + assert_eq!( + map.append_encoded(&FieldName::CacheControl, builder).unwrap_err().kind(), + InsertErrorKind::InvalidValue + ); + assert_eq!(map, original); + assert_eq!( + map.set_access_control_max_age_duration(duration).unwrap_err().kind(), + InsertErrorKind::InvalidValue + ); + assert_eq!(map, original); + } + + for seconds in [0, 2, u64::MAX] { + let duration = Duration::from_secs(seconds); + map.set_access_control_max_age_duration(duration).unwrap(); + assert_eq!(map[header::ACCESS_CONTROL_MAX_AGE].as_bytes(), seconds.to_string().as_bytes()); + map.set_cache_control(CacheControlOwned::builder().max_age(duration)).unwrap(); + assert_eq!(map[header::CACHE_CONTROL].as_bytes(), format!("max-age={seconds}").as_bytes()); + } +} diff --git a/crates/http_headers/tests/error_and_encoded_values.rs b/crates/http_headers/tests/error_and_encoded_values.rs new file mode 100644 index 000000000..cf5a4b951 --- /dev/null +++ b/crates/http_headers/tests/error_and_encoded_values.rs @@ -0,0 +1,160 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Public API tests for decode/insert errors and encoded-value collections. + +#![cfg_attr(coverage_nightly, feature(coverage_attribute))] + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod decode_error { + use std::collections::HashSet; + use std::error::Error; + use std::fmt::{self, Write}; + + use http_headers::sink::{InsertError, InsertErrorKind}; + use http_headers::{DecodeError, DecodeErrorKind, FieldName}; + + struct FailAfter { + writes_left: usize, + } + + impl Write for FailAfter { + fn write_str(&mut self, _value: &str) -> fmt::Result { + if self.writes_left == 0 { + Err(fmt::Error) + } else { + self.writes_left -= 1; + Ok(()) + } + } + } + + #[test] + fn error_accessors_display_and_index_saturation_are_structured() { + let unindexed = DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax); + assert_eq!(unindexed.header(), &FieldName::ContentType); + assert_eq!(unindexed.value_index(), None); + assert_eq!(unindexed.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!(unindexed.to_string(), "invalid content-type header: invalid syntax"); + + let indexed = unindexed.at_value(7); + assert_eq!(indexed.value_index(), Some(7)); + assert_eq!(indexed.to_string(), "invalid content-type header: invalid syntax at value 7"); + + let saturated = unindexed.at_value(usize::MAX); + assert_eq!(saturated.value_index(), Some(u32::MAX as usize - 1)); + let error: &dyn Error = &saturated; + assert!(error.source().is_none()); + } + + #[test] + fn every_decode_kind_has_stable_nonsensitive_text() { + let cases = [ + (DecodeErrorKind::MissingValue, "missing value"), + (DecodeErrorKind::UnexpectedMultipleValues, "unexpected multiple values"), + (DecodeErrorKind::InvalidSyntax, "invalid syntax"), + (DecodeErrorKind::InvalidUtf8, "invalid UTF-8"), + (DecodeErrorKind::InvalidToken, "invalid token"), + (DecodeErrorKind::InvalidNumber, "invalid number"), + (DecodeErrorKind::UnterminatedQuote, "unterminated quoted string"), + (DecodeErrorKind::SourceLimitExceeded, "source limit exceeded"), + ]; + for (kind, expected) in cases { + assert_eq!(kind.to_string(), expected); + } + } + + #[test] + fn display_propagates_formatter_failures() { + let error = DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax).at_value(3); + let mut immediate = FailAfter { writes_left: 0 }; + assert!(immediate.write_fmt(format_args!("{error}")).is_err()); + + let mut after_message = FailAfter { writes_left: 4 }; + assert!(after_message.write_fmt(format_args!("{error}")).is_err()); + + let kind = InsertErrorKind::InvalidValue; + let error = InsertError::new(kind); + assert!(immediate.write_fmt(format_args!("{kind}")).is_err()); + assert!(immediate.write_fmt(format_args!("{error}")).is_err()); + } + + #[test] + fn insertion_errors_preserve_compact_distinct_failure_categories() { + const ERROR: InsertError = InsertError::new(InsertErrorKind::CapacityExceeded); + const KIND: InsertErrorKind = ERROR.kind(); + assert_eq!(KIND, InsertErrorKind::CapacityExceeded); + assert!(size_of::() <= size_of::()); + + let cases = [ + (InsertErrorKind::InvalidValue, "invalid field value"), + ( + InsertErrorKind::InvalidEncoding, + "encoded field length does not match the announced length", + ), + (InsertErrorKind::AllocationFailed, "field buffer capacity could not be reserved"), + (InsertErrorKind::CapacityExceeded, "field value or container capacity exceeded"), + ]; + let mut distinct = HashSet::new(); + for (kind, message) in cases { + let error = InsertError::new(kind); + let copies = [error; 2]; + assert_eq!(copies[0], copies[1]); + assert_eq!(error.kind(), kind); + assert_eq!(kind.to_string(), message); + assert_eq!(error.to_string(), message); + assert!(error.source().is_none()); + assert!(format!("{error:?}").contains(&format!("{kind:?}"))); + assert!(distinct.insert(error)); + assert!(!distinct.insert(error)); + } + assert_eq!(distinct.len(), 4); + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod encoded_values { + use std::collections::hash_map::DefaultHasher; + use std::hash::{Hash, Hasher}; + + use http_headers::FieldValue; + use http_headers::sink::EncodedValues; + + #[test] + fn iterators_support_mixed_direction_iteration() { + let values = ["first", "second", "third"] + .into_iter() + .map(FieldValue::from_static) + .collect::(); + + let mut borrowed = values.iter(); + assert_eq!(borrowed.next_back().expect("last value"), "third"); + assert_eq!(borrowed.next().expect("first value"), "first"); + assert_eq!(borrowed.next_back().expect("middle value"), "second"); + assert!(borrowed.next().is_none()); + + let mut owned = values.into_iter(); + assert_eq!(owned.next_back().expect("last owned value"), "third"); + assert_eq!(owned.next().expect("first owned value"), "first"); + assert_eq!(owned.next_back().expect("middle owned value"), "second"); + assert!(owned.next().is_none()); + } + + #[test] + fn comparison_and_hashing_ignore_internal_representation() { + let single = EncodedValues::single(FieldValue::from_static("gzip")); + let vector = EncodedValues::from_vec(vec![FieldValue::from_static("gzip")]); + let greater = EncodedValues::from_vec(vec![FieldValue::from_static("gzip, br")]); + + assert_eq!(single, vector); + assert!(single < greater); + + let mut single_hash = DefaultHasher::new(); + single.hash(&mut single_hash); + let mut vector_hash = DefaultHasher::new(); + vector.hash(&mut vector_hash); + assert_eq!(single_hash.finish(), vector_hash.finish()); + } +} diff --git a/crates/http_headers/tests/feature_boundaries.rs b/crates/http_headers/tests/feature_boundaries.rs new file mode 100644 index 000000000..501e6ae89 --- /dev/null +++ b/crates/http_headers/tests/feature_boundaries.rs @@ -0,0 +1,34 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Integration coverage for independently selectable header families. + +#[cfg(any(feature = "headers-conditional", feature = "headers-range"))] +use http_headers::Field; + +#[cfg(any(feature = "headers-conditional", feature = "headers-range"))] +fn assert_field() {} + +#[cfg(feature = "headers-range")] +#[test] +fn range_family_exposes_range_headers() { + use http_headers::headers::{AcceptRanges, ContentRange, Range}; + + assert_field::(); + assert_field::(); + assert_field::(); +} + +#[cfg(feature = "headers-conditional")] +#[test] +fn conditional_family_includes_etag_headers() { + use http_headers::headers::{ETag, IfMatch, IfModifiedSince, IfNoneMatch, IfRange, IfUnmodifiedSince, LastModified}; + + assert_field::(); + assert_field::(); + assert_field::(); + assert_field::(); + assert_field::(); + assert_field::(); + assert_field::(); +} diff --git a/crates/http_headers/tests/field_lines_iteration.rs b/crates/http_headers/tests/field_lines_iteration.rs new file mode 100644 index 000000000..cd7368ebc --- /dev/null +++ b/crates/http_headers/tests/field_lines_iteration.rs @@ -0,0 +1,81 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Borrowed field-line iteration preserves raw lines and their metadata. + +use http_headers::source::FieldLines; +use http_headers::{FieldName, FieldSensitivity, FieldValue, FieldValueRef}; + +#[test] +fn borrowed_iteration_matches_repeated_without_parsing_or_validating() { + for bytes in [b"".as_slice(), b"a, b", b"raw\r\n"] { + let lines = FieldLines::single(&FieldName::SetCookie, bytes); + let mut iter = (&lines).into_iter(); + assert_eq!(iter.size_hint(), (1, Some(1))); + assert_eq!(iter.next().unwrap().as_bytes(), bytes); + assert_eq!(iter.size_hint(), (0, Some(0))); + assert!(iter.next().is_none()); + assert!(iter.next().is_none()); + assert_eq!((&lines).into_iter().collect::>(), lines.repeated().collect::>()); + assert_eq!(lines.iter().collect::>(), lines.repeated().collect::>()); + } +} + +#[test] +fn owned_and_borrowed_storage_keep_line_boundaries_sensitivity_and_lifetimes() { + let values = [ + FieldValue::from_static("a=1, b=2").with_sensitivity(FieldSensitivity::Sensitive), + FieldValue::from_static(""), + FieldValue::from_static("c=3"), + ]; + let expected = [b"a=1, b=2".as_slice(), b"", b"c=3"]; + let iter = { + let lines = FieldLines::from_slice(&FieldName::SetCookie, &values).unwrap(); + lines.iter() + }; + let copied: Vec<_> = iter.collect(); + assert_eq!(copied.iter().map(|value| value.as_bytes()).collect::>(), expected); + assert_eq!( + copied.iter().map(|value| value.is_sensitive()).collect::>(), + [true, false, false] + ); + for (value, original) in copied.iter().zip(&values) { + assert_eq!(value.as_bytes().as_ptr(), original.as_bytes().as_ptr()); + } + + let refs: Vec<_> = values.iter().map(FieldValue::as_field_value_ref).collect(); + let lines = FieldLines::from_borrowed(&FieldName::SetCookie, &refs).unwrap(); + let mut seen = Vec::new(); + for value in &lines { + seen.push((value.as_bytes(), value.is_sensitive())); + } + assert_eq!(seen, [(expected[0], true), (expected[1], false), (expected[2], false)]); + assert_eq!((&lines).into_iter().map(FieldValueRef::as_bytes).collect::>(), expected); + assert_eq!( + lines + .iter() + .map(|value| (value.as_bytes(), value.is_sensitive())) + .collect::>(), + seen + ); +} + +#[cfg(feature = "http")] +#[test] +fn http_line_iteration_preserves_repetition_and_sensitivity() { + let mut map = http::HeaderMap::new(); + let mut first = http::HeaderValue::from_static("a=1, b=2"); + first.set_sensitive(true); + map.append(http::header::SET_COOKIE, first); + map.append(http::header::SET_COOKIE, http::HeaderValue::from_static("")); + let lines = FieldLines::from_http(&FieldName::SetCookie, map.get_all(http::header::SET_COOKIE)).unwrap(); + let values: Vec<_> = (&lines).into_iter().collect(); + assert_eq!( + values.iter().map(|value| value.as_bytes()).collect::>(), + [b"a=1, b=2".as_slice(), b""] + ); + assert_eq!(values.iter().map(|value| value.is_sensitive()).collect::>(), [true, false]); + assert_eq!(values, lines.repeated().collect::>()); + assert_eq!(values, lines.iter().collect::>()); + assert_eq!(lines.iter().map(FieldValueRef::is_sensitive).collect::>(), [true, false]); +} diff --git a/crates/http_headers/tests/field_value_text.rs b/crates/http_headers/tests/field_value_text.rs new file mode 100644 index 000000000..ec223e050 --- /dev/null +++ b/crates/http_headers/tests/field_value_text.rs @@ -0,0 +1,79 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Fallible text access and validated ownership conversions for field values. + +use bytes::Bytes; +use http_headers::{FieldSensitivity, FieldValue, FieldValueRef, InvalidFieldValue}; + +#[test] +fn try_as_str_borrows_utf8_in_each_owned_representation() { + for value in [ + FieldValue::from_static("gzip"), + FieldValue::from_bytes("münich").unwrap(), + FieldValue::try_from("x".repeat(65)).unwrap(), + FieldValue::from_shared(Bytes::from_static(b"shared")).unwrap(), + ] { + let value = value.with_sensitivity(FieldSensitivity::Sensitive); + let text = value.try_as_str().unwrap(); + assert_eq!(text.as_bytes(), value.as_bytes()); + assert_eq!(text.as_ptr(), value.as_bytes().as_ptr()); + assert!(value.is_sensitive()); + } +} + +#[test] +fn try_as_str_reports_invalid_utf8_without_changing_wire_bytes() { + let value = FieldValue::from_bytes(b"a\xff").unwrap(); + let error = value.try_as_str().unwrap_err(); + assert_eq!(error.valid_up_to(), 1); + assert_eq!(error.error_len(), Some(1)); + assert_eq!(value.as_bytes(), b"a\xff"); +} + +#[test] +fn to_str_borrows_utf8_without_validating_the_field_value_grammar() { + for text in ["", "gzip", "münich", "bad\nvalue"] { + let value = FieldValueRef::new(text.as_bytes()).with_sensitivity(FieldSensitivity::Sensitive); + let borrowed = value.to_str().unwrap(); + assert_eq!(borrowed, text); + assert_eq!(borrowed.as_ptr(), text.as_ptr()); + assert!(value.is_sensitive()); + } + + let value = FieldValueRef::new(b"a\xff"); + let error = value.to_str().unwrap_err(); + assert_eq!(error.valid_up_to(), 1); + assert_eq!(error.error_len(), Some(1)); + assert_eq!(value.as_bytes(), b"a\xff"); +} + +#[test] +fn borrowed_to_owned_conversions_reject_invalid_field_bytes() { + for byte in (0..=0x1f).filter(|&byte| byte != b'\t').chain([0x7f]) { + let embedded = [b'a', byte, b'b']; + for wire in [&embedded[1..2], &embedded[..]] { + for sensitivity in [FieldSensitivity::NonSensitive, FieldSensitivity::Sensitive] { + let value = FieldValueRef::new(wire).with_sensitivity(sensitivity); + assert_eq!(value.try_to_field_value().unwrap_err(), InvalidFieldValue); + assert_eq!(FieldValue::try_from(value).unwrap_err(), InvalidFieldValue); + assert_eq!(value.as_bytes(), wire); + assert_eq!(value.is_sensitive(), sensitivity.is_sensitive()); + } + } + } +} + +#[test] +fn borrowed_to_owned_conversions_preserve_valid_bytes_and_sensitivity() { + let long = [b'a'; 65]; + for wire in [b"".as_slice(), b"good value", b"\t \x80\xff", "münich".as_bytes(), &long] { + for sensitivity in [FieldSensitivity::NonSensitive, FieldSensitivity::Sensitive] { + let value = FieldValueRef::new(wire).with_sensitivity(sensitivity); + for owned in [value.try_to_field_value().unwrap(), FieldValue::try_from(value).unwrap()] { + assert_eq!(owned.as_bytes(), wire); + assert_eq!(owned.is_sensitive(), sensitivity.is_sensitive()); + } + } + } +} diff --git a/crates/http_headers/tests/header_families.rs b/crates/http_headers/tests/header_families.rs new file mode 100644 index 000000000..0b058505d --- /dev/null +++ b/crates/http_headers/tests/header_families.rs @@ -0,0 +1,95 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Public header-family behavior moved out of source modules. + +#![cfg(feature = "headers-all")] +//! +//! These tests cover the `http` adapter, so the file is compiled only when the +//! `http` feature is enabled. + +#![cfg(feature = "http")] + +use http::{HeaderMap, HeaderValue}; +use http_headers::headers::{ContentLength, ContentLengthOwned, ETagOwned, ETagView, SetCookie, SetCookieOwned}; +use http_headers::{DecodeErrorKind, FieldValue, FieldValueRef}; + +const _: [(); size_of::>()] = [(); size_of::>()]; + +#[test] +fn entity_tag_comparison_and_construction() { + let strong = ETagOwned::strong("revision-42").expect("valid strong tag"); + let same = ETagOwned::strong("revision-42").expect("valid strong tag"); + let weak = ETagOwned::weak("revision-42").expect("valid weak tag"); + assert_eq!(strong.strong_eq(&same), Ok(true)); + assert_eq!(strong.strong_eq(&weak), Ok(false)); + assert_eq!(strong.weak_eq(&weak), Ok(true)); + assert_eq!(weak.opaque_tag(), Ok(b"revision-42".as_slice())); + let wire = ETagOwned::weak("abc").expect("valid weak tag").into_field_value(); + assert_eq!(wire, "W/\"abc\""); + assert_eq!( + ETagOwned::strong("bad\"tag").expect_err("quote must be rejected").kind(), + DecodeErrorKind::InvalidSyntax + ); +} + +#[test] +fn content_length_accepts_matching_and_rejects_conflicting_values() { + let mut map = HeaderMap::new(); + map.append("content-length", HeaderValue::from_static("42")); + map.append("content-length", HeaderValue::from_static("42, 42")); + assert_eq!(ContentLength::view(&map), Ok(Some(ContentLengthOwned::new(42)))); + + map.clear(); + map.append("content-length", HeaderValue::from_static("1")); + map.append("content-length", HeaderValue::from_static("2")); + assert_eq!( + ContentLength::view(&map).expect_err("conflicting lengths must fail").kind(), + DecodeErrorKind::InvalidSyntax + ); + + map.clear(); + map.insert("content-length", HeaderValue::from_static("18446744073709551616")); + assert_eq!( + ContentLength::view(&map).expect_err("overflow must fail").kind(), + DecodeErrorKind::InvalidNumber + ); +} + +#[test] +fn set_cookie_preserves_boundaries_sensitivity_and_opaque_attributes() { + let mut map = HeaderMap::new(); + map.append("set-cookie", HeaderValue::from_static("a=1; Path=/")); + map.append("set-cookie", HeaderValue::from_static("b=2; HttpOnly")); + let view = SetCookie::view(&map).expect("valid cookies").expect("cookies present"); + assert_eq!( + view.iter().map(FieldValueRef::as_bytes).collect::>(), + [b"a=1; Path=/".as_slice(), b"b=2; HttpOnly".as_slice()] + ); + + let owned = SetCookie::owned(&map).expect("valid cookies").expect("cookies present"); + assert_eq!(owned.len(), 2); + assert!(owned.iter().all(FieldValue::is_sensitive)); + + let mut inserted = SetCookieOwned::new(); + inserted.push_str("a=1").expect("valid cookie"); + inserted.push_str("b=2").expect("valid cookie"); + let mut inserted_map = HeaderMap::new(); + SetCookie::insert(&mut inserted_map, inserted).expect("map has capacity"); + assert_eq!(inserted_map.get_all("set-cookie").iter().count(), 2); + + let mut secret_map = HeaderMap::new(); + secret_map.insert("set-cookie", HeaderValue::from_static("session=super-secret")); + let debug = format!("{:?}", SetCookie::view(&secret_map).expect("valid cookie").expect("cookie present")); + assert!(!debug.contains("super-secret")); + assert!(debug.contains("value_count")); + + let mut cookies = SetCookieOwned::new(); + cookies + .push_str("session=abc; Max-Age=0") + .expect("cookie grammar belongs to the caller"); + assert_eq!( + cookies.iter().next().map(FieldValue::as_bytes), + Some(b"session=abc; Max-Age=0".as_slice()) + ); +} diff --git a/crates/http_headers/tests/host_debug.rs b/crates/http_headers/tests/host_debug.rs new file mode 100644 index 000000000..3172700d0 --- /dev/null +++ b/crates/http_headers/tests/host_debug.rs @@ -0,0 +1,124 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Sensitive Host diagnostics redact the complete authority. + +#![cfg(feature = "headers-negotiation")] + +use std::fmt::{self, Write as _}; + +use http_headers::headers::{Host, HostKind}; +use http_headers::{DecodeMode, FieldSensitivity, FieldValue, SingleValueField}; + +fn assert_redacted(value: &impl fmt::Debug, name: &str, field_type: &str) { + assert_eq!(format!("{value:?}"), format!("{name} {{ value: {field_type}(Sensitive), .. }}")); + assert_eq!( + format!("{value:#?}"), + format!("{name} {{\n value: {field_type}(Sensitive),\n ..\n}}") + ); +} + +#[test] +fn sensitive_authorities_redact_all_owned_and_borrowed_debug_fields() { + for (wire, host, mode) in [ + ("internal.example:08443", "internal.example", DecodeMode::Strict), + ("192.0.2.123:08443", "192.0.2.123", DecodeMode::Strict), + ("[2001:0DB8::123]:08443", "[2001:0DB8::123]", DecodeMode::Strict), + ("[vF.internal:node]:08443", "[vF.internal:node]", DecodeMode::Strict), + ("münich.example:08443", "münich.example", DecodeMode::Relaxed), + ] { + let value = FieldValue::from_str(wire).unwrap().with_sensitivity(FieldSensitivity::Sensitive); + let borrowed = Host::decode_view_with(value.as_field_value_ref(), mode).unwrap(); + let owned = Host::decode_owned_with(value.clone(), mode).unwrap(); + let owned_view = owned.as_view(); + + assert_redacted(&owned, "HostOwned", "FieldValue"); + for view in [&borrowed, &owned_view] { + assert_redacted(view, "HostView", "FieldValueRef"); + assert_eq!(view.host(), host); + assert_eq!(view.port(), Some("08443")); + assert_eq!(view.network_port(), Ok(Some(8443))); + assert_eq!(view.as_str().unwrap(), wire); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + assert!(view.as_field_value().is_sensitive()); + } + assert_eq!(owned.host().unwrap(), host); + assert_eq!(owned.port().unwrap(), Some("08443")); + assert_eq!(owned.network_port(), Ok(Some(8443))); + assert_eq!(owned.as_str().unwrap(), wire); + assert_eq!(owned.kind(), borrowed.kind()); + assert_eq!(owned_view.kind(), borrowed.kind()); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + assert!(owned.as_field_value().is_sensitive()); + if mode == DecodeMode::Relaxed { + let HostKind::RegisteredName(name) = borrowed.kind() else { + panic!("the international host must remain a registered name"); + }; + assert_eq!(name.normalized(), "xn--mnich-kva.example"); + } + } +} + +#[test] +fn nonsensitive_diagnostics_retain_host_port_address_and_normalization_details() { + for (wire, host, kind, mode) in [ + ("internal.example:08443", "internal.example", "RegisteredName", DecodeMode::Strict), + ("192.0.2.123:08443", "192.0.2.123", "Ipv4(192.0.2.123)", DecodeMode::Strict), + ( + "[2001:0DB8::123]:08443", + "[2001:0DB8::123]", + "Ipv6(2001:db8::123)", + DecodeMode::Strict, + ), + ("münich.example:08443", "münich.example", "RegisteredName", DecodeMode::Relaxed), + ] { + let value = FieldValue::from_str(wire).unwrap(); + let borrowed = Host::decode_view_with(value.as_field_value_ref(), mode).unwrap(); + let owned = Host::decode_owned_with(value.clone(), mode).unwrap(); + let owned_view = owned.as_view(); + let owned_debug = format!("{owned:?}"); + assert!(owned_debug.contains(&format!("value: {value:?}"))); + assert!(owned_debug.contains("parsed: ParsedHost")); + assert!(owned_debug.contains(&format!("kind: {kind}"))); + assert!(owned_debug.contains("numeric_port: Some(Ok(8443))")); + for view in [&borrowed, &owned_view] { + let debug = format!("{view:?}"); + assert!(debug.contains(&format!("host: {host:?}"))); + assert!(debug.contains("port: Some(\"08443\")")); + assert!(debug.contains(&format!("kind: {kind}"))); + assert!(debug.contains("numeric_port: Some(Ok(8443))")); + assert!(!debug.contains("Sensitive")); + if mode == DecodeMode::Relaxed { + assert!(debug.contains("normalized: Some(\"xn--mnich-kva.example\")")); + } + } + assert!(!owned_debug.contains("Sensitive")); + if mode == DecodeMode::Relaxed { + assert!(owned_debug.contains("normalized: Some(\"xn--mnich-kva.example\")")); + } + } +} + +struct RejectWriter; + +impl fmt::Write for RejectWriter { + fn write_str(&mut self, _value: &str) -> fmt::Result { + Err(fmt::Error) + } +} + +#[test] +fn debug_propagates_formatter_errors_with_and_without_redaction() { + for sensitivity in [FieldSensitivity::Sensitive, FieldSensitivity::NonSensitive] { + let value = FieldValue::from_static("internal.example:8443").with_sensitivity(sensitivity); + let borrowed = Host::decode_view(value.as_field_value_ref()).unwrap(); + let owned = Host::decode_owned(value.clone()).unwrap(); + let owned_view = owned.as_view(); + assert_eq!(write!(&mut RejectWriter, "{owned:?}"), Err(fmt::Error)); + assert_eq!(write!(&mut RejectWriter, "{owned:#?}"), Err(fmt::Error)); + for view in [&borrowed, &owned_view] { + assert_eq!(write!(&mut RejectWriter, "{view:?}"), Err(fmt::Error)); + assert_eq!(write!(&mut RejectWriter, "{view:#?}"), Err(fmt::Error)); + } + } +} diff --git a/crates/http_headers/tests/host_origin_components.rs b/crates/http_headers/tests/host_origin_components.rs new file mode 100644 index 000000000..97ed40ba3 --- /dev/null +++ b/crates/http_headers/tests/host_origin_components.rs @@ -0,0 +1,876 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Retained Host and CORS-origin components without the external HTTP adapter. + +#![cfg(any(feature = "headers-negotiation", feature = "headers-cors"))] + +use http_headers::sink::{EncodedValues, FieldSink, InsertError}; +use http_headers::source::{FieldLines, FieldSource, MAX_CUSTOM_FIELD_BYTES}; +use http_headers::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue}; +#[cfg(feature = "headers-negotiation")] +use http_headers::{FieldValueRef, SingleValueField}; + +struct Values { + name: &'static FieldName, + values: Vec, +} + +impl Values { + #[expect(clippy::unwrap_used, reason = "test fixtures must contain valid field values")] + fn new(name: &'static FieldName, wire: &str) -> Self { + Self { + name, + values: vec![FieldValue::try_from(wire).unwrap()], + } + } +} + +impl FieldSource for Values { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::from_slice(name, &self.values)).flatten() + } +} + +impl FieldSink for Values { + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.name = name; + self.values = values.into_iter().collect(); + Ok(()) + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + assert_eq!(name, self.name); + self.values.extend(values); + Ok(()) + } + + fn remove_values(&mut self, name: &'static FieldName) { + assert_eq!(name, self.name); + self.values.clear(); + } +} + +#[cfg(feature = "headers-negotiation")] +mod host { + use std::error::Error; + use std::net::{Ipv4Addr, Ipv6Addr}; + + use http_headers::headers::{Host, HostKind, HostOwned, HostPortView, PortConversionError, PortConversionErrorKind}; + + use super::*; + + #[test] + fn empty_host_is_valid_across_borrowed_owned_and_source_paths() { + let view = ::decode_view(FieldValueRef::new(b"")).unwrap(); + assert_eq!(view.host(), ""); + assert_eq!(view.port(), None); + + let owned = HostOwned::try_from("").unwrap(); + assert_eq!(owned.host().unwrap(), ""); + assert_eq!(owned.port().unwrap(), None); + + let source = Values::new(&FieldName::Host, ""); + let sourced = Host::view(&source).unwrap().unwrap(); + assert_eq!(sourced.host(), ""); + assert_eq!(sourced.port(), None); + + let parsed_error = HostOwned::try_from(":443").unwrap_err().kind(); + assert_eq!(parsed_error, DecodeErrorKind::InvalidSyntax); + assert_eq!(HostOwned::with_port("", 443).unwrap_err().kind(), parsed_error); + assert_eq!( + HostOwned::from_parts(view.kind(), Some(HostPortView::new("443").unwrap())) + .unwrap_err() + .kind(), + parsed_error + ); + } + + #[test] + fn ascii_names_borrow_and_ipv4_classification_does_not_coerce_registered_names() { + for wire in [ + "Example.COM", + "exa_mple~host", + "example%20host", + "!$&'()*+,;=", + "127.1", + "192.168.001.1", + "256.0.0.1", + "127.0.0.1.", + "4294967295", + ] { + let view = ::decode_view(FieldValueRef::new(wire.as_bytes())).unwrap(); + let owned = HostOwned::try_from(wire).unwrap(); + let HostKind::RegisteredName(name) = view.kind() else { + panic!("registered name was coerced: {wire}"); + }; + assert_eq!(name.as_str(), wire); + assert_eq!(name.to_string(), wire); + assert_eq!(name.normalized(), wire); + assert_eq!(name.as_str().as_ptr(), wire.as_ptr()); + assert_eq!(name.normalized().as_ptr(), wire.as_ptr()); + assert_eq!(owned.kind(), view.kind()); + assert_eq!(owned.as_view().kind(), view.kind()); + assert_eq!(HostOwned::from_parts(view.kind(), None).unwrap(), owned); + } + + for (wire, address) in [("0.0.0.0", Ipv4Addr::UNSPECIFIED), ("192.0.2.1", Ipv4Addr::new(192, 0, 2, 1))] { + let owned = HostOwned::try_from(wire).unwrap(); + assert_eq!(owned.kind(), HostKind::Ipv4(address)); + assert_eq!(owned.as_view().kind(), HostKind::Ipv4(address)); + assert_eq!(HostOwned::from_ipv4(address, None), owned); + } + } + + #[test] + fn retained_ipv6_and_ipvfuture_preserve_brackets_and_unbounded_versions() { + let wire = "[2001:0DB8:0:0:0:0:0:1]:000443"; + let address = Ipv6Addr::new(0x2001, 0xdb8, 0, 0, 0, 0, 0, 1); + let owned = HostOwned::try_from(wire).unwrap(); + let view = ::decode_view(FieldValueRef::new(wire.as_bytes())).unwrap(); + assert_eq!(owned.kind(), HostKind::Ipv6(address)); + assert_eq!(view.kind(), owned.kind()); + assert_eq!(view.host(), "[2001:0DB8:0:0:0:0:0:1]"); + assert_eq!(view.port(), Some("000443")); + assert_eq!(view.network_port(), Ok(Some(443))); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!(HostOwned::from_ipv6(address, Some(443)).as_str().unwrap(), "[2001:db8::1]:443"); + + let wire = "[VFFFFFFFFFFFFFFFFFFFFFFFFFFFF.alpha:beta!$&'()*+,;=_~]:65536"; + let owned = HostOwned::try_from(wire).unwrap(); + let view = owned.as_view(); + let HostKind::IpvFuture(future) = view.kind() else { + panic!("IPvFuture was not retained"); + }; + assert_eq!(future.version(), "FFFFFFFFFFFFFFFFFFFFFFFFFFFF"); + assert_eq!(future.address(), "alpha:beta!$&'()*+,;=_~"); + assert_eq!(owned.kind(), view.kind()); + let rebuilt = HostOwned::from_parts(view.kind(), view.port_view()).unwrap(); + assert_eq!(rebuilt.as_str().unwrap(), wire.replacen('V', "v", 1)); + assert_eq!(rebuilt.network_port().unwrap_err().kind(), PortConversionErrorKind::Overflow); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + } + + #[test] + fn ipvfuture_dispatch_preserves_components_and_literal_errors() { + for (literal, version, address) in [ + ("[v1.a]", "1", "a"), + ("[VfF.a:!$&'()*+,;=_~]", "fF", "a:!$&'()*+,;=_~"), + ("[v000000000000000000000000FFFFFFFF.a]", "000000000000000000000000FFFFFFFF", "a"), + ] { + let wire = format!("{literal}:000443"); + let owned = HostOwned::try_from(wire.as_str()).unwrap(); + let view = ::decode_view(FieldValueRef::new(wire.as_bytes())).unwrap(); + let HostKind::IpvFuture(future) = view.kind() else { + panic!("validated IPvFuture was not retained"); + }; + assert_eq!(future.version(), version); + assert_eq!(future.address(), address); + assert_eq!(owned.kind(), view.kind()); + assert_eq!(view.host(), literal); + assert_eq!(owned.host().unwrap(), literal); + assert_eq!(view.network_port(), Ok(Some(443))); + assert_eq!(owned.network_port(), Ok(Some(443))); + assert_eq!(view.port(), Some("000443")); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + } + for (wire, expected) in [ + (b"[v.abc]".as_slice(), DecodeErrorKind::InvalidSyntax), + (b"[v1.]", DecodeErrorKind::InvalidSyntax), + (b"[vG.abc]", DecodeErrorKind::InvalidSyntax), + (b"[V1.a?]", DecodeErrorKind::InvalidSyntax), + (b"[v1.%61]", DecodeErrorKind::InvalidSyntax), + (b"[v1.a/b]", DecodeErrorKind::InvalidSyntax), + (b"[v1.a@b]", DecodeErrorKind::InvalidSyntax), + (b"[v1.a ]", DecodeErrorKind::InvalidSyntax), + (b"[v1.a\xff]", DecodeErrorKind::InvalidSyntax), + (b"[v1.a]suffix", DecodeErrorKind::InvalidSyntax), + (b"[v1.a]:bad", DecodeErrorKind::InvalidNumber), + (b"[v1.a]:bad@", DecodeErrorKind::InvalidSyntax), + ] { + assert_eq!( + ::decode_view(FieldValueRef::new(wire)) + .unwrap_err() + .kind(), + expected, + "{wire:?}" + ); + assert_eq!( + HostOwned::try_from(FieldValue::from_bytes(wire).unwrap()).unwrap_err().kind(), + expected, + "{wire:?}" + ); + } + } + + #[test] + fn ipvfuture_dispatch_leaves_ipv6_values_unchanged() { + for (wire, expected) in [ + ("[::]", Ipv6Addr::UNSPECIFIED), + ("[::1]", Ipv6Addr::LOCALHOST), + ("[DEAD:BEEF::1]", Ipv6Addr::new(0xdead, 0xbeef, 0, 0, 0, 0, 0, 1)), + ("[::ffff:192.0.2.1]", Ipv4Addr::new(192, 0, 2, 1).to_ipv6_mapped()), + ] { + let owned = HostOwned::try_from(wire).unwrap(); + let view = ::decode_view(FieldValueRef::new(wire.as_bytes())).unwrap(); + assert_eq!(owned.kind(), HostKind::Ipv6(expected)); + assert_eq!(view.kind(), HostKind::Ipv6(expected)); + assert_eq!(view.port(), None); + assert_eq!(view.network_port(), Ok(None)); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + } + } + + #[test] + fn ports_distinguish_absence_empty_zero_leading_zeros_and_overflow() { + for (text, expected) in [ + (None, Ok(None)), + (Some(""), Err(PortConversionErrorKind::Empty)), + (Some("0"), Ok(Some(0))), + (Some("00000"), Ok(Some(0))), + (Some("000000000000000000000000000000000000000000443"), Ok(Some(443))), + (Some("6553"), Ok(Some(6553))), + (Some("6554"), Ok(Some(6554))), + (Some("65530"), Ok(Some(65530))), + (Some("65534"), Ok(Some(65534))), + (Some("65535"), Ok(Some(65535))), + (Some("000000000000000000000000000000000000000065535"), Ok(Some(65535))), + (Some("65536"), Err(PortConversionErrorKind::Overflow)), + (Some("65539"), Err(PortConversionErrorKind::Overflow)), + (Some("65540"), Err(PortConversionErrorKind::Overflow)), + (Some("655350"), Err(PortConversionErrorKind::Overflow)), + ( + Some("000000000000000000000000000000000000000065536"), + Err(PortConversionErrorKind::Overflow), + ), + ( + Some("655360000000000000000000000000000000000000000"), + Err(PortConversionErrorKind::Overflow), + ), + ( + Some("999999999999999999999999999999999999999999999"), + Err(PortConversionErrorKind::Overflow), + ), + ] { + for host in ["example.com", "[::1]"] { + let wire = text.map_or_else(|| host.to_owned(), |port| format!("{host}:{port}")); + let owned = HostOwned::try_from(wire.as_str()).unwrap(); + let view = ::decode_view(FieldValueRef::new(wire.as_bytes())).unwrap(); + assert_eq!(owned.network_port().map_err(PortConversionError::kind), expected); + assert_eq!(view.network_port().map_err(PortConversionError::kind), expected); + assert_eq!(owned.port_view().map(HostPortView::as_str), text); + assert_eq!(view.port_view().map(HostPortView::as_str), text); + assert_eq!(view.port_view().map(|port| port.to_string()).as_deref(), text); + assert_eq!(owned.port().unwrap(), text); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + let typed = HostOwned::from_parts(view.kind(), text.map(|text| HostPortView::new(text).unwrap())).unwrap(); + assert_eq!(typed, owned); + } + } + } + + #[test] + fn port_conversion_errors_describe_empty_and_overflow_without_a_source() { + for (text, kind, message) in [ + ("", PortConversionErrorKind::Empty, "the URI port is empty"), + ("65536", PortConversionErrorKind::Overflow, "the URI port exceeds u16::MAX"), + ] { + let error = HostPortView::new(text).unwrap().to_u16().unwrap_err(); + assert_eq!(error.kind(), kind); + assert_eq!(error.to_string(), message); + assert!(error.source().is_none()); + } + } + + #[test] + fn port_overflow_never_masks_invalid_byte_errors() { + for prefix in ["0", "65535", "65536", "999999999999999999999999999"] { + for (suffix, name_error, literal_error) in [ + (b"x".as_slice(), DecodeErrorKind::InvalidNumber, DecodeErrorKind::InvalidNumber), + (b"\t", DecodeErrorKind::InvalidNumber, DecodeErrorKind::InvalidNumber), + (b"/", DecodeErrorKind::InvalidNumber, DecodeErrorKind::InvalidNumber), + (b"x:", DecodeErrorKind::InvalidSyntax, DecodeErrorKind::InvalidNumber), + (b":x", DecodeErrorKind::InvalidSyntax, DecodeErrorKind::InvalidNumber), + (b"x@", DecodeErrorKind::InvalidSyntax, DecodeErrorKind::InvalidSyntax), + (b"x:@", DecodeErrorKind::InvalidSyntax, DecodeErrorKind::InvalidSyntax), + (b"x\x7f", DecodeErrorKind::InvalidSyntax, DecodeErrorKind::InvalidSyntax), + (b"x\x00", DecodeErrorKind::InvalidSyntax, DecodeErrorKind::InvalidSyntax), + (b"x\n", DecodeErrorKind::InvalidSyntax, DecodeErrorKind::InvalidSyntax), + (b"x\xff", DecodeErrorKind::InvalidSyntax, DecodeErrorKind::InvalidSyntax), + ] { + for (host, expected) in [("example.com", name_error), ("[::1]", literal_error)] { + let mut wire = format!("{host}:{prefix}").into_bytes(); + wire.extend_from_slice(suffix); + assert_eq!( + ::decode_view(FieldValueRef::new(&wire)) + .unwrap_err() + .kind(), + expected, + "{wire:?}" + ); + if let Ok(value) = FieldValue::from_bytes(&wire) { + assert_eq!(HostOwned::try_from(value).unwrap_err().kind(), expected, "{wire:?}"); + } + } + } + } + } + + #[test] + fn relaxed_idna_is_retained_and_owned_views_share_normalized_storage() { + let wire = "münich.example:000443"; + let bytes = FieldValue::try_from(wire).unwrap(); + assert_eq!( + ::decode_view(bytes.as_field_value_ref()) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ::decode_owned(bytes.clone()).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + let owned = ::decode_owned_with(bytes.clone(), DecodeMode::Relaxed).unwrap(); + let view = ::decode_view_with(bytes.as_field_value_ref(), DecodeMode::Relaxed).unwrap(); + let borrowed = owned.as_view(); + let HostKind::RegisteredName(name) = owned.kind() else { + panic!("international name was not retained"); + }; + let HostKind::RegisteredName(borrowed_name) = borrowed.kind() else { + panic!("borrowed international name was not retained"); + }; + assert_eq!(name.as_str(), "münich.example"); + assert_eq!(name.normalized(), "xn--mnich-kva.example"); + assert_eq!(name.normalized().as_ptr(), borrowed_name.normalized().as_ptr()); + assert_eq!(view.kind(), owned.kind()); + assert_eq!(view.network_port(), Ok(Some(443))); + assert_eq!(HostOwned::from_parts(view.kind(), view.port_view()).unwrap(), owned); + drop(borrowed); + assert_eq!(owned.into_field_value().as_bytes(), wire.as_bytes()); + let mut sink = Values::new(&FieldName::Host, "placeholder"); + view.insert_into(&mut sink).unwrap(); + assert_eq!(sink.values[0].as_bytes(), wire.as_bytes()); + + for wire in ["\u{200d}.example", "\u{00ad}", "münich@example", "[münich]", "münich:bad"] { + let value = FieldValue::try_from(wire).unwrap(); + assert_eq!( + ::decode_owned_with(value.clone(), DecodeMode::Relaxed) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ::decode_view_with(value.as_field_value_ref(), DecodeMode::Relaxed) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + } + + #[test] + fn typed_composition_preserves_idna_delimiter_validation() { + let wire = "example\u{ff1a}80"; + let value = ::decode_owned_with(FieldValue::try_from(wire).unwrap(), DecodeMode::Relaxed).unwrap(); + let HostKind::RegisteredName(name) = value.kind() else { + panic!("international registered name was not retained"); + }; + assert_eq!(name.normalized(), "example:80"); + assert_eq!(value.network_port(), Ok(None)); + assert_eq!(HostOwned::from_parts(value.kind(), None).unwrap(), value); + assert_eq!( + HostOwned::from_parts(value.kind(), Some(HostPortView::new("443").unwrap())) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ::decode_owned_with(FieldValue::try_from(format!("{wire}:443")).unwrap(), DecodeMode::Relaxed) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + for (wire, normalized, error) in [ + ("\u{ff3b}\u{ff1a}\u{ff1a}1\u{ff3d}", "[::1]", None), + ( + "\u{ff3b}\u{ff1a}\u{ff1a}1\u{ff3d}\u{ff1a}80", + "[::1]:80", + Some(DecodeErrorKind::InvalidNumber), + ), + ] { + let value = ::decode_owned_with(FieldValue::try_from(wire).unwrap(), DecodeMode::Relaxed).unwrap(); + let HostKind::RegisteredName(name) = value.kind() else { + panic!("international registered name was not retained"); + }; + assert_eq!(name.normalized(), normalized); + let typed = HostOwned::from_parts(value.kind(), Some(HostPortView::new("443").unwrap())); + let decoded = + ::decode_owned_with(FieldValue::try_from(format!("{wire}:443")).unwrap(), DecodeMode::Relaxed); + if let Some(error) = error { + assert_eq!(typed.unwrap_err().kind(), error); + assert_eq!(decoded.unwrap_err().kind(), error); + } else { + assert_eq!(typed.unwrap(), decoded.unwrap()); + } + } + } + + #[test] + fn exact_port_error_precedence_and_source_round_trips_are_preserved() { + for (wire, kind) in [ + ("host:999999999999999x", DecodeErrorKind::InvalidNumber), + ("host:999999999999999x:", DecodeErrorKind::InvalidSyntax), + ("host:999999999999999@", DecodeErrorKind::InvalidSyntax), + ("[::1]:999999999999999x", DecodeErrorKind::InvalidNumber), + ("[::1]:999999999999999:", DecodeErrorKind::InvalidNumber), + ("[::1]:999999999999999@", DecodeErrorKind::InvalidSyntax), + ("[::1]suffix", DecodeErrorKind::InvalidSyntax), + ] { + let value = FieldValue::try_from(wire).unwrap(); + assert_eq!(HostOwned::try_from(value.clone()).unwrap_err().kind(), kind, "{wire}"); + assert_eq!( + ::decode_view(value.as_field_value_ref()) + .unwrap_err() + .kind(), + kind, + "{wire}" + ); + } + for (host, kind) in [ + ("host:bad", DecodeErrorKind::InvalidSyntax), + ("host:", DecodeErrorKind::InvalidSyntax), + ("[::1]:", DecodeErrorKind::InvalidNumber), + ] { + assert_eq!(HostOwned::with_port(host, 443).unwrap_err().kind(), kind); + } + + let mut source = Values::new(&FieldName::Host, "[2001:0DB8::1]:000443"); + let owned = ::owned(&source).unwrap().unwrap(); + assert_eq!(::view(&source).unwrap().unwrap().kind(), owned.kind()); + let mut borrowed_sink = Values::new(&FieldName::Host, "placeholder"); + ::view(&source) + .unwrap() + .unwrap() + .insert_into(&mut borrowed_sink) + .unwrap(); + assert_eq!(borrowed_sink.values, source.values); + ::insert(&mut source, owned).unwrap(); + assert_eq!(source.values[0].as_bytes(), b"[2001:0DB8::1]:000443"); + source.values.push(source.values[0].clone()); + assert_eq!( + ::view(&source).unwrap_err().kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!( + ::owned(&source).unwrap_err().kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + source.values = vec![FieldValue::try_from("a".repeat(MAX_CUSTOM_FIELD_BYTES)).unwrap()]; + assert_eq!( + ::view(&source).unwrap().unwrap().host().len(), + MAX_CUSTOM_FIELD_BYTES + ); + source.values = vec![FieldValue::try_from("a".repeat(MAX_CUSTOM_FIELD_BYTES + 1)).unwrap()]; + assert_eq!( + ::view(&source).unwrap_err().kind(), + DecodeErrorKind::SourceLimitExceeded + ); + assert_eq!( + ::owned(&source).unwrap_err().kind(), + DecodeErrorKind::SourceLimitExceeded + ); + } + + #[test] + fn invalid_utf8_preserves_the_existing_entry_point_error_precedence() { + let value = FieldValue::from_bytes(b"\xff").unwrap(); + assert_eq!( + ::decode_view(value.as_field_value_ref()) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!( + ::decode_view_with(value.as_field_value_ref(), mode) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidUtf8 + ); + assert_eq!( + ::decode_owned_with(value.clone(), mode) + .unwrap_err() + .kind(), + if mode == DecodeMode::Strict { + DecodeErrorKind::InvalidSyntax + } else { + DecodeErrorKind::InvalidUtf8 + } + ); + } + } +} + +#[cfg(feature = "headers-cors")] +mod origin { + use std::net::{Ipv4Addr, Ipv6Addr}; + + use http_headers::headers::{ + AccessControlAllowOrigin, AccessControlAllowOriginKind, AccessControlAllowOriginOwned, OriginDomainView, OriginHost, OriginScheme, + }; + + use super::*; + + #[test] + fn wildcard_null_and_origins_have_independent_structured_and_wire_values() { + for (wire, expected) in [ + (" \t*\t ", AccessControlAllowOriginKind::Wildcard), + ("\tnull ", AccessControlAllowOriginKind::Null), + ] { + let source = Values::new(&FieldName::AccessControlAllowOrigin, wire); + let owned = ::owned(&source).unwrap().unwrap(); + let view = ::view(&source).unwrap().unwrap(); + assert_eq!(view.kind(), expected); + assert_eq!(owned.kind(), expected); + assert_eq!(owned.as_view().kind(), expected); + assert_eq!(view.origin(), None); + assert_eq!(owned.origin().unwrap(), None); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + } + assert_ne!( + AccessControlAllowOriginOwned::wildcard(), + AccessControlAllowOriginOwned::try_from(" * ").unwrap() + ); + assert_ne!( + AccessControlAllowOriginOwned::null(), + AccessControlAllowOriginOwned::try_from(" null ").unwrap() + ); + } + + #[test] + fn all_schemes_expose_explicit_and_effective_ports_and_omit_default_ports() { + for (scheme, prefix, default) in [ + (OriginScheme::Ftp, "ftp", 21), + (OriginScheme::Http, "http", 80), + (OriginScheme::Https, "https", 443), + (OriginScheme::Ws, "ws", 80), + (OriginScheme::Wss, "wss", 443), + ] { + assert_eq!(scheme.to_string(), prefix); + for port in [None, Some(0), Some(8443), Some(u16::MAX)] { + let wire = port.map_or_else( + || format!("{prefix}://example.com."), + |port| format!("{prefix}://example.com.:{port}"), + ); + let source = Values::new(&FieldName::AccessControlAllowOrigin, &wire); + let owned = ::owned(&source).unwrap().unwrap(); + let view = ::view(&source).unwrap().unwrap(); + let AccessControlAllowOriginKind::Origin(tuple) = owned.kind() else { + panic!("serialized origin was not retained"); + }; + assert_eq!(view.kind(), owned.kind()); + assert_eq!(owned.as_view().kind(), view.kind()); + assert_eq!(tuple.as_str(), wire); + assert_eq!(tuple.to_string(), wire); + assert_eq!(tuple.scheme(), scheme); + assert_eq!(tuple.port(), port); + assert_eq!(tuple.effective_port(), port.unwrap_or(default)); + let OriginHost::Domain(domain) = tuple.host() else { + panic!("origin domain was not retained"); + }; + assert_eq!(domain.as_str(), "example.com."); + assert_eq!(domain.to_string(), "example.com."); + assert_eq!( + AccessControlAllowOriginOwned::from_parts(scheme, tuple.host(), port).unwrap(), + owned + ); + } + + let domain = OriginHost::Domain(OriginDomainView::new("example.com").unwrap()); + let typed = AccessControlAllowOriginOwned::from_parts(scheme, domain, Some(default)).unwrap(); + assert_eq!(typed.as_str().unwrap(), format!("{prefix}://example.com")); + let AccessControlAllowOriginKind::Origin(tuple) = typed.kind() else { + panic!("typed origin was not retained"); + }; + assert_eq!(tuple.port(), None); + assert_eq!(tuple.effective_port(), default); + assert_eq!( + AccessControlAllowOriginOwned::try_from(format!("{prefix}://example.com:{default}")) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } + } + + #[test] + fn ip_addresses_are_retained_and_typed_ipv6_uses_url_canonicalization() { + let ipv4 = Ipv4Addr::new(192, 0, 2, 128); + let ipv6 = ipv4.to_ipv6_mapped(); + for (wire, host) in [ + ("https://192.0.2.128:8443", OriginHost::Ipv4(ipv4)), + ("https://192.0.2.128.:8443", OriginHost::Ipv4(ipv4)), + ("https://[::ffff:c000:280]:8443", OriginHost::Ipv6(ipv6)), + ("https://[::1]:8443", OriginHost::Ipv6(Ipv6Addr::LOCALHOST)), + ] { + let source = Values::new(&FieldName::AccessControlAllowOrigin, wire); + let owned = ::owned(&source).unwrap().unwrap(); + let view = ::view(&source).unwrap().unwrap(); + let AccessControlAllowOriginKind::Origin(tuple) = owned.kind() else { + panic!("IP origin was not retained"); + }; + assert_eq!(tuple.host(), host); + assert_eq!(tuple.port(), Some(8443)); + assert_eq!(tuple.effective_port(), 8443); + assert_eq!(view.kind(), owned.kind()); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + } + for (port, expected, explicit) in [ + (None, "https://192.0.2.128", None), + (Some(443), "https://192.0.2.128", None), + (Some(8443), "https://192.0.2.128:8443", Some(8443)), + ] { + let typed = AccessControlAllowOriginOwned::from_parts(OriginScheme::Https, OriginHost::Ipv4(ipv4), port).unwrap(); + assert_eq!(typed.as_str().unwrap(), expected); + let AccessControlAllowOriginKind::Origin(tuple) = typed.kind() else { + panic!("typed IPv4 origin was not retained"); + }; + assert_eq!(tuple.host(), OriginHost::Ipv4(ipv4)); + assert_eq!(tuple.port(), explicit); + assert_eq!(tuple.effective_port(), explicit.unwrap_or(443)); + assert_eq!(typed, AccessControlAllowOriginOwned::try_from(expected).unwrap()); + } + let typed = AccessControlAllowOriginOwned::from_parts(OriginScheme::Https, OriginHost::Ipv6(ipv6), Some(443)).unwrap(); + assert_eq!(typed.as_str().unwrap(), "https://[::ffff:c000:280]"); + assert_eq!(typed, AccessControlAllowOriginOwned::try_from("https://[::ffff:c000:280]").unwrap()); + let ties = Ipv6Addr::new(0x2001, 0, 0, 1, 0, 0, 1, 1); + assert_eq!( + AccessControlAllowOriginOwned::from_parts(OriginScheme::Http, OriginHost::Ipv6(ties), None) + .unwrap() + .as_str() + .unwrap(), + "http://[2001::1:0:0:1:1]" + ); + } + + #[test] + fn canonical_ipv6_ascii_boundaries_preserve_components_and_serialization() { + for (literal, expected) in [ + ("::", Ipv6Addr::UNSPECIFIED), + ("::1", Ipv6Addr::LOCALHOST), + ("1:2:3:4:5:6:7:8", Ipv6Addr::new(1, 2, 3, 4, 5, 6, 7, 8)), + ("2001::1:0:0:1:1", Ipv6Addr::new(0x2001, 0, 0, 1, 0, 0, 1, 1)), + ("ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff", Ipv6Addr::from([u16::MAX; 8])), + ] { + let serialized = format!("https://[{literal}]:8443"); + let wire = format!(" \t{serialized}\t "); + let source = Values::new(&FieldName::AccessControlAllowOrigin, &wire); + let owned = ::owned(&source).unwrap().unwrap(); + let view = ::view(&source).unwrap().unwrap(); + let AccessControlAllowOriginKind::Origin(tuple) = view.kind() else { + panic!("canonical IPv6 origin was not retained"); + }; + assert_eq!(tuple.host(), OriginHost::Ipv6(expected)); + assert_eq!(tuple.scheme(), OriginScheme::Https); + assert_eq!(tuple.port(), Some(8443)); + assert_eq!(tuple.effective_port(), 8443); + assert_eq!(tuple.as_str(), serialized); + assert_eq!(owned.kind(), view.kind()); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + let typed = AccessControlAllowOriginOwned::from_parts(OriginScheme::Https, OriginHost::Ipv6(expected), Some(8443)).unwrap(); + assert_eq!(typed.as_str().unwrap(), serialized); + assert_eq!(typed.kind(), view.kind()); + } + } + + #[test] + fn ipv6_ascii_bounds_do_not_change_noncanonical_or_utf8_errors() { + for literal in [ + "", + "ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff0", + "0ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff", + "ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff:", + "ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff::", + "0:0:0:0:0:0:0:1", + "ABCD::1", + "2001:0db8::1", + "::ffff:192.0.2.1", + "::1%eth0", + "::1\t", + "\u{e9}", + ] { + let wire = format!("https://[{literal}]:8443"); + let source = Values::new(&FieldName::AccessControlAllowOrigin, &wire); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!( + ::view_with(&source, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax, + "{wire}" + ); + assert_eq!( + ::owned_with(&source, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax, + "{wire}" + ); + } + } + for prefix in ["", "ffff:ffff:ffff:ffff:ffff:ffff:ffff:ffff"] { + let mut wire = format!("https://[{prefix}").into_bytes(); + wire.push(0xff); + wire.extend_from_slice(b"]:8443"); + let source = Values { + name: &FieldName::AccessControlAllowOrigin, + values: vec![FieldValue::from_bytes(&wire).unwrap()], + }; + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!( + ::view_with(&source, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidUtf8 + ); + assert_eq!( + ::owned_with(&source, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidUtf8 + ); + } + } + } + + #[test] + fn broader_host_grammar_and_noncanonical_origin_forms_remain_rejected() { + for wire in [ + "custom://example.com", + "file://example.com", + "HTTPS://example.com", + "https://Example.com", + "https://münich.example", + "https://example%20host", + "https://exa_mple", + "https://[v1.alpha]", + "https://127.1", + "https://192.168.001.1", + "https://[2001:0db8::1]", + "https://[2001:DB8::1]", + "https://[::ffff:192.0.2.128]", + "https://example.com:443", + "https://example.com:01", + "https://example.com:65536", + "https://example.com:", + "https://user@example.com", + "https://example.com/", + "https://example.com?q", + "https://example.com#fragment", + ] { + let source = Values::new(&FieldName::AccessControlAllowOrigin, wire); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!( + ::view_with(&source, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax, + "{wire}" + ); + assert_eq!( + ::owned_with(&source, mode).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax, + "{wire}" + ); + } + } + for domain in [ + "Example.com", + "münich.example", + "127.0.0.1", + "-invalid", + "a..b", + "name%", + "", + "a".repeat(64).as_str(), + ] { + assert_eq!(OriginDomainView::new(domain).unwrap_err().kind(), DecodeErrorKind::InvalidSyntax); + } + } + + #[test] + fn typed_construction_preserves_http_authority_length_boundaries() { + for (last_label, accepted) in [(59, true), (60, false), (61, false)] { + let domain = format!( + "{}.{}.{}.{}", + "a".repeat(63), + "b".repeat(63), + "c".repeat(63), + "d".repeat(last_label) + ); + let host = OriginHost::Domain(OriginDomainView::new(&domain).unwrap()); + let typed = AccessControlAllowOriginOwned::from_parts(OriginScheme::Https, host, Some(0)); + let parsed = AccessControlAllowOriginOwned::try_from(format!("https://{domain}:0")); + if accepted { + assert_eq!(typed.unwrap(), parsed.unwrap()); + } else { + assert_eq!(typed.unwrap_err().kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!(parsed.unwrap_err().kind(), DecodeErrorKind::InvalidSyntax); + } + let default_port = AccessControlAllowOriginOwned::from_parts(OriginScheme::Https, host, Some(443)).unwrap(); + assert_eq!( + default_port, + AccessControlAllowOriginOwned::try_from(format!("https://{domain}")).unwrap() + ); + let ftp = AccessControlAllowOriginOwned::from_parts(OriginScheme::Ftp, host, Some(0)).unwrap(); + assert_eq!(ftp, AccessControlAllowOriginOwned::try_from(format!("ftp://{domain}:0")).unwrap()); + } + } + + #[test] + fn whitespace_encoding_duplicate_errors_and_source_limits_are_preserved() { + let wire = " \thttps://[2001:db8::1]:65535\t "; + let mut source = Values::new(&FieldName::AccessControlAllowOrigin, wire); + let owned = ::owned(&source).unwrap().unwrap(); + assert_eq!( + ::view(&source).unwrap().unwrap().kind(), + owned.kind() + ); + let mut borrowed_sink = Values::new(&FieldName::AccessControlAllowOrigin, "*"); + ::view(&source) + .unwrap() + .unwrap() + .insert_into(&mut borrowed_sink) + .unwrap(); + assert_eq!(borrowed_sink.values, source.values); + ::insert(&mut source, owned).unwrap(); + assert_eq!(source.values[0].as_bytes(), wire.as_bytes()); + source.values.push(source.values[0].clone()); + assert_eq!( + ::view(&source).unwrap_err().kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + assert_eq!( + ::owned(&source).unwrap_err().kind(), + DecodeErrorKind::UnexpectedMultipleValues + ); + source.values = vec![FieldValue::try_from(format!("{}*", " ".repeat(MAX_CUSTOM_FIELD_BYTES - 1))).unwrap()]; + assert_eq!( + ::view(&source).unwrap().unwrap().kind(), + AccessControlAllowOriginKind::Wildcard + ); + source.values = vec![FieldValue::try_from(format!("{}*", " ".repeat(MAX_CUSTOM_FIELD_BYTES))).unwrap()]; + assert_eq!( + ::view(&source).unwrap_err().kind(), + DecodeErrorKind::SourceLimitExceeded + ); + assert_eq!( + ::owned(&source).unwrap_err().kind(), + DecodeErrorKind::SourceLimitExceeded + ); + source.values = vec![FieldValue::from_bytes(b"\xff").unwrap()]; + assert_eq!( + ::view(&source).unwrap_err().kind(), + DecodeErrorKind::InvalidUtf8 + ); + assert_eq!( + ::owned(&source).unwrap_err().kind(), + DecodeErrorKind::InvalidUtf8 + ); + } +} diff --git a/crates/http_headers/tests/hsts_syntax.rs b/crates/http_headers/tests/hsts_syntax.rs new file mode 100644 index 000000000..9ed477315 --- /dev/null +++ b/crates/http_headers/tests/hsts_syntax.rs @@ -0,0 +1,54 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! RFC 6797 section 6.1 permits empty directives between semicolons. + +#![cfg(feature = "headers-security")] + +use std::time::Duration; + +use http_headers::headers::{StrictTransportSecurity, StrictTransportSecurityOwned}; +use http_headers::{DecodeErrorKind, FieldValue, FieldValueRef, SingleValueField}; + +#[test] +fn empty_directives_are_valid_when_max_age_is_present() { + for (wire, preload) in [ + (";max-age=60", false), + ("max-age=60;", false), + ("max-age=60;;preload", true), + (" \t; \tmax-age=60 ; ; preload ; \t", true), + (";max-age=60;;future=\"a;b\";", false), + ] { + let owned = StrictTransportSecurityOwned::try_from(wire).unwrap(); + assert_eq!(owned.max_age(), Duration::from_mins(1)); + assert_eq!(owned.preload(), preload); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + let view = StrictTransportSecurity::decode_view(FieldValueRef::new(wire.as_bytes())).unwrap(); + assert_eq!(view.max_age(), owned.max_age()); + assert_eq!(view.preload(), preload); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!( + ::decode_owned(FieldValue::from_str(wire).unwrap()).unwrap(), + owned + ); + let directives = owned.directives().map(|item| item.unwrap().as_bytes()).collect::>(); + let expected: &[&[u8]] = if preload { + &[b"max-age=60", b"preload"] + } else if wire.contains("future") { + &[b"max-age=60", b"future=\"a;b\""] + } else { + &[b"max-age=60"] + }; + assert_eq!(directives, expected); + } +} + +#[test] +fn empty_directives_do_not_supply_the_required_max_age() { + for wire in ["", ";", " ; \t; ", ";;preload;"] { + assert_eq!( + StrictTransportSecurityOwned::try_from(wire).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + } +} diff --git a/crates/http_headers/tests/http_map.rs b/crates/http_headers/tests/http_map.rs new file mode 100644 index 000000000..704433827 --- /dev/null +++ b/crates/http_headers/tests/http_map.rs @@ -0,0 +1,283 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! `http::HeaderMap` integration tests. + +#![cfg(feature = "headers-all")] +#![cfg(feature = "http")] + +use std::sync::LazyLock; + +use http::{HeaderMap, HeaderValue}; +use http_headers::headers::{Accept, Allow, SetCookie, SetCookieOwned, UserAgent}; +use http_headers::sink::{ + EncodedValues, FieldEncodeOutput, FieldEncoder, FieldSink as _, FieldSinkExt, FieldValueWriter, InsertError, InsertErrorKind, +}; +use http_headers::source::FieldSource; +use http_headers::{DecodeError, DecodeErrorKind, Field, FieldName, FieldSensitivity, FieldValue, FieldValueRef, SingleValueField}; + +static TRACE_ID: LazyLock = LazyLock::new(|| FieldName::from_static("x-trace-id")); + +struct TraceId(FieldValue); + +struct BytesEncoder { + expected: usize, + bytes: &'static [u8], + sensitive: bool, +} + +impl FieldEncoder for BytesEncoder { + fn encode(self, output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + let sensitivity = if self.sensitive { + http_headers::FieldSensitivity::Sensitive + } else { + http_headers::FieldSensitivity::NonSensitive + }; + let mut writer = output.begin_value(self.expected, sensitivity)?; + writer.write_bytes(self.bytes)?; + writer.finish() + } +} + +#[derive(Clone, Copy)] +#[expect( + dead_code, + reason = "the field proves the view carries the borrowed ref; decode_owned uses the owned value" +)] +struct TraceIdView<'a>(FieldValueRef<'a>); + +impl SingleValueField for TraceId { + type View<'a> = TraceIdView<'a>; + type Owned = Self; + + fn name() -> &'static FieldName { + &TRACE_ID + } + + fn decode_view(value: FieldValueRef<'_>) -> Result, DecodeError> { + Ok(TraceIdView(value)) + } + + fn decode_owned(value: FieldValue) -> Result { + Ok(Self(value)) + } + + fn as_field_value(value: &Self::Owned) -> &FieldValue { + &value.0 + } + + fn into_field_value(value: Self::Owned) -> FieldValue { + value.0 + } +} + +#[test] +fn known_and_custom_names_reach_the_same_storage() { + let mut map = HeaderMap::new(); + map.insert(http::header::USER_AGENT, HeaderValue::from_static("client/1")); + map.insert("x-trace-id", HeaderValue::from_static("abc")); + + assert!(map.contains(&FieldName::UserAgent)); + assert!(map.contains(&TRACE_ID)); + assert_eq!( + FieldSource::lines(&map, ::name()) + .expect("trace ID present") + .len(), + 1 + ); + assert_eq!( + FieldSource::lines(&map, ::name()) + .expect("user agent present") + .len(), + 1 + ); + assert_eq!( + UserAgent::view(&map) + .expect("valid user agent") + .expect("user agent present") + .as_bytes(), + b"client/1" + ); +} + +#[test] +fn repeated_negotiation_errors_preserve_their_error_details() { + let mut map = HeaderMap::new(); + map.append(http::header::ACCEPT, HeaderValue::from_static("text/plain")); + map.append(http::header::ACCEPT, HeaderValue::from_static("\"unterminated")); + + let error = Accept::owned(&map).expect_err("an unterminated quote must be rejected"); + assert_eq!(error.kind(), DecodeErrorKind::UnterminatedQuote); + assert_eq!(error.value_index(), Some(1)); + + map.insert(http::header::ACCEPT, HeaderValue::from_static("invalid")); + let error = Accept::owned(&map).expect_err("an invalid media range must be rejected"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); +} + +#[test] +fn setting_empty_values_removes_the_header() { + let mut map = HeaderMap::new(); + map.insert(http::header::USER_AGENT, HeaderValue::from_static("client/1")); + map.set_values(::name(), EncodedValues::new()) + .expect("removal always fits"); + + assert!(!map.contains(&FieldName::UserAgent)); +} + +#[test] +fn repeated_field_lines_survive_a_round_trip() { + let mut map = HeaderMap::new(); + let mut encoded = EncodedValues::new(); + encoded.push(FieldValue::from_static("a=1")); + encoded.push(FieldValue::from_static("b=2")); + map.set_values(::name(), encoded) + .expect("an empty map has capacity"); + + assert_eq!(map.get_all(http::header::SET_COOKIE).iter().count(), 2); + assert_eq!(SetCookie::view(&map).expect("valid cookies").expect("cookies present").len(), 2); +} + +fn assert_sensitive_cookie_values(map: &HeaderMap, expected: &[&[u8]]) { + let values = map.get_all(http::header::SET_COOKIE).iter().collect::>(); + assert_eq!(values.iter().map(|value| value.as_bytes()).collect::>(), expected); + for value in values { + assert!(value.is_sensitive()); + let debug = format!("{value:?}"); + assert!( + !debug + .as_bytes() + .windows(value.as_bytes().len()) + .any(|bytes| bytes == value.as_bytes()) + ); + } +} + +#[test] +fn typed_cookie_replacement_restores_cleared_sensitivity() { + let mut map = HeaderMap::new(); + map.append(http::header::SET_COOKIE, HeaderValue::from_static("stale=1")); + + let mut replacement = SetCookieOwned::new(); + replacement.push_str("replacement-secret=first").expect("valid cookie"); + replacement.push_str("replacement-secret=second").expect("valid cookie"); + replacement + .iter_mut() + .for_each(|value| value.set_sensitivity(FieldSensitivity::NonSensitive)); + + SetCookie::insert(&mut map, replacement).expect("the map has capacity"); + + assert_sensitive_cookie_values( + &map, + &[b"replacement-secret=first".as_slice(), b"replacement-secret=second".as_slice()], + ); +} + +#[test] +fn fluent_cookie_append_preserves_lines_and_restores_cleared_sensitivity() { + let mut map = HeaderMap::new(); + let mut first = SetCookieOwned::new(); + first.push_str("a=1").expect("valid cookie"); + first.push_str("b=2").expect("valid cookie"); + first + .iter_mut() + .for_each(|value| value.set_sensitivity(FieldSensitivity::NonSensitive)); + let mut second = SetCookieOwned::new(); + second.push_str("c=3").expect("valid cookie"); + second.push_str("d=4").expect("valid cookie"); + second + .iter_mut() + .for_each(|value| value.set_sensitivity(FieldSensitivity::NonSensitive)); + + map.append_set_cookie(first) + .expect("an empty map has capacity") + .append_set_cookie(second) + .expect("the map has capacity"); + + assert_sensitive_cookie_values(&map, &[b"a=1".as_slice(), b"b=2".as_slice(), b"c=3".as_slice(), b"d=4".as_slice()]); +} + +#[test] +fn deferred_map_encoding_covers_numeric_writer_and_custom_removal() { + let mut map = HeaderMap::new(); + map.set_content_length(42) + .expect("numeric field fits") + .set_access_control_max_age(600) + .expect("numeric field fits"); + assert_eq!(map[http::header::CONTENT_LENGTH], "42"); + + map.set_encoded( + &TRACE_ID, + BytesEncoder { + expected: 3, + bytes: b"abc", + sensitive: true, + }, + ) + .expect("custom value fits"); + assert!(map[TRACE_ID.as_str()].is_sensitive()); + map.remove_values(&TRACE_ID); + assert!(!map.contains_key(TRACE_ID.as_str())); +} + +#[test] +fn deferred_map_rejects_bad_encoders_without_replacement() { + let mut map = HeaderMap::new(); + map.insert(http::header::USER_AGENT, HeaderValue::from_static("original")); + + for (encoder, expected_kind) in [ + BytesEncoder { + expected: 2, + bytes: b"x", + sensitive: false, + }, + BytesEncoder { + expected: 1, + bytes: b"\n", + sensitive: false, + }, + BytesEncoder { + expected: 0, + bytes: b"x", + sensitive: false, + }, + BytesEncoder { + expected: usize::MAX, + bytes: b"x", + sensitive: false, + }, + ] + .into_iter() + .zip([ + InsertErrorKind::InvalidEncoding, + InsertErrorKind::InvalidValue, + InsertErrorKind::InvalidEncoding, + InsertErrorKind::CapacityExceeded, + ]) { + assert_eq!( + map.set_encoded(&FieldName::UserAgent, encoder), + Err(InsertError::new(expected_kind)) + ); + assert_eq!(map[http::header::USER_AGENT], "original"); + } +} + +#[test] +fn validated_map_uses_optimized_token_list_decoders() { + let mut map = HeaderMap::new(); + map.insert(http::header::ALLOW, HeaderValue::from_static("GET")); + + assert!(Allow::view(&map).expect("valid borrowed list").is_some()); + assert!(Allow::owned(&map).expect("valid owned list").is_some()); + assert!( + Allow::owned_with(&map, http_headers::DecodeMode::Relaxed) + .expect("valid relaxed owned list") + .is_some() + ); + + map.insert(http::header::ALLOW, HeaderValue::from_static("GET, bad token")); + Allow::owned(&map).expect_err("a token-list member cannot contain a space"); +} diff --git a/crates/http_headers/tests/http_messages.rs b/crates/http_headers/tests/http_messages.rs new file mode 100644 index 000000000..64c4f6a80 --- /dev/null +++ b/crates/http_headers/tests/http_messages.rs @@ -0,0 +1,88 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Direct `http::Request` and `http::Response` field-adapter contracts. + +#![cfg(feature = "headers-all")] +#![cfg(feature = "http")] +#![expect(clippy::unwrap_used, reason = "test failures provide sufficient context")] + +use http_headers::headers::{UserAgent, UserAgentOwned}; +use http_headers::sink::{EncodedValues, FieldEncodeOutput, FieldEncoder, FieldSink, InsertError, InsertErrorKind}; +use http_headers::{FieldName, FieldSensitivity, FieldValue, FieldValueRef}; + +struct RejectEncoder; + +impl FieldEncoder for RejectEncoder { + fn encode(self, _output: &mut O) -> Result<(), InsertError> + where + O: FieldEncodeOutput, + { + Err(InsertError::new(InsertErrorKind::InvalidEncoding)) + } +} + +fn exercise(message: &mut impl FieldSink) { + message + .set_values(&FieldName::UserAgent, EncodedValues::single(FieldValue::from_static("client/1"))) + .unwrap(); + assert_eq!(UserAgent::view(message).unwrap().unwrap().as_bytes(), b"client/1"); + + UserAgent::insert(message, UserAgentOwned::try_from_static("replacement").unwrap()).unwrap(); + message + .append_values(&FieldName::UserAgent, EncodedValues::single(FieldValue::from_static("additional"))) + .unwrap(); + assert_eq!( + message + .lines(&FieldName::UserAgent) + .unwrap() + .repeated() + .map(FieldValueRef::as_bytes) + .collect::>(), + [b"replacement".as_slice(), b"additional".as_slice()] + ); + + message + .set_encoded( + &FieldName::Authorization, + FieldValueRef::new(b"Basic dXNlcjpwYXNz").with_sensitivity(FieldSensitivity::Sensitive), + ) + .unwrap(); + assert!( + message + .lines(&FieldName::Authorization) + .unwrap() + .exactly_one() + .unwrap() + .is_sensitive() + ); + + assert_eq!( + message.set_encoded(&FieldName::UserAgent, RejectEncoder), + Err(InsertError::new(InsertErrorKind::InvalidEncoding)) + ); + assert_eq!( + message.append_encoded(&FieldName::UserAgent, RejectEncoder), + Err(InsertError::new(InsertErrorKind::InvalidEncoding)) + ); + assert_eq!( + message + .lines(&FieldName::UserAgent) + .unwrap() + .repeated() + .map(FieldValueRef::as_bytes) + .collect::>(), + [b"replacement".as_slice(), b"additional".as_slice()] + ); + + message.remove_values(&FieldName::UserAgent); + message.remove_values(&FieldName::Authorization); + assert!(!message.contains(&FieldName::UserAgent)); + assert!(!message.contains(&FieldName::Authorization)); +} + +#[test] +fn request_and_response_delegate_to_their_header_maps() { + exercise(&mut http::Request::new(())); + exercise(&mut http::Response::new(())); +} diff --git a/crates/http_headers/tests/location.rs b/crates/http_headers/tests/location.rs new file mode 100644 index 000000000..0b80e6c7f --- /dev/null +++ b/crates/http_headers/tests/location.rs @@ -0,0 +1,935 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Independent semantic expectations for Location and component construction. + +#![cfg(feature = "headers-location")] + +use std::hash::{DefaultHasher, Hash, Hasher}; + +use http_headers::headers::{Location, LocationOwned, UriAuthority, UriReference}; +use http_headers::source::{FieldLines, FieldSource, MAX_CUSTOM_FIELD_BYTES}; +use http_headers::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, SingleValueField}; + +fn fingerprint(value: &impl Hash) -> u64 { + let mut hasher = DefaultHasher::new(); + value.hash(&mut hasher); + hasher.finish() +} + +#[derive(Debug, Eq, PartialEq)] +struct Components<'a> { + scheme: Option<&'a str>, + authority: Option<(Option<&'a str>, &'a str, Option<&'a str>)>, + path: &'a str, + query: Option<&'a str>, + fragment: Option<&'a str>, +} + +fn components(uri: UriReference<'_>) -> Components<'_> { + Components { + scheme: uri.scheme(), + authority: uri + .authority() + .map(|authority| (authority.userinfo(), authority.host(), authority.port())), + path: uri.path(), + query: uri.query(), + fragment: uri.fragment(), + } +} + +#[test] +fn complete_reference_forms_have_independent_expected_components() { + let cases = [ + ("", None, None, "", None, None), + ("/", None, None, "/", None, None), + ("a", None, None, "a", None, None), + ("../next", None, None, "../next", None, None), + ("./a:b", None, None, "./a:b", None, None), + ("/a:b", None, None, "/a:b", None, None), + ("a/b:c", None, None, "a/b:c", None, None), + ("?", None, None, "", Some(""), None), + ("#", None, None, "", None, Some("")), + ("?#", None, None, "", Some(""), Some("")), + ("?q=a/b?c", None, None, "", Some("q=a/b?c"), None), + ("#a/b?c", None, None, "", None, Some("a/b?c")), + ("/a?#", None, None, "/a", Some(""), Some("")), + ("foo:", Some("foo"), None, "", None, None), + ("foo:/", Some("foo"), None, "/", None, None), + ("mailto:user@example.com", Some("mailto"), None, "user@example.com", None, None), + ("urn:a:b:c", Some("urn"), None, "a:b:c", None, None), + ("//", None, Some((None, "", None)), "", None, None), + ("///path", None, Some((None, "", None)), "/path", None, None), + ("//@:", None, Some((Some(""), "", Some(""))), "", None, None), + ("//host/path", None, Some((None, "host", None)), "/path", None, None), + ("https://host", Some("https"), Some((None, "host", None)), "", None, None), + ("https://host:", Some("https"), Some((None, "host", Some(""))), "", None, None), + ( + "https://host:000443/?#", + Some("https"), + Some((None, "host", Some("000443"))), + "/", + Some(""), + Some(""), + ), + ( + "HTTPS://USER:secret@EXAMPLE.com:999999999999999999999/a?b#c", + Some("HTTPS"), + Some((Some("USER:secret"), "EXAMPLE.com", Some("999999999999999999999"))), + "/a", + Some("b"), + Some("c"), + ), + ( + "https://u%40p:p%3Ass@h%6Fst/a%2fb%FF?q=%23%00#%3F", + Some("https"), + Some((Some("u%40p:p%3Ass"), "h%6Fst", None)), + "/a%2fb%FF", + Some("q=%23%00"), + Some("%3F"), + ), + ( + "//[2001:db8::1]:8443/path", + None, + Some((None, "[2001:db8::1]", Some("8443"))), + "/path", + None, + None, + ), + ( + "//[::ffff:192.0.2.1]", + None, + Some((None, "[::ffff:192.0.2.1]", None)), + "", + None, + None, + ), + ("//[vF.host:!$&]:", None, Some((None, "[vF.host:!$&]", Some(""))), "", None, None), + ("//[V1.future]", None, Some((None, "[V1.future]", None)), "", None, None), + ("//192.0.2.1", None, Some((None, "192.0.2.1", None)), "", None, None), + ("//999.0.2.1", None, Some((None, "999.0.2.1", None)), "", None, None), + ]; + for (wire, scheme, authority, path, query, fragment) in cases { + let expected = Components { + scheme, + authority, + path, + query, + fragment, + }; + let field = FieldValue::from_str(wire).unwrap(); + let owned = LocationOwned::try_from(field.clone()).unwrap(); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let view = ::decode_view_with(field.as_field_value_ref(), mode).unwrap(); + let mode_owned = ::decode_owned_with(field.clone(), mode).unwrap(); + assert_eq!(components(view.uri_reference()), expected, "{wire}"); + assert_eq!(components(mode_owned.uri_reference()), expected, "{wire}"); + assert_eq!(view.uri_reference().as_str(), wire); + assert_eq!(mode_owned.uri_reference(), owned.uri_reference()); + assert_eq!(fingerprint(&view.uri_reference()), fingerprint(&mode_owned.uri_reference())); + assert!(!view.was_normalized()); + assert!(!mode_owned.was_normalized()); + assert_eq!(view.as_bytes(), wire.as_bytes()); + assert_eq!(view.uri_reference().as_str().as_ptr(), view.as_str().unwrap().as_ptr()); + } + let uri = owned.uri_reference(); + let rebuilt = LocationOwned::from_components(uri.scheme(), uri.authority(), uri.path(), uri.query(), uri.fragment()).unwrap(); + assert_eq!(rebuilt.as_str().unwrap(), wire); + assert_eq!(components(rebuilt.uri_reference()), expected, "{wire}"); + assert_eq!(rebuilt, owned); + assert_eq!(fingerprint(&rebuilt), fingerprint(&owned)); + assert!(rebuilt.into_field_value().is_sensitive()); + } +} + +#[test] +fn simple_authority_boundaries_exclude_query_and_fragment_delimiters() { + for (wire, port, path, query, fragment) in [ + ("https://host?next=//other:99/a", None, "", Some("next=//other:99/a"), None), + ("https://host#next://other:99/a", None, "", None, Some("next://other:99/a")), + ( + "https://host:?next=/a:b#fragment://x:y/z", + Some(""), + "", + Some("next=/a:b"), + Some("fragment://x:y/z"), + ), + ( + "https://host:01234?next=/a:b#fragment", + Some("01234"), + "", + Some("next=/a:b"), + Some("fragment"), + ), + ("https://host:?#", Some(""), "", Some(""), Some("")), + ("https://host:/?#", Some(""), "/", Some(""), Some("")), + ( + "https://host/path//to:part?x:/#fragment?y://", + None, + "/path//to:part", + Some("x:/"), + Some("fragment?y://"), + ), + ( + "https://host:9999999999999999999999/path/to:part?next=//other:99#fragment://x", + Some("9999999999999999999999"), + "/path/to:part", + Some("next=//other:99"), + Some("fragment://x"), + ), + ] { + assert!(http_headers_simd::as_simple_uri_reference(wire.as_bytes()).is_some()); + let expected = Components { + scheme: Some("https"), + authority: Some((None, "host", port)), + path, + query, + fragment, + }; + let field = FieldValue::from_str(wire).unwrap(); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let view = ::decode_view_with(field.as_field_value_ref(), mode).unwrap(); + let owned = ::decode_owned_with(field.clone(), mode).unwrap(); + assert_eq!(components(view.uri_reference()), expected, "{wire}"); + assert_eq!(components(owned.uri_reference()), expected, "{wire}"); + assert_eq!(view.as_str().unwrap(), wire); + assert_eq!(owned.as_str().unwrap(), wire); + assert!(!view.was_normalized()); + assert!(!owned.was_normalized()); + } + } +} + +#[test] +fn simple_authority_boundaries_cover_vector_edges_and_long_ports() { + let long_port = "9".repeat(4097); + for length in [0, 1, 14, 15, 16, 17, 30, 31, 32, 33, 62, 63, 64, 65] { + let host = "h".repeat(length + 1); + for port in [None, Some(""), Some("000443"), Some(long_port.as_str())] { + let port_wire = port.map_or_else(String::new, |port| format!(":{port}")); + for path in [String::new(), String::from("/"), format!("/{}", "p".repeat(length))] { + let wire = format!("https://{host}{port_wire}{path}?next=//other:99/a#fragment://x:80/z"); + assert!(http_headers_simd::as_simple_uri_reference(wire.as_bytes()).is_some()); + let expected = Components { + scheme: Some("https"), + authority: Some((None, host.as_str(), port)), + path: &path, + query: Some("next=//other:99/a"), + fragment: Some("fragment://x:80/z"), + }; + let field = FieldValue::from_str(&wire).unwrap(); + let view = ::decode_view(field.as_field_value_ref()).unwrap(); + let owned = LocationOwned::try_from(field.clone()).unwrap(); + assert_eq!(components(view.uri_reference()), expected); + assert_eq!(components(owned.uri_reference()), expected); + assert_eq!(view.as_str().unwrap(), wire); + assert_eq!(owned.as_str().unwrap(), wire); + } + } + } +} + +#[test] +fn general_parser_lends_original_encoded_text_and_complete_components() { + for (wire, expected) in [ + ( + "../a%FF?x=%00#%80", + Components { + scheme: None, + authority: None, + path: "../a%FF", + query: Some("x=%00"), + fragment: Some("%80"), + }, + ), + ( + "mailto:a%FF@example.com", + Components { + scheme: Some("mailto"), + authority: None, + path: "a%FF@example.com", + query: None, + fragment: None, + }, + ), + ( + "//user%FF@[v1.future]:000443/a%00?%FF#%80", + Components { + scheme: None, + authority: Some((Some("user%FF"), "[v1.future]", Some("000443"))), + path: "/a%00", + query: Some("%FF"), + fragment: Some("%80"), + }, + ), + ] { + assert!(http_headers_simd::as_simple_uri_reference(wire.as_bytes()).is_none()); + let field = FieldValue::from_str(wire).unwrap(); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let view = ::decode_view_with(field.as_field_value_ref(), mode).unwrap(); + let owned = ::decode_owned_with(field.clone(), mode).unwrap(); + assert_eq!(components(view.uri_reference()), expected); + assert_eq!(components(owned.uri_reference()), expected); + assert_eq!(view.as_str().unwrap(), wire); + assert_eq!(view.as_str().unwrap().as_ptr(), field.as_bytes().as_ptr()); + assert_eq!(view.uri_reference().as_str().as_ptr(), field.as_bytes().as_ptr()); + assert_eq!(owned.uri_reference().as_str().as_ptr(), owned.as_bytes().as_ptr()); + assert!(!view.was_normalized()); + assert!(!owned.was_normalized()); + let rebuilt = LocationOwned::from_components( + expected.scheme, + view.uri_reference().authority(), + expected.path, + expected.query, + expected.fragment, + ) + .unwrap(); + assert_eq!(components(rebuilt.uri_reference()), expected); + assert_eq!(rebuilt.uri_reference().as_str(), wire); + assert_eq!(rebuilt, owned); + assert_eq!(owned.into_field_value().as_bytes(), wire.as_bytes()); + } + } +} + +#[expect(clippy::unwrap_used, reason = "this assertion helper is used only by integration tests")] +fn assert_invalid_location_bytes(bytes: &[u8]) { + let field = FieldValue::try_from(bytes.to_vec()).unwrap(); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let view_error = ::decode_view_with(field.as_field_value_ref(), mode).unwrap_err(); + let owned_error = ::decode_owned_with(field.clone(), mode).unwrap_err(); + assert_eq!(view_error, owned_error); + assert_eq!(view_error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!(view_error.header(), &FieldName::Location); + assert_eq!(view_error.value_index(), None); + } +} + +#[test] +fn literal_unicode_is_rejected_in_every_component_and_mode() { + for unicode in ["\u{80}", "é", "\u{301}", "界", "🦀", "\u{feff}", "\u{10ffff}"] { + for (prefix, suffix) in [ + ("", ":path"), + ("//", "@host/path"), + ("//", "/path"), + ("//host:", "/path"), + ("//[v1.", "]/path"), + ("/", ""), + ("?q=", ""), + ("#", ""), + (r"https:\\", r"\path"), + (r"https:\\host\", ""), + (r"/a\?q=", ""), + (r"/a\#", ""), + ] { + let wire = format!("{prefix}{unicode}{suffix}"); + assert_invalid_location_bytes(wire.as_bytes()); + } + } +} + +#[test] +fn invalid_utf8_is_rejected_across_ascii_vector_boundaries() { + for invalid in [ + b"\x80".as_slice(), + b"\xff", + b"\xc0\xaf", + b"\xed\xa0\x80", + b"\xf4\x90\x80\x80", + b"\xe2\x82", + ] { + for padding in [0, 14, 15, 16, 30, 31, 32, 62, 63, 64] { + for backslash in [false, true] { + let mut wire = format!("/{}", "a".repeat(padding)).into_bytes(); + if backslash { + wire.push(b'\\'); + } + wire.extend_from_slice(invalid); + wire.extend_from_slice(b"?q=x#f"); + assert_invalid_location_bytes(&wire); + } + } + } +} + +#[test] +fn percent_encoded_octets_stay_opaque_in_strict_and_normalized_references() { + for encoded in [ + "%00", + "%7F", + "%80", + "%FF", + "%c0%af", + "%ED%A0%80", + "%F4%90%80%80", + "%C3%A9", + "%E7%95%8C", + "%F0%9F%A6%80", + ] { + let userinfo = format!("u{encoded}"); + let host = format!("h{encoded}"); + let path = format!("/{encoded}"); + let query = format!("q={encoded}"); + let wire = format!("//{userinfo}@{host}:{path}?{query}#{encoded}"); + let backslash_wire = format!(r"\\{userinfo}@{host}:\{encoded}?{query}#{encoded}"); + let expected = Components { + scheme: None, + authority: Some((Some(userinfo.as_str()), host.as_str(), Some(""))), + path: &path, + query: Some(&query), + fragment: Some(encoded), + }; + for (raw, mode, normalized) in [ + (wire.as_str(), DecodeMode::Strict, false), + (wire.as_str(), DecodeMode::Relaxed, false), + (backslash_wire.as_str(), DecodeMode::Relaxed, true), + ] { + let field = FieldValue::from_str(raw).unwrap(); + let view = ::decode_view_with(field.as_field_value_ref(), mode).unwrap(); + let owned = ::decode_owned_with(field.clone(), mode).unwrap(); + assert_eq!(components(view.uri_reference()), expected); + assert_eq!(components(owned.uri_reference()), expected); + assert_eq!(view.uri_reference().as_str(), wire); + assert_eq!(owned.uri_reference().as_str(), wire); + assert_eq!(view.as_str().unwrap(), raw); + assert_eq!(owned.as_str().unwrap(), raw); + assert_eq!(view.was_normalized(), normalized); + assert_eq!(owned.was_normalized(), normalized); + let forwarded = owned.into_field_value(); + assert_eq!(forwarded.as_bytes(), raw.as_bytes()); + assert!(forwarded.is_sensitive()); + } + let field = FieldValue::from_str(&backslash_wire).unwrap(); + assert_eq!( + ::decode_view(field.as_field_value_ref()) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + } +} + +#[test] +fn byte_general_parsing_preserves_invalid_syntax_errors() { + for bytes in [ + b"\xff".as_slice(), + b"../\xff", + b"https://host/\xff", + b"/a\\\xff", + b"https:\\\\host\\\xc3", + "https://host/café".as_bytes(), + "https:\\\\host\\café".as_bytes(), + b"../%GG", + b"/a\\%GG", + ] { + let field = FieldValue::try_from(bytes.to_vec()).unwrap(); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let view_error = ::decode_view_with(field.as_field_value_ref(), mode).unwrap_err(); + let owned_error = ::decode_owned_with(field.clone(), mode).unwrap_err(); + assert_eq!(view_error, owned_error); + assert_eq!(view_error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!(view_error.header(), &FieldName::Location); + assert_eq!(view_error.value_index(), None); + } + } +} + +#[test] +fn byte_general_parsing_retains_normalized_encoded_components() { + let wire = r"https:\\user:p%FFss@[2001:db8::1]:\a%2Fb?%FF#%80"; + let field = FieldValue::from_str(wire).unwrap(); + let view = ::decode_view_with(field.as_field_value_ref(), DecodeMode::Relaxed).unwrap(); + let owned = ::decode_owned_with(field.clone(), DecodeMode::Relaxed).unwrap(); + let expected = Components { + scheme: Some("https"), + authority: Some((Some("user:p%FFss"), "[2001:db8::1]", Some(""))), + path: "/a%2Fb", + query: Some("%FF"), + fragment: Some("%80"), + }; + assert_eq!(components(view.uri_reference()), expected); + assert_eq!(components(owned.uri_reference()), expected); + assert_eq!( + view.uri_reference().as_str(), + concat!("https://", "user:p%FFss@", "[2001:db8::1]:/a%2Fb?%FF#%80") + ); + assert_eq!(owned.uri_reference().as_str(), view.uri_reference().as_str()); + assert_eq!(view.as_str().unwrap(), wire); + assert_eq!(owned.as_str().unwrap(), wire); + assert!(view.was_normalized()); + assert!(owned.was_normalized()); + let forwarded = owned.into_field_value(); + assert_eq!(forwarded.as_bytes(), wire.as_bytes()); + assert!(forwarded.is_sensitive()); +} + +#[test] +fn owned_ascii_projection_retains_encoded_octets_across_vector_boundaries() { + let encoded = "%00%7f%80%FF"; + let tail = "?x=%FE#%FD"; + let fixed_length = 1 + encoded.len() + tail.len(); + for length in [fixed_length, 31, 32, 33, 63, 64, 65, 127, 128, 129, 255, 256, 257] { + let path = format!("/{}{encoded}", "a".repeat(length - fixed_length)); + let wire = format!("{path}{tail}"); + let owned = LocationOwned::try_from(wire.as_str()).unwrap(); + let expected = Components { + scheme: None, + authority: None, + path: &path, + query: Some("x=%FE"), + fragment: Some("%FD"), + }; + for _ in 0..2 { + assert_eq!(components(owned.uri_reference()), expected); + assert_eq!(owned.uri_reference().as_str(), wire); + assert_eq!(owned.uri_reference().as_str().as_ptr(), owned.as_bytes().as_ptr()); + assert_eq!(owned.as_str().unwrap(), wire); + assert!(!owned.was_normalized()); + } + let clone = owned.clone(); + drop(owned); + assert_eq!(components(clone.uri_reference()), expected); + let field = clone.into_field_value(); + assert_eq!(field.as_bytes(), wire.as_bytes()); + assert!(field.is_sensitive()); + } + for wire in ["", "/", "?", "#", "?#", "//", "//@:", "/%FF", "?%FF", "#%80"] { + let owned = LocationOwned::try_from(wire).unwrap(); + assert_eq!(owned.uri_reference().as_str(), wire); + assert_eq!(owned.as_str().unwrap(), wire); + } +} + +#[test] +fn owned_ascii_projection_keeps_normalized_backing_separate() { + let wire = r"/a\%FF?x=%80#%00"; + let field = FieldValue::from_str(wire).unwrap(); + let owned = ::decode_owned_with(field, DecodeMode::Relaxed).unwrap(); + let expected = Components { + scheme: None, + authority: None, + path: "/a/%FF", + query: Some("x=%80"), + fragment: Some("%00"), + }; + assert_eq!(components(owned.uri_reference()), expected); + assert_eq!(owned.uri_reference().as_str(), "/a/%FF?x=%80#%00"); + assert_eq!(owned.as_str().unwrap(), wire); + assert!(owned.was_normalized()); + let clone = owned.clone(); + drop(owned); + assert_eq!(components(clone.uri_reference()), expected); + assert_eq!(clone.as_str().unwrap(), wire); + let forwarded = clone.into_field_value(); + assert_eq!(forwarded.as_bytes(), wire.as_bytes()); + assert!(forwarded.is_sensitive()); +} + +#[test] +fn relaxed_backslashes_retain_wire_and_normalized_component_boundaries() { + for (wire, semantic, host, path, query, fragment) in [ + (r"/a\b", "/a/b", None, "/a/b", None, None), + (r"\\host\a", "//host/a", Some("host"), "/a", None, None), + (r"https:\\host\a", "https://host/a", Some("host"), "/a", None, None), + ( + r"https://host\@other/x", + "https://host/@other/x", + Some("host"), + "/@other/x", + None, + None, + ), + ( + r"https://host\?q=\#f=\", + "https://host/?q=/#f=/", + Some("host"), + "/", + Some("q=/"), + Some("f=/"), + ), + (r"\?#\", "/?#/", None, "/", Some(""), Some("/")), + ] { + let field = FieldValue::from_str(wire).unwrap(); + assert_eq!( + ::decode_view(field.as_field_value_ref()) + .unwrap_err() + .kind(), + DecodeErrorKind::InvalidSyntax + ); + let view = ::decode_view_with(field.as_field_value_ref(), DecodeMode::Relaxed).unwrap(); + let owned = ::decode_owned_with(field.clone(), DecodeMode::Relaxed).unwrap(); + assert!(view.was_normalized()); + assert!(owned.was_normalized()); + assert_eq!(view.as_str().unwrap(), wire); + assert_eq!(owned.as_bytes(), wire.as_bytes()); + assert_eq!(fingerprint(&owned), fingerprint(&field)); + assert_eq!(fingerprint(&view), fingerprint(&wire)); + let expected = LocationOwned::try_from(semantic).unwrap(); + for uri in [view.uri_reference(), owned.uri_reference()] { + assert_eq!(uri.as_str(), semantic); + assert_eq!(uri.authority().map(UriAuthority::host), host); + assert_eq!(uri.path(), path); + assert_eq!(uri.query(), query); + assert_eq!(uri.fragment(), fragment); + assert_eq!(uri, expected.uri_reference()); + assert_eq!(fingerprint(&uri), fingerprint(&expected.uri_reference())); + } + assert_ne!(owned, expected); + let cloned = view.clone(); + let wire_borrow = view.as_str().unwrap(); + drop(view); + assert_eq!(wire_borrow, wire); + assert_eq!(cloned.uri_reference().as_str(), semantic); + let moved = vec![cloned].pop().unwrap(); + assert_eq!(moved.uri_reference().as_str(), semantic); + let cloned_owned = owned.clone(); + drop(owned); + assert_eq!(cloned_owned.uri_reference().as_str(), semantic); + let encoded = cloned_owned.into_field_value(); + assert_eq!(encoded.as_bytes(), wire.as_bytes()); + assert!(encoded.is_sensitive()); + } +} + +#[test] +fn validated_construction_checks_component_and_contextual_grammar() { + let authority = UriAuthority::new(Some("user:p%40ss"), "[::1]", Some("")).unwrap(); + let constructed = LocationOwned::from_components(Some("https"), Some(authority), "/p%2Fq", Some("a?b/c"), Some("")).unwrap(); + assert_eq!(constructed.as_str().unwrap(), "https://user:p%40ss@[::1]:/p%2Fq?a?b/c#"); + assert_eq!( + constructed.uri_reference().authority(), + Some(UriAuthority::new(Some("user:p%40ss"), "[::1]", Some("")).unwrap()) + ); + for (scheme, authority, path, query, fragment) in [ + (Some(""), None, "", None, None), + (Some("1http"), None, "", None, None), + (Some("ht%74p"), None, "", None, None), + (Some("ht tp"), None, "", None, None), + (None, Some(authority), "relative", None, None), + (Some("s"), Some(authority), "relative", None, None), + (None, None, "//host", None, None), + (Some("s"), None, "//host", None, None), + (None, None, "s:opaque", None, None), + (None, None, "a%2Fb:c", None, None), + (None, None, "/bad?path", None, None), + (None, None, "/bad#path", None, None), + (None, None, "/bad%2", None, None), + (None, None, "/bad space", None, None), + (None, None, "/café", None, None), + (None, None, "", Some("bad#query"), None), + (None, None, "", Some("bad%xx"), None), + (None, None, "", None, Some("bad#fragment")), + (None, None, "", None, Some("bad\\fragment")), + (None, None, "", None, Some("\n")), + ] { + let error = LocationOwned::from_components(scheme, authority, path, query, fragment).unwrap_err(); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!(error.header(), &FieldName::Location); + assert_eq!(error.value_index(), None); + } + let colon = LocationOwned::from_components(None, None, "a%3Ab/c:d", None, None).unwrap(); + assert_eq!(colon.uri_reference().path(), "a%3Ab/c:d"); + let authority_path = LocationOwned::from_components(None, Some(authority), "//p", None, None).unwrap(); + assert_eq!(authority_path.uri_reference().path(), "//p"); +} + +#[test] +fn authority_construction_preserves_full_generic_host_and_port_grammar() { + for host in [ + "", + "host", + "999.0.0.1", + "h%ffst", + "!$&'()*+,;=", + "[::]", + "[::ffff:192.0.2.1]", + "[v1.a:b]", + "[VF.!$&'()*+,;=:_~-]", + ] { + for userinfo in [None, Some(""), Some("user:p%40ss")] { + for port in [None, Some(""), Some("00080"), Some("9999999999999999999999999999999999")] { + let authority = UriAuthority::new(userinfo, host, port).unwrap(); + assert_eq!(authority.userinfo(), userinfo); + assert_eq!(authority.host(), host); + assert_eq!(authority.port(), port); + let built = LocationOwned::from_components(None, Some(authority), "", None, None).unwrap(); + let parsed = LocationOwned::try_from(built.as_str().unwrap()).unwrap(); + assert_eq!(parsed.uri_reference().authority(), Some(authority)); + assert_eq!(fingerprint(&parsed.uri_reference().authority().unwrap()), fingerprint(&authority)); + } + } + } + for (userinfo, host, port) in [ + (Some("a@b"), "host", None), + (Some("a%2"), "host", None), + (Some("a/b"), "host", None), + (None, "[v.a]", None), + (None, "[v1.]", None), + (None, "[vG.a]", None), + (None, "[v1.a%20]", None), + (None, "[::gg]", None), + (None, "[::1", None), + (None, "::1", None), + (None, "[fe80::1%25eth0]", None), + (None, "h%zzst", None), + (None, "host/path", None), + (None, "host?query", None), + (None, "host#fragment", None), + (None, "user@host", None), + (None, "münich.example", None), + (None, "host:443", None), + (None, "host", Some("a")), + (None, "host", Some("-1")), + (None, "host", Some(" 80")), + (None, "host", Some("%38%30")), + ] { + assert_eq!( + UriAuthority::new(userinfo, host, port).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + } +} + +#[test] +fn component_validators_match_the_full_parser_without_delimiter_reinterpretation() { + let candidates = (0_u8..=127) + .map(|byte| char::from(byte).to_string()) + .chain(["", "%00", "%2f", "%FF", "%", "%0", "%GG", "é"].into_iter().map(str::to_owned)); + for candidate in candidates { + let path = format!("/a{candidate}"); + let wire = format!("https://host{path}"); + let expected = + fluent_uri::Uri::parse(&wire).is_ok_and(|uri| uri.path().as_str() == path && uri.query().is_none() && uri.fragment().is_none()); + let authority = UriAuthority::new(None, "host", None).unwrap(); + assert_eq!( + LocationOwned::from_components(Some("https"), Some(authority), &path, None, None).is_ok(), + expected, + "path {candidate:?}" + ); + + let wire = format!("https://host/?{candidate}"); + let expected = fluent_uri::Uri::parse(&wire) + .is_ok_and(|uri| uri.query().is_some_and(|query| query.as_str() == candidate) && uri.fragment().is_none()); + assert_eq!( + LocationOwned::from_components(Some("https"), Some(authority), "/", Some(&candidate), None).is_ok(), + expected, + "query {candidate:?}" + ); + + let wire = format!("https://host/#{candidate}"); + let expected = fluent_uri::Uri::parse(&wire).is_ok_and(|uri| uri.fragment().is_some_and(|fragment| fragment.as_str() == candidate)); + assert_eq!( + LocationOwned::from_components(Some("https"), Some(authority), "/", None, Some(&candidate)).is_ok(), + expected, + "fragment {candidate:?}" + ); + + let wire = format!("//{candidate}@host"); + let expected = fluent_uri::Uri::parse(&wire).is_ok_and(|uri| { + uri.authority().is_some_and(|authority| { + authority.userinfo().is_some_and(|userinfo| userinfo.as_str() == candidate) && authority.host().as_str() == "host" + }) && uri.path().as_str().is_empty() + && uri.query().is_none() + && uri.fragment().is_none() + }); + assert_eq!( + UriAuthority::new(Some(&candidate), "host", None).is_ok(), + expected, + "userinfo {candidate:?}" + ); + + let wire = format!("//{candidate}"); + let expected = fluent_uri::Uri::parse(&wire).is_ok_and(|uri| { + uri.authority().is_some_and(|authority| { + authority.userinfo().is_none() && authority.host().as_str() == candidate && authority.port().is_none() + }) && uri.path().as_str().is_empty() + && uri.query().is_none() + && uri.fragment().is_none() + }); + assert_eq!(UriAuthority::new(None, &candidate, None).is_ok(), expected, "host {candidate:?}"); + + let scheme = format!("s{candidate}"); + let wire = format!("{scheme}:"); + let expected = fluent_uri::Uri::parse(&wire) + .is_ok_and(|uri| uri.scheme().is_some_and(|parsed| parsed.as_str() == scheme) && uri.path().as_str().is_empty()); + assert_eq!( + LocationOwned::from_components(Some(&scheme), None, "", None, None).is_ok(), + expected, + "scheme {candidate:?}" + ); + + let host = format!("[v1.{candidate}]"); + let wire = format!("//{host}"); + let expected = fluent_uri::Uri::parse(&wire).is_ok_and(|uri| { + uri.authority() + .is_some_and(|authority| authority.host().as_str() == host && authority.port().is_none()) + }); + assert_eq!(UriAuthority::new(None, &host, None).is_ok(), expected, "IPvFuture {candidate:?}"); + + let wire = format!("//host:{candidate}"); + let expected = fluent_uri::Uri::parse(&wire).is_ok_and(|uri| { + uri.authority() + .is_some_and(|authority| authority.port() == Some(candidate.as_str())) + && uri.path().as_str().is_empty() + && uri.query().is_none() + && uri.fragment().is_none() + }); + assert_eq!( + UriAuthority::new(None, "host", Some(&candidate)).is_ok(), + expected, + "port {candidate:?}" + ); + } +} + +#[test] +fn sensitive_debug_and_exact_spelling_equality() { + let field = FieldValue::from_static("https://private:secret@host/private?secret#private"); + let view = ::decode_view(field.as_field_value_ref()).unwrap(); + let owned = LocationOwned::try_from(field.clone()).unwrap(); + for debug in [ + format!("{view:?}"), + format!("{owned:?}"), + format!("{:?}", view.uri_reference()), + format!("{:?}", view.uri_reference().authority().unwrap()), + ] { + assert!(debug.contains("redacted")); + assert!(!debug.contains("secret")); + assert!(!debug.contains("private")); + } + assert_eq!(fingerprint(&view), fingerprint(&view.clone())); + assert_eq!(owned, owned.clone()); + assert_eq!(fingerprint(&owned), fingerprint(&field)); + assert_eq!(fingerprint(&view), fingerprint(&view.as_str().unwrap())); + let other_case = LocationOwned::try_from("HTTPS://private:secret@host/private?secret#private").unwrap(); + assert_ne!(owned.uri_reference(), other_case.uri_reference()); + let no_port = UriAuthority::new(None, "", None).unwrap(); + assert_ne!(no_port, UriAuthority::new(None, "", Some("")).unwrap()); + assert_ne!(no_port, UriAuthority::new(Some(""), "", None).unwrap()); +} + +#[test] +fn borrowed_equality_and_hashing_preserve_exact_wire_spelling() { + for (left_wire, right_wire, equal) in [("/a", "/a", true), ("/a", "/A", false), ("/%61", "/a", false)] { + let left_field = FieldValue::from_str(left_wire).unwrap(); + let right_field = FieldValue::from_str(right_wire).unwrap(); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let left = ::decode_view_with(left_field.as_field_value_ref(), mode).unwrap(); + let right = ::decode_view_with(right_field.as_field_value_ref(), mode).unwrap(); + assert_eq!(left == right, equal); + assert_eq!(fingerprint(&left), fingerprint(&left_wire)); + assert_eq!(fingerprint(&right), fingerprint(&right_wire)); + } + } +} + +struct Source<'a>(&'a [FieldValue]); + +impl FieldSource for Source<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::Location) + .then(|| FieldLines::from_slice(name, self.0)) + .flatten() + } +} + +#[test] +fn source_boundaries_and_error_behavior_are_preserved() { + assert!(Location::view(&Source(&[])).unwrap().is_none()); + for wire in ["bad%xx", "https://host/a b", "a:b#c#d", r"bad\%xx", "//[::gg]", "//host:a"] { + let values = [FieldValue::from_str(wire).unwrap()]; + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let view_error = Location::view_with(&Source(&values), mode).unwrap_err(); + let owned_error = Location::owned_with(&Source(&values), mode).unwrap_err(); + assert_eq!(view_error, owned_error); + assert_eq!(view_error.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!(view_error.header(), &FieldName::Location); + assert_eq!(view_error.value_index(), None); + } + } + let duplicates = [FieldValue::from_static("%xx"), FieldValue::from_static("/valid")]; + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let error = Location::view_with(&Source(&duplicates), mode).unwrap_err(); + assert_eq!(error.kind(), DecodeErrorKind::UnexpectedMultipleValues); + assert_eq!(error.value_index(), None); + assert_eq!(error, Location::owned_with(&Source(&duplicates), mode).unwrap_err()); + } + let mut wire = String::from("/"); + wire.extend(std::iter::repeat_n('a', MAX_CUSTOM_FIELD_BYTES - 1)); + let boundary = [FieldValue::from_str(&wire).unwrap()]; + let source = Source(&boundary); + assert_eq!(Location::view(&source).unwrap().unwrap().uri_reference().path(), wire); + wire.push('a'); + let over = [FieldValue::from_str(&wire).unwrap()]; + assert_eq!( + Location::view(&Source(&over)).unwrap_err().kind(), + DecodeErrorKind::SourceLimitExceeded + ); + assert_eq!( + Location::owned(&Source(&over)).unwrap_err().kind(), + DecodeErrorKind::SourceLimitExceeded + ); + assert_eq!(LocationOwned::try_from(wire.as_str()).unwrap().uri_reference().path(), wire); + assert_eq!( + LocationOwned::from_components(None, None, &wire, None, None) + .unwrap() + .uri_reference() + .path(), + wire + ); +} + +struct RawSource<'a>(&'a [u8]); + +impl FieldSource for RawSource<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::Location).then(|| FieldLines::single(name, self.0)) + } +} + +#[test] +fn raw_source_has_the_same_semantics_and_preserves_error_preflight() { + let source = RawSource(br"\\user:secret@host:0123\path?x=%FF#"); + let view = Location::view_with(&source, DecodeMode::Relaxed).unwrap().unwrap(); + let owned = Location::owned_with(&source, DecodeMode::Relaxed).unwrap().unwrap(); + assert_eq!( + components(view.uri_reference()), + Components { + scheme: None, + authority: Some((Some("user:secret"), "host", Some("0123"))), + path: "/path", + query: Some("x=%FF"), + fragment: Some(""), + } + ); + assert_eq!(owned.uri_reference(), view.uri_reference()); + assert_eq!(owned.as_bytes(), source.0); + assert!(owned.into_field_value().is_sensitive()); + + for bytes in [b"/line\nbreak".as_slice(), b"/line\rbreak".as_slice(), &[0xff]] { + let source = RawSource(bytes); + let view_error = Location::view_with(&source, DecodeMode::Relaxed).unwrap_err(); + let owned_error = Location::owned_with(&source, DecodeMode::Relaxed).unwrap_err(); + assert_eq!(view_error, owned_error); + assert_eq!(view_error.header(), &FieldName::Location); + assert_eq!(view_error.kind(), DecodeErrorKind::InvalidSyntax); + } +} + +#[cfg(feature = "http")] +#[test] +fn http_forwarding_retains_original_sensitive_wire() { + let mut headers = http::HeaderMap::new(); + headers.insert("location", http::HeaderValue::from_static(r"https:\\private\path?secret#fragment")); + let view = Location::view_with(&headers, DecodeMode::Relaxed).unwrap().unwrap(); + let owned = Location::owned_with(&headers, DecodeMode::Relaxed).unwrap().unwrap(); + assert_eq!(view.uri_reference(), owned.uri_reference()); + let mut borrowed_sink = http::HeaderMap::new(); + view.insert_into(&mut borrowed_sink).unwrap(); + assert_eq!(borrowed_sink["location"].as_bytes(), headers["location"].as_bytes()); + assert!(borrowed_sink["location"].is_sensitive()); + let mut owned_sink = http::HeaderMap::new(); + Location::insert(&mut owned_sink, owned).unwrap(); + assert_eq!(owned_sink["location"], borrowed_sink["location"]); + assert!(owned_sink["location"].is_sensitive()); +} diff --git a/crates/http_headers/tests/moved_public_api.rs b/crates/http_headers/tests/moved_public_api.rs new file mode 100644 index 000000000..663a61b20 --- /dev/null +++ b/crates/http_headers/tests/moved_public_api.rs @@ -0,0 +1,1674 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Public-API-only header tests moved from library source modules. + +#![cfg(feature = "headers-all")] +//! +//! These tests cover the `http` adapter, so the file is compiled only when the +//! `http` feature is enabled. + +#![cfg(feature = "http")] + +use std::time::{Duration, UNIX_EPOCH}; + +use http::header::*; +use http::{HeaderMap, HeaderName, HeaderValue, Method}; +use http_headers::headers::*; +use http_headers::*; + +mod authorization { + use super::*; + + #[test] + fn schemes_are_matched_case_insensitively_and_may_be_short() { + let mut map = HeaderMap::new(); + map.insert("authorization", HeaderValue::from_static("bEaReR abc.def")); + let view = Authorization::::view(&map) + .expect("valid bearer authorization") + .expect("authorization present"); + assert_eq!(view.token(), b"abc.def"); + + for value in ["Basic X", "Basic ", "Bearer ", "Basic", "Bear"] { + let mut map = HeaderMap::new(); + map.insert( + "authorization", + HeaderValue::from_str(value).expect("fixture is a legal field value"), + ); + let error = Authorization::::view(&map).expect_err("truncated authorization values must be rejected"); + assert_eq!(error.kind(), http_headers::DecodeErrorKind::InvalidSyntax, "{value}"); + } + } + + #[test] + fn borrowed_view_rejects_malformed_basic_base64() { + for value in ["Basic YTpi=", "Basic YQ======", "Basic YTp=", "Basic Oh==", "Basic YTp?"] { + let mut map = HeaderMap::new(); + map.insert( + "authorization", + HeaderValue::from_str(value).expect("fixture is a legal field value"), + ); + let error = Authorization::::view(&map).expect_err("malformed Basic base64 must be rejected"); + assert_eq!(error.kind(), http_headers::DecodeErrorKind::InvalidSyntax, "{value}"); + } + } + + #[test] + fn borrowed_view_rejects_basic_payload_without_colon() { + let mut map = HeaderMap::new(); + map.insert("authorization", HeaderValue::from_static("Basic bm8tY29sb24=")); + let error = Authorization::::view(&map).expect_err("decoded Basic credentials require a colon"); + assert_eq!(error.kind(), http_headers::DecodeErrorKind::InvalidSyntax); + } + + #[test] + fn basic_credentials_debug_is_redacted() { + let authorization = AuthorizationOwned::::basic(b"private-user", b"private-password").expect("valid basic credentials"); + let mut map = HeaderMap::new(); + Authorization::::insert(&mut map, authorization).expect("an empty header map has capacity"); + let authorization = Authorization::::view(&map) + .expect("valid basic authorization") + .expect("authorization present"); + let mut credentials = BasicCredentials::new(); + let credentials = authorization.extract(&mut credentials).expect("valid basic credentials"); + let debug = format!("{credentials:?}"); + assert!(!debug.contains("private-user")); + assert!(!debug.contains("private-password")); + } + + #[test] + fn basic_constructor_matches_rfc_encoding() { + let authorization = AuthorizationOwned::::basic(b"Aladdin", b"open sesame").expect("valid basic credentials"); + assert_eq!(authorization.as_field_value().as_bytes(), b"Basic QWxhZGRpbjpvcGVuIHNlc2FtZQ=="); + } + + #[test] + fn bearer_accepts_padding_and_long_tokens() { + for token in [ + "YWJj=", + "YWJjZA==", + "abcdefghijklmnopqrstuvwxyz0123456789", + "abcdefghijklmnopqrstuvwxyz012345678=", + "a", + ] { + let mut map = HeaderMap::new(); + map.insert( + "authorization", + HeaderValue::from_str(&format!("Bearer {token}")).expect("fixture is a legal field value"), + ); + let view = Authorization::::view(&map) + .expect("valid bearer authorization") + .expect("authorization present"); + assert_eq!(view.token(), token.as_bytes(), "{token}"); + } + } + + #[test] + fn bearer_rejects_malformed_tokens() { + for token in [ + "", + "=", + "==", + "ab=c", + "abcdefghijklmnopqrstuvwxyz01234=6", + "abcdefghijklmnopqrstuvwxyz0123456 ", + "ab c", + ] { + let mut map = HeaderMap::new(); + map.insert( + "authorization", + HeaderValue::from_str(&format!("Bearer {token}")).expect("fixture is a legal field value"), + ); + let error = Authorization::::view(&map).expect_err("malformed bearer tokens must be rejected"); + assert_eq!(error.kind(), http_headers::DecodeErrorKind::InvalidSyntax); + } + } + + #[test] + fn bearer_is_sensitive_and_borrowed() { + let authorization = AuthorizationOwned::::bearer("abc.def").expect("valid bearer token"); + assert!(authorization.as_field_value().is_sensitive()); + let mut map = HeaderMap::new(); + Authorization::::insert(&mut map, authorization).expect("an empty header map has capacity"); + let view = Authorization::::view(&map) + .expect("valid bearer authorization") + .expect("authorization present"); + assert_eq!(view.token(), b"abc.def"); + assert!(map.get("authorization").is_some_and(HeaderValue::is_sensitive)); + } + + #[test] + fn basic_extracts_into_reusable_credentials() { + let authorization = AuthorizationOwned::::basic(b"Aladdin", b"open sesame").expect("valid basic credentials"); + let mut map = HeaderMap::new(); + Authorization::::insert(&mut map, authorization).expect("an empty header map has capacity"); + let authorization = Authorization::::view(&map) + .expect("valid basic authorization") + .expect("authorization present"); + let mut credentials = BasicCredentials::new(); + let credentials = authorization.extract(&mut credentials).expect("valid basic credentials"); + assert_eq!(credentials.username(), b"Aladdin"); + assert_eq!(credentials.password(), b"open sesame"); + } +} + +mod cache_control { + use super::*; + + #[test] + fn parses_multiple_lines_and_preserves_extensions() { + let mut map = HeaderMap::new(); + map.append("cache-control", HeaderValue::from_static("no-cache, x-mode=\"fast, safe\"")); + map.append("cache-control", HeaderValue::from_static("max-age=30")); + let view = CacheControl::view(&map) + .expect("valid cache control") + .expect("cache control present"); + assert!(view.no_cache()); + assert_eq!(view.max_age(), Some(Duration::from_secs(30))); + assert!( + view.directives() + .any(|directive| { directive.name() == "x-mode" && directive.value() == Some(b"\"fast, safe\"".as_slice()) }) + ); + let owned = CacheControl::owned(&map) + .expect("valid cache control") + .expect("cache control present"); + assert!(owned.no_cache()); + assert_eq!(owned.max_age(), Some(Duration::from_secs(30))); + let mut encoded = HeaderMap::new(); + CacheControl::insert(&mut encoded, owned).expect("an empty header map has capacity"); + assert_eq!(encoded["cache-control"], "no-cache, x-mode=\"fast, safe\", max-age=30"); + } + + #[test] + fn recipients_ignore_empty_list_elements_but_senders_reject_them() { + let mut map = HeaderMap::new(); + map.append("cache-control", HeaderValue::from_static(",, no-cache,")); + map.append("cache-control", HeaderValue::from_static(",,,")); + let view = CacheControl::view(&map) + .expect("empty list elements are allowed") + .expect("cache control present"); + assert!(view.no_cache()); + + map.clear(); + map.insert("cache-control", HeaderValue::from_static(",,,")); + let view = CacheControl::view(&map) + .expect("empty recipient list is valid") + .expect("cache control present"); + assert_eq!(view.directives().count(), 0); + let error = CacheControlOwned::try_from(",,,").expect_err("senders must not generate empty list elements"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + } + + #[test] + fn rejects_whitespace_around_equals_for_senders_and_recipients() { + for (value, kind) in [ + ("max-age =30", DecodeErrorKind::InvalidToken), + ("max-age= 30", DecodeErrorKind::InvalidSyntax), + ] { + let mut map = HeaderMap::new(); + map.insert("cache-control", HeaderValue::from_str(value).expect("valid field bytes")); + let error = CacheControl::view(&map).expect_err("whitespace around equals is not in the grammar"); + assert_eq!(error.kind(), kind); + let error = CacheControlOwned::try_from(value).expect_err("sender constructor must reject invalid grammar"); + assert_eq!(error.kind(), kind); + } + } + + #[test] + fn accepts_obs_text_and_quoted_delta_seconds() { + let mut map = HeaderMap::new(); + map.insert( + "cache-control", + HeaderValue::from_bytes(b"max-age=\"30\", x-note=\"\xff\"").expect("valid field value"), + ); + let view = CacheControl::view(&map) + .expect("valid cache control") + .expect("cache control present"); + assert_eq!(view.max_age(), Some(Duration::from_secs(30))); + let extension = view + .directives() + .find(|directive| directive.name() == "x-note") + .expect("extension present"); + assert_eq!(extension.value(), Some(b"\"\xff\"".as_slice())); + assert_eq!( + extension.value_str().expect_err("obs-text is not UTF-8").kind(), + DecodeErrorKind::InvalidUtf8 + ); + } + + #[test] + fn summary_uses_first_valid_max_age() { + let mut map = HeaderMap::new(); + map.insert( + "cache-control", + HeaderValue::from_static("max-age=invalid, max-age=\"45\", max-age=90"), + ); + let view = CacheControl::view(&map) + .expect("valid cache control") + .expect("cache control present"); + assert_eq!(view.max_age(), Some(Duration::from_secs(45))); + } + + #[test] + fn builder_rejects_empty_and_keeps_extensions() { + let error = CacheControlOwned::builder().build().expect_err("empty cache control must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + let header = CacheControlOwned::builder() + .private() + .max_age(Duration::from_mins(5)) + .extension_value("stale-while-revalidate", "30") + .build() + .expect("nonempty cache control"); + assert_eq!(header.max_age(), Some(Duration::from_mins(5))); + assert!( + header + .directives() + .any(|directive| { directive.name() == "stale-while-revalidate" && directive.value() == Some(b"30".as_slice()) }) + ); + } +} + +mod conditional { + use super::*; + + #[test] + fn dates_preserve_valid_obsolete_wire_forms_and_construct_imf_fixdate() { + let obsolete = "Sunday, 06-Nov-94 08:49:37 GMT"; + let header = LastModifiedOwned::try_from(obsolete).expect("valid HTTP-date"); + assert_eq!(header.as_field_value(), obsolete); + let mut encoded = HeaderMap::new(); + LastModified::insert(&mut encoded, header).expect("an empty header map has capacity"); + assert_eq!( + encoded.get_all(LAST_MODIFIED).iter().map(HeaderValue::as_bytes).collect::>(), + vec![obsolete.as_bytes()] + ); + + let date = UNIX_EPOCH + Duration::from_secs(784_111_777); + let canonical = IfModifiedSinceOwned::new(date).expect("representable date"); + assert_eq!(canonical.as_field_value(), "Sun, 06 Nov 1994 08:49:37 GMT"); + assert_eq!(canonical.date(), date); + + let last_modified = LastModifiedOwned::new(date).expect("representable date"); + assert_eq!(last_modified.date(), date); + let if_unmodified = IfUnmodifiedSinceOwned::new(date).expect("representable date"); + assert_eq!(if_unmodified.date(), date); + } + + #[test] + fn parses_wildcards_and_tag_lists_across_field_lines() { + let wildcard = IfMatchOwned::try_from("*").expect("valid wildcard"); + assert!(wildcard.is_wildcard()); + assert_eq!(wildcard.tags().count(), 0); + + let first = ETagOwned::try_from("\"first\"").expect("valid entity tag"); + let second = ETagOwned::try_from("W/\"second\"").expect("valid entity tag"); + let if_match = IfMatchOwned::from_tags([first.clone(), second.clone()]).expect("borrowed array of tags"); + assert_eq!(if_match.tags().count(), 2); + let if_none_match = IfNoneMatchOwned::from_tags(vec![first, second]).expect("owned vector of tags"); + assert_eq!(if_none_match.tags().count(), 2); + + let mut map = HeaderMap::new(); + map.append("if-match", HeaderValue::from_static(",,,")); + map.append("if-match", HeaderValue::from_static("\"one\", W/\"two\"")); + map.append("if-match", HeaderValue::from_static("\"comma,slash\\\"")); + let view = IfMatch::view(&map).expect("valid field").expect("field present"); + let tags: Vec<_> = view.tags().map(|tag| (tag.opaque_tag().to_vec(), tag.is_weak())).collect(); + assert_eq!( + tags, + vec![ + (b"one".to_vec(), false), + (b"two".to_vec(), true), + (b"comma,slash\\".to_vec(), false), + ] + ); + let owned = IfMatch::owned(&map).expect("valid field").expect("field present"); + assert_eq!(owned.tags().count(), 3); + let mut encoded = HeaderMap::new(); + IfMatch::insert(&mut encoded, owned).expect("an empty header map has capacity"); + assert_eq!(encoded.get_all(IF_MATCH).iter().count(), 3); + } + + #[test] + fn decodes_owned_tag_lists_from_one_and_many_field_lines() { + let mut single = HeaderMap::new(); + single.append("if-match", HeaderValue::from_static("\"one,two\", W/\"three\"")); + let owned = IfMatch::owned(&single).expect("valid field").expect("field present"); + assert!(!owned.is_wildcard()); + let tags: Vec<_> = owned.tags().map(|tag| (tag.opaque_tag().to_vec(), tag.is_weak())).collect(); + assert_eq!(tags, vec![(b"one,two".to_vec(), false), (b"three".to_vec(), true)]); + + let mut many = HeaderMap::new(); + many.append("if-none-match", HeaderValue::from_static("\"one\"")); + many.append("if-none-match", HeaderValue::from_static("W/\"two\"")); + let owned = IfNoneMatch::owned(&many).expect("valid field").expect("field present"); + assert_eq!(owned.tags().count(), 2); + + let mut broken = HeaderMap::new(); + broken.append("if-match", HeaderValue::from_static("\"one\"")); + broken.append("if-match", HeaderValue::from_static("nonsense")); + let error = IfMatch::owned(&broken).expect_err("the second field line is invalid"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + + let mut wildcard_conflict = HeaderMap::new(); + wildcard_conflict.append("if-match", HeaderValue::from_static("*")); + wildcard_conflict.append("if-match", HeaderValue::from_static("\"one\"")); + let error = IfMatch::owned(&wildcard_conflict).expect_err("a wildcard never shares a field with a tag"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + + assert!( + IfMatch::owned(&HeaderMap::new()) + .expect("an absent field is not an error") + .is_none() + ); + } + + #[test] + fn borrowed_and_owned_if_range_views_agree() { + let mut map = HeaderMap::new(); + map.insert("if-range", HeaderValue::from_static("\"revision\"")); + let view = IfRange::view(&map).expect("valid field").expect("field present"); + let borrowed = view.value(); + let owned = IfRange::owned(&map).expect("valid field").expect("field present"); + assert_eq!(owned.value(), Ok(borrowed)); + } + + #[test] + fn rejects_mixed_wildcards_malformed_tags_and_missing_lists() { + for wire in ["*, \"tag\"", "\"unterminated", "w/\"lowercase\"", "\"bad tag\""] { + let error = IfNoneMatchOwned::try_from(wire).expect_err("invalid tag list"); + assert!(matches!( + error.kind(), + DecodeErrorKind::InvalidSyntax | DecodeErrorKind::MissingValue + )); + } + let error = IfMatchOwned::try_from(",,,").expect_err("empty list must fail"); + assert_eq!(error.kind(), DecodeErrorKind::MissingValue); + + let mut map = HeaderMap::new(); + map.append("if-none-match", HeaderValue::from_static("*")); + map.append("if-none-match", HeaderValue::from_static("*")); + let error = IfNoneMatch::view(&map).expect_err("duplicate wildcard must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + } + + #[test] + fn if_range_distinguishes_strong_tags_and_dates() { + let tag = IfRangeOwned::try_from("\"revision\"").expect("strong tag"); + let value = tag.value().expect("consistent stored range"); + assert!(matches!(value, IfRangeValueView::EntityTag(_))); + if let IfRangeValueView::EntityTag(tag) = value { + assert_eq!(tag.opaque_tag(), b"revision"); + assert!(!tag.is_weak()); + } + + let date = IfRangeOwned::try_from("Sun, 06 Nov 1994 08:49:37 GMT").expect("date"); + assert!(matches!(date.value().expect("consistent value"), IfRangeValueView::Date(_))); + let error = IfRangeOwned::try_from("W/\"weak\"").expect_err("weak If-Range must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + } +} + +mod content_type { + use super::*; + + #[test] + fn rejects_invalid_parameter_syntax() { + let error = ContentTypeOwned::try_from("text/plain; charset").expect_err("missing parameter value must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + let error = ContentTypeOwned::try_from("text/plain; q=\"unterminated").expect_err("unterminated quote must fail"); + assert_eq!(error.kind(), DecodeErrorKind::UnterminatedQuote); + for (value, kind) in [ + ("text/plain; charset =utf-8", DecodeErrorKind::InvalidSyntax), + ("text/plain; charset= utf-8", DecodeErrorKind::InvalidToken), + ] { + let error = ContentTypeOwned::try_from(value).expect_err("whitespace around equals must fail"); + assert_eq!(error.kind(), kind, "{value}"); + } + } + + #[test] + fn exposes_borrowed_components_and_parameters() { + let content_type = ContentTypeOwned::try_from("text/html; charset=utf-8; level=\"1\"").expect("valid media type"); + assert_eq!(content_type.type_(), Ok("text")); + assert_eq!(content_type.subtype(), Ok("html")); + assert_eq!(content_type.parameter("charset"), Ok(Some(b"utf-8".as_slice()))); + assert_eq!(content_type.parameter("level"), Ok(Some(b"\"1\"".as_slice()))); + assert_eq!(content_type.parameter("missing"), Ok(None)); + } + + #[test] + fn accepts_empty_parameter_slots() { + for value in ["text/plain;", "text/plain; ; charset=utf-8;"] { + let content_type = ContentTypeOwned::try_from(value).expect("empty slots are valid"); + let parameters: Result, _> = content_type.parameters().collect(); + let parameters = parameters.expect("stored media type remains valid"); + if value.contains("charset") { + assert_eq!(parameters.len(), 1); + assert_eq!(parameters[0].name(), "charset"); + } else { + assert!(parameters.is_empty()); + } + } + } +} + +mod cors { + use std::collections::hash_map::DefaultHasher; + use std::hash::{Hash, Hasher}; + + use super::*; + + #[test] + fn header_name_lists_borrow_preserve_case_and_duplicates() { + let mut map = HeaderMap::new(); + map.append(ACCESS_CONTROL_REQUEST_HEADERS, HeaderValue::from_static("X-Trace, content-type")); + map.append(ACCESS_CONTROL_REQUEST_HEADERS, HeaderValue::from_static("x-trace")); + let raw = map.get(ACCESS_CONTROL_REQUEST_HEADERS).expect("inserted value"); + let view = AccessControlRequestHeaders::view(&map) + .expect("valid header-name list") + .expect("present header-name list"); + let names: Vec<&str> = view.iter().map(FieldNameView::as_str).collect(); + assert_eq!(names, ["X-Trace", "content-type", "x-trace"]); + assert_eq!(view.len(), 3); + assert!(!view.is_empty()); + assert_eq!(view.field_values().count(), 2); + assert!(format!("{view:?}").contains("header_name_count")); + assert_eq!(names[0].as_ptr(), raw.as_bytes().as_ptr()); + assert!(view.iter().next().expect("first name").eq_ignore_ascii_case("x-trace")); + assert_eq!( + view.iter().next().expect("first name").try_to_field_name().expect("validated name"), + "x-trace" + ); + assert_eq!( + view.iter() + .next() + .expect("first name") + .try_to_http_header_name() + .expect("validated HTTP name"), + "x-trace" + ); + + map.clear(); + map.insert(ACCESS_CONTROL_EXPOSE_HEADERS, HeaderValue::from_static("X-Trace")); + let exposed = AccessControlExposeHeaders::view(&map) + .expect("valid exposed-header list") + .expect("present exposed-header list"); + assert_eq!(exposed.len(), 1); + assert!(!exposed.is_empty()); + assert_eq!(exposed.field_values().count(), 1); + assert!(!exposed.contains_wildcard()); + assert!(format!("{exposed:?}").contains("header_name_count")); + assert_eq!(exposed.iter().map(FieldNameView::as_str).collect::>(), ["X-Trace"]); + map.insert(ACCESS_CONTROL_EXPOSE_HEADERS, HeaderValue::from_static("*")); + let wildcard = AccessControlExposeHeaders::view(&map) + .expect("valid wildcard exposed-header list") + .expect("present wildcard exposed-header list"); + assert!(wildcard.is_wildcard()); + + map.clear(); + map.insert(ACCESS_CONTROL_ALLOW_HEADERS, HeaderValue::from_static("*")); + let allowed = AccessControlAllowHeaders::view(&map) + .expect("valid allowed-header list") + .expect("present allowed-header list"); + assert_eq!(allowed.len(), 1); + assert!(!allowed.is_empty()); + assert_eq!(allowed.field_values().count(), 1); + assert!(allowed.contains_wildcard()); + assert!(allowed.is_wildcard()); + assert!(format!("{allowed:?}").contains("header_name_count")); + } + + #[test] + fn list_constructors_validate_tokens_and_keep_wildcard_syntactic() { + let methods = AccessControlAllowMethodsOwned::from_methods([Method::GET.as_str(), "X-CUSTOM", "GET"]).expect("valid methods"); + assert_eq!( + methods.iter().map(MethodView::as_str).collect::>(), + ["GET", "X-CUSTOM", "GET"] + ); + AccessControlAllowMethodsOwned::from_methods(["GET", "not a method"]).expect_err("methods must be tokens"); + + let headers = AccessControlExposeHeadersOwned::from_header_names([http::header::CONTENT_TYPE.as_str(), "X-Extension"]) + .expect("valid field names"); + assert_eq!( + headers.iter().map(FieldNameView::as_str).collect::>(), + ["content-type", "X-Extension"] + ); + assert_eq!(headers.len(), 2); + assert!(!headers.is_empty()); + assert_eq!(headers.field_values().count(), 1); + assert_eq!(headers.clone().into_field_values().len(), 1); + assert!(format!("{headers:?}").contains("header_name_count")); + let expose_adopted = + AccessControlExposeHeadersOwned::try_from(FieldValue::from_static("x-trace")).expect("single exposed field line"); + assert_eq!(expose_adopted.len(), 1); + let expose_lines = + AccessControlExposeHeadersOwned::try_from(vec![FieldValue::from_static("x-trace"), FieldValue::from_static("content-type")]) + .expect("repeated exposed field lines"); + assert_eq!(expose_lines.len(), 2); + let expose_wildcard = AccessControlExposeHeadersOwned::wildcard(); + assert!(expose_wildcard.contains_wildcard()); + assert!(expose_wildcard.is_wildcard()); + assert!(AccessControlExposeHeadersOwned::empty().is_empty()); + AccessControlExposeHeadersOwned::from_header_names(["bad name"]).expect_err("field names must be tokens"); + + let wildcard = AccessControlAllowHeadersOwned::wildcard(); + assert!(wildcard.is_wildcard()); + let mixed = AccessControlAllowHeadersOwned::from_header_names(["*", "authorization"]).expect("wildcard is also a field-name token"); + assert!(mixed.contains_wildcard()); + assert!(!mixed.is_wildcard()); + } + + #[test] + fn list_constructors_execute_external_generic_shapes() { + let single = AccessControlAllowHeadersOwned::from_header_names(["content-type"]).expect("single borrowed field name"); + assert_eq!(single, single.clone()); + assert_eq!(single.len(), 1); + assert!(!single.is_empty()); + assert_eq!(single.field_values().count(), 1); + assert_eq!(single.clone().into_field_values().len(), 1); + assert!(format!("{single:?}").contains("header_name_count")); + assert!(AccessControlAllowHeadersOwned::empty().is_empty()); + let allow_lines = + AccessControlAllowHeadersOwned::try_from(vec![FieldValue::from_static("content-type"), FieldValue::from_static("x-trace")]) + .expect("repeated allowed field lines"); + assert_eq!(allow_lines.len(), 2); + let adopted = AccessControlAllowHeadersOwned::try_from(FieldValue::from_static("content-type")).expect("single adopted field line"); + assert_eq!(adopted, single); + assert_eq!(single.iter().next().expect("field name").as_str(), "content-type"); + + let three = AccessControlAllowHeadersOwned::from_header_names(["content-type", "authorization", "x-trace-id"]) + .expect("three borrowed field names"); + assert_eq!(three.len(), 3); + let mut first_hash = DefaultHasher::new(); + single.hash(&mut first_hash); + let mut second_hash = DefaultHasher::new(); + single.clone().hash(&mut second_hash); + assert_eq!(first_hash.finish(), second_hash.finish()); + + let owned = AccessControlRequestHeadersOwned::from_header_names(vec![String::from("content-type"), String::from("x-trace-id")]) + .expect("owned field names"); + assert_eq!(owned.len(), 2); + assert!(!owned.is_empty()); + assert_eq!(owned.field_values().count(), 1); + assert_eq!(owned.clone().into_field_values().len(), 1); + assert!(format!("{owned:?}").contains("header_name_count")); + let request_adopted = + AccessControlRequestHeadersOwned::try_from(FieldValue::from_static("content-type")).expect("single request field line"); + assert_eq!(request_adopted.len(), 1); + AccessControlRequestHeadersOwned::from_header_names(Vec::::new()).expect_err("request header list must not be empty"); + AccessControlRequestHeadersOwned::from_header_names(vec![String::from("bad name")]).expect_err("owned invalid field name"); + AccessControlRequestHeadersOwned::try_from(vec![FieldValue::from_static("")]) + .expect_err("adopted request-header lists must contain an item"); + + let method = AccessControlAllowMethodsOwned::from_methods(["GET"]).expect("single borrowed method"); + assert_eq!(method.iter().next().expect("method").as_str(), "GET"); + let standard_methods = + AccessControlAllowMethodsOwned::from_methods(["GET", "PUT", "HEAD", "POST", "PATCH", "TRACE", "DELETE", "CONNECT", "OPTIONS"]) + .expect("standard methods"); + assert_eq!(standard_methods.len(), 9); + let common_names = AccessControlExposeHeadersOwned::from_header_names([ + "accept", + "authorization", + "cache-control", + "content-length", + "content-type", + "etag", + "origin", + "x-requested-with", + ]) + .expect("common field names"); + assert_eq!(common_names.len(), 8); + let owned_methods = + AccessControlAllowMethodsOwned::from_methods(vec![String::from("GET"), String::from("X-CUSTOM")]).expect("owned methods"); + assert_eq!(owned_methods.len(), 2); + AccessControlAllowMethodsOwned::from_methods(vec![String::from("bad method")]).expect_err("owned invalid method"); + } + + #[test] + fn malformed_list_items_report_the_field_line() { + let mut map = HeaderMap::new(); + map.append(ACCESS_CONTROL_ALLOW_METHODS, HeaderValue::from_static("GET")); + map.append(ACCESS_CONTROL_ALLOW_METHODS, HeaderValue::from_static("bad method")); + let error = AccessControlAllowMethods::view(&map).expect_err("invalid method token must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidToken); + assert_eq!(error.value_index(), Some(1)); + + map.clear(); + map.insert(ACCESS_CONTROL_ALLOW_HEADERS, HeaderValue::from_static("\"x-header\"")); + let error = AccessControlAllowHeaders::view(&map).expect_err("quoted field name must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidToken); + } + + #[test] + fn singleton_cors_headers_reject_multiple_field_lines() { + let mut map = HeaderMap::new(); + for (name, value) in [ + (ACCESS_CONTROL_ALLOW_ORIGIN, "*"), + (ACCESS_CONTROL_ALLOW_CREDENTIALS, "true"), + (ACCESS_CONTROL_MAX_AGE, "60"), + (ACCESS_CONTROL_REQUEST_METHOD, "PATCH"), + ] { + map.clear(); + map.append(name.clone(), HeaderValue::from_str(value).expect("static test value")); + map.append(name.clone(), HeaderValue::from_str(value).expect("static test value")); + let error = if name == ACCESS_CONTROL_ALLOW_ORIGIN { + AccessControlAllowOrigin::view(&map).expect_err("duplicate singleton must fail") + } else if name == ACCESS_CONTROL_ALLOW_CREDENTIALS { + AccessControlAllowCredentials::view(&map).expect_err("duplicate singleton must fail") + } else if name == ACCESS_CONTROL_MAX_AGE { + AccessControlMaxAge::view(&map).expect_err("duplicate singleton must fail") + } else { + AccessControlRequestMethod::view(&map).expect_err("duplicate singleton must fail") + }; + assert_eq!(error.kind(), DecodeErrorKind::UnexpectedMultipleValues); + } + } + + #[test] + fn request_method_borrows_and_preserves_extensions() { + let mut map = HeaderMap::new(); + map.insert(ACCESS_CONTROL_REQUEST_METHOD, HeaderValue::from_static("X-REINDEX")); + let raw = map.get(ACCESS_CONTROL_REQUEST_METHOD).expect("inserted value"); + let view = AccessControlRequestMethod::view(&map) + .expect("valid request method") + .expect("present request method"); + assert_eq!(view.method().as_str(), "X-REINDEX"); + assert_eq!(view.method().as_bytes().as_ptr(), raw.as_bytes().as_ptr()); + assert_eq!(view.method().try_to_method().expect("validated method").as_str(), "X-REINDEX"); + + for invalid in ["", "GET, POST", "bad method", "\"PATCH\""] { + AccessControlRequestMethodOwned::try_from(invalid).expect_err("request method must be one token"); + } + } + + #[test] + fn method_lists_preserve_lines_extensions_empty_members_and_duplicates() { + let mut map = HeaderMap::new(); + map.append(ACCESS_CONTROL_ALLOW_METHODS, HeaderValue::from_static("GET, X-PURGE,,")); + map.append(ACCESS_CONTROL_ALLOW_METHODS, HeaderValue::from_static("PATCH, GET")); + let view = AccessControlAllowMethods::view(&map) + .expect("valid method list") + .expect("present method list"); + let methods: Vec<&str> = view.iter().map(MethodView::as_str).collect(); + assert_eq!(methods, ["GET", "X-PURGE", "PATCH", "GET"]); + assert_eq!(view.field_values().count(), 2); + + let owned = AccessControlAllowMethods::owned(&map) + .expect("valid method list") + .expect("present method list"); + let mut encoded = HeaderMap::new(); + AccessControlAllowMethods::insert(&mut encoded, owned).expect("an empty header map has capacity"); + assert_eq!( + encoded + .get_all(ACCESS_CONTROL_ALLOW_METHODS) + .iter() + .map(HeaderValue::as_bytes) + .collect::>(), + [b"GET, X-PURGE,,".as_slice(), b"PATCH, GET".as_slice(),] + ); + } + + #[test] + fn max_age_handles_boundaries_overflow_and_duplicate_lines() { + let maximum = AccessControlMaxAgeOwned::try_from(u64::MAX.to_string()).expect("u64 maximum is valid"); + assert_eq!(maximum.seconds(), u64::MAX); + assert_eq!(maximum.duration(), Duration::from_secs(u64::MAX)); + + for invalid in ["", "-1", "+1", "1.0", "18446744073709551616"] { + let error = AccessControlMaxAgeOwned::try_from(invalid).expect_err("invalid delta-seconds must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber); + } + + let padded = AccessControlMaxAgeOwned::try_from(" 00060 ").expect("OWS and leading zeroes are valid"); + assert_eq!(padded.seconds(), 60); + assert_eq!(padded, AccessControlMaxAgeOwned::new(60)); + assert_eq!(padded.into_field_value(), "60"); + } + + #[test] + fn allow_origin_enforces_serialized_origin_grammar() { + for invalid in [ + "", + "NULL", + "HTTPS://example.com", + "https://Example.com", + "https://example.com/path", + "https://user@example.com", + "https://example.com:123456", + "https://example.com:65536", + "https://example.com:00443", + "https://example.com:443", + "https://01.2.3.4", + "https://256.2.3.4", + "https://[2001:0db8::1]", + "https://[2001:db8:0:0:0:0::1]", + "https://[2001:db8:0:0:0:0:0:1]", + "https://[2001:db8::192.0.2.1]", + ] { + assert!( + AccessControlAllowOriginOwned::try_from(invalid).is_err(), + "{invalid:?} must be rejected" + ); + } + for valid in [ + "http://127.0.0.1", + "https://example.com", + "wss://a-b.0:8080", + "https://[2001:db8::1]", + "https://[2001::1:0:0:1:1]", + ] { + assert!( + AccessControlAllowOriginOwned::from_origin(valid).is_ok(), + "{valid:?} must be accepted" + ); + } + } + + #[test] + fn response_lists_allow_empty_but_request_headers_require_a_member() { + let mut map = HeaderMap::new(); + map.insert(ACCESS_CONTROL_ALLOW_METHODS, HeaderValue::from_static("")); + assert!( + AccessControlAllowMethods::view(&map) + .expect("empty #method is valid") + .expect("present header") + .is_empty() + ); + + map.clear(); + map.insert(ACCESS_CONTROL_ALLOW_HEADERS, HeaderValue::from_static(",,,")); + assert!( + AccessControlAllowHeaders::view(&map) + .expect("empty #field-name is valid") + .expect("present header") + .is_empty() + ); + + map.clear(); + map.insert(ACCESS_CONTROL_EXPOSE_HEADERS, HeaderValue::from_static("")); + assert!( + AccessControlExposeHeaders::view(&map) + .expect("empty #field-name is valid") + .expect("present header") + .is_empty() + ); + + map.clear(); + map.insert(ACCESS_CONTROL_REQUEST_HEADERS, HeaderValue::from_static(",,,")); + let error = AccessControlRequestHeaders::view(&map).expect_err("1#field-name requires a member"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + } + + #[test] + fn owned_list_round_trip_retains_all_field_lines() { + let original = vec![FieldValue::from_static("X-A, x-a"), FieldValue::from_static("X-B,, X-C")]; + let header = AccessControlExposeHeadersOwned::from_field_values(original.clone()).expect("valid field lines"); + assert_eq!(header.clone().into_field_values(), original); + + let mut map = HeaderMap::new(); + AccessControlExposeHeaders::insert(&mut map, header).expect("empty map has insertion capacity"); + let round_trip: Vec = map.get_all(ACCESS_CONTROL_EXPOSE_HEADERS).iter().map(FieldValue::from).collect(); + assert_eq!(round_trip, original); + } + + #[test] + fn credentials_are_case_sensitive_and_policy_independent() { + assert_eq!( + AccessControlAllowCredentialsOwned::try_from(" \ttrue\t ") + .expect("OWS-framed true is valid") + .into_field_value(), + "true" + ); + for invalid in ["", "True", "TRUE", "false", "true, true"] { + AccessControlAllowCredentialsOwned::try_from(invalid).expect_err("credentials value must be exactly true"); + } + + let origin = AccessControlAllowOriginOwned::wildcard(); + let credentials = AccessControlAllowCredentialsOwned::allow(); + assert!(origin.is_wildcard()); + assert_eq!(credentials.into_field_value(), "true"); + } + + #[test] + fn allow_origin_borrows_and_classifies_values() { + let mut map = HeaderMap::new(); + map.insert(ACCESS_CONTROL_ALLOW_ORIGIN, HeaderValue::from_static("https://api.example:8443")); + let raw = map.get(ACCESS_CONTROL_ALLOW_ORIGIN).expect("inserted value"); + let view = AccessControlAllowOrigin::view(&map).expect("valid origin").expect("present origin"); + let origin = view.origin().expect("serialized origin"); + assert_eq!(origin, "https://api.example:8443"); + assert_eq!(origin.as_ptr(), raw.as_bytes().as_ptr()); + assert!(!view.is_wildcard()); + assert!(!view.is_null()); + + assert!(AccessControlAllowOriginOwned::wildcard().is_wildcard()); + assert!(AccessControlAllowOriginOwned::null().is_null()); + AccessControlAllowOriginOwned::from_origin("custom+v1://example").expect_err("unsupported schemes have opaque origins"); + } +} + +mod negotiation { + use super::*; + + #[test] + fn encoding_and_language_accept_unknown_values() { + let encodings = AcceptEncodingOwned::try_from("gzip, br;q=0.8, x-custom").expect("valid content codings"); + assert_eq!(encodings.items().count(), 3); + + let languages = AcceptLanguageOwned::try_from("en-US, fr;q=0.7, x-private").expect("valid language ranges"); + assert_eq!(languages.items().count(), 3); + let error = AcceptLanguageOwned::try_from("toolongprimarytag").expect_err("primary subtag is limited to eight letters"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidToken); + let error = AcceptEncodingOwned::try_from("gzip;level=1").expect_err("only a quality parameter is allowed"); + assert_eq!(error.header().as_str(), "accept-encoding"); + } + + #[test] + fn server_is_opaque_nonempty_and_singleton() { + let server = ServerOwned::try_from("example/1.0 (edge)").expect("valid opaque server value"); + assert_eq!(server.as_bytes(), b"example/1.0 (edge)"); + let error = ServerOwned::try_from(" \t").expect_err("blank Server value must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + + let mut map = HeaderMap::new(); + map.append("server", HeaderValue::from_static("one")); + map.append("server", HeaderValue::from_static("two")); + let error = Server::view(&map).expect_err("Server is a singleton"); + assert_eq!(error.kind(), DecodeErrorKind::UnexpectedMultipleValues); + } + + #[test] + fn host_rejects_userinfo_bad_ports_and_duplicates() { + let error = HostOwned::try_from("user@example.com").expect_err("userinfo is forbidden"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + let error = HostOwned::try_from("example.com:http").expect_err("port must be decimal"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber); + + let mut map = HeaderMap::new(); + map.append("host", HeaderValue::from_static("example.com")); + map.append("host", HeaderValue::from_static("example.org")); + let error = Host::view(&map).expect_err("Host is a singleton"); + assert_eq!(error.kind(), DecodeErrorKind::UnexpectedMultipleValues); + } + + #[test] + fn host_supports_names_ports_and_ip_literals() { + let host = HostOwned::with_port("example.com", 8443).expect("valid host and port"); + assert_eq!(host.host(), Ok("example.com")); + assert_eq!(host.port(), Ok(Some("8443"))); + + let ipv6 = HostOwned::try_from("[2001:db8::1]:443").expect("valid IPv6 host"); + assert_eq!(ipv6.host(), Ok("[2001:db8::1]")); + assert_eq!(ipv6.port(), Ok(Some("443"))); + + let future = HostOwned::try_from("[v1.fe80]:80").expect("valid IPvFuture host"); + assert_eq!(future.host(), Ok("[v1.fe80]")); + + // The owned form rebuilds its offsets from the borrowed slices. + let mut map = HeaderMap::new(); + map.insert(http::header::HOST, HeaderValue::from_static("example.com:8443")); + let owned = Host::owned(&map).expect("valid authority").expect("host present"); + assert_eq!(owned.host(), Ok("example.com")); + assert_eq!(owned.port(), Ok(Some("8443"))); + + map.insert(http::header::HOST, HeaderValue::from_static("[2001:db8::1]")); + let bare = Host::owned(&map).expect("valid authority").expect("host present"); + assert_eq!(bare.host(), Ok("[2001:db8::1]")); + assert_eq!(bare.port(), Ok(None)); + + // An authority may end at its colon, which is the case that separates + // a derived port offset from a stored one. + for (wire, port) in [("example.com:", Some("")), ("example.com", None), ("[2001:db8::1]:", Some(""))] { + map.insert(http::header::HOST, HeaderValue::from_static(wire)); + let owned = Host::owned(&map).expect("valid authority").expect("host present"); + assert_eq!(owned.port(), Ok(port), "{wire}"); + assert_eq!( + owned.port(), + Ok(Host::view(&map).expect("valid authority").expect("host present").port()), + "{wire}" + ); + } + } + + #[test] + fn accept_preserves_quoted_commas_and_extensions() { + let mut map = HeaderMap::new(); + map.append( + "accept", + HeaderValue::from_static("text/html;level=1, application/json;q=0.9;profile=\"a,b\""), + ); + map.append("accept", HeaderValue::from_static("*/*;q=0.1")); + let view = Accept::view(&map).expect("valid Accept").expect("Accept present"); + let items: Vec<_> = view.items().collect(); + assert_eq!(items.len(), 3); + assert_eq!(items[1], b"application/json;q=0.9;profile=\"a,b\""); + + let owned = Accept::owned(&map).expect("valid Accept").expect("Accept present"); + let mut encoded = HeaderMap::new(); + Accept::insert(&mut encoded, owned).expect("an empty map has capacity"); + assert_eq!(encoded.get_all("accept").iter().count(), 2); + } + + #[test] + fn allow_and_vary_preserve_raw_tokens() { + let allow = AllowOwned::try_from("GET, PATCH, X-CUSTOM").expect("valid method list"); + assert_eq!( + allow.items().collect::>(), + vec![b"GET".as_slice(), b"PATCH".as_slice(), b"X-CUSTOM".as_slice()] + ); + + let empty = AllowOwned::try_from("").expect("an empty Allow value is valid"); + assert_eq!(empty.items().count(), 0); + + let vary = VaryOwned::try_from("accept-encoding, x-tenant, *").expect("valid field-name list"); + assert_eq!(vary.items().count(), 3); + + let error = VaryOwned::try_from("\"unterminated").expect_err("quoted strings are not field names"); + assert_eq!(error.header().as_str(), "vary"); + } + + #[test] + fn accept_rejects_invalid_wildcards_and_quality() { + for value in ["*/json", "text/html;q=1.1", "text/html;q =0.5", "text/html;q= 0.5", "text"] { + let error = AcceptOwned::try_from(value).expect_err("invalid Accept grammar"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax, "{value}"); + } + } + + #[test] + fn relaxed_mode_accepts_only_documented_quality_deviations() { + fn assert_relaxed(name: &'static HeaderName, value: &'static str) { + let mut map = HeaderMap::new(); + map.insert(name, HeaderValue::from_static(value)); + assert!(H::view(&map).err().is_some(), "{name}: {value}"); + assert!( + H::owned_with(&map, DecodeMode::Relaxed).expect("relaxed quality syntax").is_some(), + "{name}: {value}" + ); + } + + assert_relaxed::(&ACCEPT, "text/html; q = .12345"); + assert_relaxed::(&ACCEPT_ENCODING, "gzip; q=.5"); + assert_relaxed::(&ACCEPT_LANGUAGE, "en-US; Q = 1.0000"); + } + + #[test] + fn relaxed_mode_accepts_documented_interoperability_deviations() { + fn assert_relaxed(name: &'static HeaderName, value: HeaderValue) { + let mut map = HeaderMap::new(); + map.insert(name, value); + assert!(H::view(&map).err().is_some(), "{name}"); + assert!( + H::view_with(&map, DecodeMode::Relaxed) + .expect("relaxed borrowed decoding succeeds") + .is_some() + ); + assert!( + H::owned_with(&map, DecodeMode::Relaxed) + .expect("relaxed owned decoding succeeds") + .is_some() + ); + } + + assert_relaxed::(&ETAG, HeaderValue::from_static("w/\"revision\"")); + assert_relaxed::(&IF_MATCH, HeaderValue::from_static("w/\"revision\"")); + assert_relaxed::(&CONTENT_TYPE, HeaderValue::from_static("text / html")); + assert_relaxed::(&RANGE, HeaderValue::from_static("bytes = 0 - 499")); + assert_relaxed::(&CONTENT_RANGE, HeaderValue::from_static("bytes 0 - 499 / 1234")); + assert_relaxed::(&LAST_MODIFIED, HeaderValue::from_static("Sun, 6 Nov 1994 8:49:37 UTC")); + assert_relaxed::(&HOST, HeaderValue::from_bytes(b"caf\xc3\xa9.example").expect("valid field bytes")); + assert_relaxed::(&LOCATION, HeaderValue::from_static("/a\\b\\c")); + } + + #[test] + fn relaxed_mode_preserves_semantics_and_original_wire_values() { + let mut map = HeaderMap::new(); + map.insert(ETAG, HeaderValue::from_static("w/\"revision\"")); + let etag = ETag::owned_with(&map, DecodeMode::Relaxed) + .expect("relaxed ETag decoding succeeds") + .expect("ETag is present"); + assert!(etag.is_weak()); + assert_eq!(etag.opaque_tag().expect("opaque tag metadata is valid"), b"revision"); + + map.clear(); + map.insert(IF_MATCH, HeaderValue::from_static("w/\"revision\"")); + let if_match = IfMatch::view_with(&map, DecodeMode::Relaxed) + .expect("relaxed If-Match decoding succeeds") + .expect("If-Match is present"); + let tag = if_match.tags().next().expect("one conditional tag"); + assert!(tag.is_weak()); + assert_eq!(tag.opaque_tag(), b"revision"); + + map.clear(); + map.insert(RANGE, HeaderValue::from_static("bytes = 0 - 499")); + let range = Range::view_with(&map, DecodeMode::Relaxed) + .expect("relaxed Range decoding succeeds") + .expect("Range is present"); + assert_eq!( + range.byte_ranges().expect("byte ranges are available").collect::>(), + vec![ByteRangeSpec::FromTo { first: 0, last: 499 }] + ); + + map.clear(); + map.insert(CONTENT_RANGE, HeaderValue::from_static("bytes 0 - 499 / 1234")); + let content_range = ContentRange::view_with(&map, DecodeMode::Relaxed) + .expect("relaxed Content-Range decoding succeeds") + .expect("Content-Range is present"); + assert_eq!( + content_range.byte_range(), + Some(ByteContentRange::Satisfied { + first: 0, + last: 499, + complete_length: Some(1234), + }) + ); + + map.clear(); + map.insert(LAST_MODIFIED, HeaderValue::from_static("Sun, 6 Nov 1994 8:49:37 UTC")); + let modified = LastModified::view_with(&map, DecodeMode::Relaxed) + .expect("relaxed Last-Modified decoding succeeds") + .expect("Last-Modified is present"); + assert_eq!(modified.date(), UNIX_EPOCH + Duration::from_secs(784_111_777)); + + map.clear(); + map.insert( + HOST, + HeaderValue::from_bytes(b"caf\xc3\xa9.example:443").expect("valid field bytes"), + ); + let host = Host::view_with(&map, DecodeMode::Relaxed) + .expect("relaxed Host decoding succeeds") + .expect("Host is present"); + assert_eq!(host.host().as_bytes(), b"caf\xc3\xa9.example"); + assert_eq!(host.port(), Some("443")); + drop(host); + + map.clear(); + map.insert(LOCATION, HeaderValue::from_static("/a\\b\\c")); + let location = Location::view_with(&map, DecodeMode::Relaxed) + .expect("relaxed Location decoding succeeds") + .expect("Location is present"); + assert_eq!(location.as_str().expect("Location remains valid UTF-8"), "/a\\b\\c"); + } + + #[test] + fn relaxed_mode_preserves_structural_validation() { + for value in ["*/json", "text/html;q=1.01", "text/html;q=-0.5"] { + let mut map = HeaderMap::new(); + map.insert(ACCEPT, HeaderValue::from_static(value)); + Accept::view_with(&map, DecodeMode::Relaxed).expect_err("relaxed mode must preserve structural validation"); + } + + let mut map = HeaderMap::new(); + map.insert(ACCEPT_LANGUAGE, HeaderValue::from_static("en_US;q=.5")); + AcceptLanguage::view_with(&map, DecodeMode::Relaxed).expect_err("underscores are not language-range separators"); + + map.insert(HOST, HeaderValue::from_static("user@example.com")); + Host::view_with(&map, DecodeMode::Relaxed).expect_err("relaxed host decoding must still reject user-info"); + + map.insert(ETAG, HeaderValue::from_static("\"space inside\"")); + ETag::view_with(&map, DecodeMode::Relaxed).expect_err("relaxed ETag decoding must retain the etagc grammar"); + + map.insert(RANGE, HeaderValue::from_static("bytes = +1 - 2")); + Range::view_with(&map, DecodeMode::Relaxed).expect_err("relaxed ranges must retain digit-only positions"); + + map.insert(CONTENT_RANGE, HeaderValue::from_static("bytes 0 - 1 / 2")); + ContentRange::view_with(&map, DecodeMode::Relaxed).expect_err("relaxed Content-Range keeps its unit separator strict"); + + map.insert( + LAST_MODIFIED, + HeaderValue::from_bytes(b"Sun,\t6 Nov 1994 8:49:37 UTC").expect("horizontal tabs are legal field bytes"), + ); + LastModified::view_with(&map, DecodeMode::Relaxed).expect_err("relaxed dates keep internal separators strict"); + } + + #[test] + fn source_exposes_relaxed_decoding() { + let mut map = HeaderMap::new(); + map.insert(ACCEPT_ENCODING, HeaderValue::from_static("br; q = .75")); + + AcceptEncoding::view(&map).expect_err("ordinary facade decoding remains strict"); + assert!( + AcceptEncoding::view_with(&map, DecodeMode::Relaxed) + .expect("relaxed quality syntax") + .is_some() + ); + } +} + +mod range { + use super::*; + + #[test] + fn parses_rfc_byte_range_examples() { + let range = RangeOwned::try_from("bytes=0-499, 500-999, -500, 9500-").expect("valid byte ranges"); + let specs: Vec<_> = range.byte_ranges().expect("byte unit").collect(); + assert_eq!( + specs, + vec![ + ByteRangeSpec::FromTo { first: 0, last: 499 }, + ByteRangeSpec::FromTo { first: 500, last: 999 }, + ByteRangeSpec::Suffix { length: 500 }, + ByteRangeSpec::From { first: 9500 }, + ] + ); + assert_eq!(range.as_field_value(), "bytes=0-499, 500-999, -500, 9500-"); + } + + #[test] + fn accept_ranges_supports_lists_extensions_and_none() { + let mut map = HeaderMap::new(); + map.append("accept-ranges", HeaderValue::from_static("bytes")); + map.append("accept-ranges", HeaderValue::from_static("example")); + let view = AcceptRanges::view(&map).expect("valid units").expect("field present"); + assert_eq!(view.units().collect::>(), vec!["bytes", "example"]); + assert!(!view.is_none()); + let owned = AcceptRanges::owned(&map).expect("valid units").expect("field present"); + assert_eq!(owned.units().collect::>(), vec!["bytes", "example"]); + let mut encoded = HeaderMap::new(); + AcceptRanges::insert(&mut encoded, owned).expect("an empty header map has capacity"); + assert_eq!(encoded[ACCEPT_RANGES], "bytes, example"); + + assert!(AcceptRangesOwned::none().is_none()); + let error = AcceptRangesOwned::try_from("none, bytes").expect_err("none is exclusive"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + } + + #[test] + fn borrowed_and_owned_range_views_agree() { + let mut map = HeaderMap::new(); + map.insert("range", HeaderValue::from_static("bytes=0-99, -10")); + let view = Range::view(&map).expect("valid range").expect("field present"); + let borrowed: Vec<_> = view.byte_ranges().expect("byte unit").collect(); + let owned = Range::owned(&map).expect("valid range").expect("field present"); + let owned_specs: Vec<_> = owned.byte_ranges().expect("byte unit").collect(); + assert_eq!(borrowed, owned_specs); + assert_eq!(owned.as_field_value(), "bytes=0-99, -10"); + } + + #[test] + fn range_constructor_round_trips_and_preserves_extensions() { + let range = RangeOwned::bytes([ + ByteRangeSpec::from_range(0..=99).expect("ordered"), + ByteRangeSpec::starting_at(200), + ByteRangeSpec::suffix(50), + ]) + .expect("valid set"); + assert_eq!(range.as_field_value(), "bytes=0-99, 200-, -50"); + + let extension = RangeOwned::extension("example-unit", "opaque=payload").expect("valid extension range"); + assert_eq!(extension.unit(), Ok("example-unit")); + assert!(extension.byte_ranges().is_none()); + assert_eq!(extension.extension_range_set(), Some(b"opaque=payload".as_slice())); + assert_eq!(extension.as_field_value(), "example-unit=opaque=payload"); + } + + #[test] + fn rejects_malformed_inverted_overflowing_and_duplicate_ranges() { + for (wire, kind) in [ + ("bytes=", DecodeErrorKind::MissingValue), + ("bytes=10-9", DecodeErrorKind::InvalidSyntax), + ("bytes=1-2-3", DecodeErrorKind::InvalidSyntax), + ("bytes=-", DecodeErrorKind::InvalidNumber), + ("bytes=18446744073709551616-", DecodeErrorKind::InvalidNumber), + ("bytes =0-1", DecodeErrorKind::InvalidToken), + ] { + let error = RangeOwned::try_from(wire).expect_err("invalid byte range"); + assert_eq!(error.kind(), kind, "{wire}"); + } + let error = RangeOwned::try_from("bytes=18446744073709551616-").expect_err("overflow must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber); + + let mut map = HeaderMap::new(); + map.append("range", HeaderValue::from_static("bytes=0-1")); + map.append("range", HeaderValue::from_static("bytes=2-3")); + let error = Range::view(&map).expect_err("Range is a singleton field"); + assert_eq!(error.kind(), DecodeErrorKind::UnexpectedMultipleValues); + } + + #[test] + fn parses_satisfied_unknown_and_unsatisfied_content_ranges() { + let satisfied = ContentRangeOwned::try_from("bytes 0-499/1234").expect("satisfied"); + assert_eq!( + satisfied.byte_range(), + Some(ByteContentRange::Satisfied { + first: 0, + last: 499, + complete_length: Some(1234), + }) + ); + + let unknown = ContentRangeOwned::try_from("bytes 0-499/*").expect("unknown length"); + assert_eq!( + unknown.byte_range(), + Some(ByteContentRange::Satisfied { + first: 0, + last: 499, + complete_length: None, + }) + ); + + let unsatisfied = ContentRangeOwned::try_from("bytes */1234").expect("unsatisfied"); + assert_eq!( + unsatisfied.byte_range(), + Some(ByteContentRange::Unsatisfied { complete_length: 1234 }) + ); + assert_eq!( + ContentRangeOwned::unsatisfied_bytes(1234) + .expect("valid unsatisfied range") + .as_field_value(), + "bytes */1234" + ); + + let extension = ContentRangeOwned::extension("example", "opaque response").expect("valid extension"); + assert_eq!(extension.unit(), Ok("example")); + assert_eq!(extension.byte_range(), None); + assert_eq!(extension.extension_payload(), Some(b"opaque response".as_slice())); + assert_eq!(extension.as_field_value(), "example opaque response"); + } + + #[test] + fn rejects_invalid_and_overflowing_content_ranges() { + for (wire, kind) in [ + ("bytes 500-499/1234", DecodeErrorKind::InvalidSyntax), + ("bytes 0-499/499", DecodeErrorKind::InvalidSyntax), + ("bytes */*", DecodeErrorKind::InvalidNumber), + ("bytes 0-1/18446744073709551616", DecodeErrorKind::InvalidNumber), + ("bytes 0-1-2/3", DecodeErrorKind::InvalidSyntax), + ] { + let error = ContentRangeOwned::try_from(wire).expect_err("invalid content range"); + assert_eq!(error.kind(), kind, "{wire}"); + } + let error = ContentRangeOwned::try_from("bytes */18446744073709551616").expect_err("overflow must fail"); + assert_eq!(error.kind(), DecodeErrorKind::InvalidNumber); + } +} + +mod security { + use super::*; + + #[test] + fn referrer_policy_parses_fallback_lists_and_selects_last() { + let mut map = HeaderMap::new(); + map.append("referrer-policy", HeaderValue::from_static("no-referrer-when-downgrade, origin")); + map.append("referrer-policy", HeaderValue::from_static("strict-origin-when-cross-origin")); + let view = ReferrerPolicy::view(&map).expect("valid policy list").expect("policy present"); + let policies: Result, _> = view.policies().collect(); + assert_eq!( + policies, + Ok(vec![ + ReferrerPolicyValue::NoReferrerWhenDowngrade, + ReferrerPolicyValue::Origin, + ReferrerPolicyValue::StrictOriginWhenCrossOrigin, + ]) + ); + assert_eq!(view.preferred(), Ok(ReferrerPolicyValue::StrictOriginWhenCrossOrigin)); + assert_eq!( + ReferrerPolicy::owned(&map) + .expect("valid policy list") + .expect("policy present") + .preferred(), + Ok(ReferrerPolicyValue::StrictOriginWhenCrossOrigin) + ); + } + + #[test] + fn hsts_is_a_singleton_header() { + let mut map = HeaderMap::new(); + map.append("strict-transport-security", HeaderValue::from_static("max-age=60")); + map.append("strict-transport-security", HeaderValue::from_static("max-age=120")); + StrictTransportSecurity::view(&map).expect_err("HSTS is a singleton"); + } + + #[test] + fn nosniff_is_exact_and_singleton() { + assert_eq!(XContentTypeOptionsOwned::nosniff().as_field_value().as_bytes(), b"nosniff"); + XContentTypeOptionsOwned::try_from("nosniff").expect("nosniff is valid"); + XContentTypeOptionsOwned::try_from("NoSniff").expect_err("value is case-sensitive"); + XContentTypeOptionsOwned::try_from("nosniff ").expect_err("trailing space must fail"); + let mut map = HeaderMap::new(); + map.append("x-content-type-options", HeaderValue::from_static("nosniff")); + map.append("x-content-type-options", HeaderValue::from_static("nosniff")); + XContentTypeOptions::view(&map).expect_err("X-Content-Type-Options is a singleton"); + } + + #[test] + fn hsts_preserves_wire_and_enforces_known_directives() { + let wire = "MAX-AGE=60 ; includeSubDomains ; x-vendor=\"a;b\""; + let hsts = StrictTransportSecurityOwned::try_from(wire).expect("valid HSTS"); + assert_eq!(hsts.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!(hsts.max_age(), Duration::from_mins(1)); + assert!(hsts.include_subdomains()); + assert!(!hsts.preload()); + StrictTransportSecurityOwned::try_from("includeSubDomains").expect_err("max-age is required"); + StrictTransportSecurityOwned::try_from("max-age=1; max-age=2").expect_err("duplicate max-age must fail"); + let quoted = StrictTransportSecurityOwned::try_from("max-age=\"10\"").expect("quoted decimal is valid"); + assert_eq!(quoted.max_age(), Duration::from_secs(10)); + let escaped = StrictTransportSecurityOwned::try_from("max-age=\"\\10\"").expect("quoted pairs unescape"); + assert_eq!(escaped.max_age(), Duration::from_secs(10)); + assert_eq!( + StrictTransportSecurityOwned::try_from("max-age=\"ten\"") + .expect_err("unescaped value must contain only digits") + .kind(), + http_headers::DecodeErrorKind::InvalidNumber + ); + StrictTransportSecurityOwned::try_from("max-age =10").expect_err("whitespace before equals must fail"); + StrictTransportSecurityOwned::try_from("max-age= 10").expect_err("whitespace after equals must fail"); + StrictTransportSecurityOwned::try_from("max-age=10; preload=yes").expect_err("preload does not take a value"); + StrictTransportSecurityOwned::try_from("max-age=18446744073709551616").expect_err("overflow must fail"); + StrictTransportSecurityOwned::try_from("max-age=10; includeSubDomains; includeSubDomains") + .expect_err("duplicate includeSubDomains must fail"); + } + + #[test] + fn csp_is_opaque_safe_policy_text_and_preserves_multiple_lines() { + let csp = ContentSecurityPolicyOwned::new("default-src 'self'") + .and_then(|policy| policy.with_policy("frame-ancestors 'none'")) + .expect("valid policies"); + let policies: Vec<_> = csp.policies().collect(); + assert_eq!( + policies, + vec![b"default-src 'self'".as_slice(), b"frame-ancestors 'none'".as_slice()] + ); + ContentSecurityPolicyOwned::new("default-src 'self'\nscript-src *").expect_err("newlines are forbidden"); + } + + #[test] + fn csp_exposes_non_utf8_policy_bytes_fallibly() { + let policy = + ContentSecurityPolicyOwned::from_bytes([b'd', b'e', b'f', b'a', b'u', b'l', b't', 0xff]).expect("obs-text is field-value safe"); + assert_eq!(policy.policies().next(), Some(&b"default\xff"[..])); + assert!(policy.policy_strs().next().is_some_and(|policy| policy.is_err())); + } + + #[test] + fn hsts_builder_and_accessors_cover_common_directives() { + let hsts = StrictTransportSecurityOwned::builder(Duration::from_hours(8760)) + .include_subdomains() + .preload() + .extension_value("x-rollout", "\"stable\"") + .build() + .expect("valid HSTS construction"); + assert_eq!(hsts.max_age(), Duration::from_hours(8760)); + assert!(hsts.include_subdomains()); + assert!(hsts.preload()); + assert_eq!( + hsts.as_field_value().as_bytes(), + b"max-age=31536000; includeSubDomains; preload; x-rollout=\"stable\"" + ); + let directives: Result, _> = hsts.directives().collect(); + let directives = directives.expect("valid stored directives"); + assert_eq!(directives.len(), 4); + assert_eq!(directives[3].name(), "x-rollout"); + assert_eq!(directives[3].value(), Some(b"\"stable\"".as_slice())); + } + + #[test] + fn referrer_policy_preserves_unknown_extension_tokens() { + let mut map = HeaderMap::new(); + map.insert("referrer-policy", HeaderValue::from_static("strict-origin, future-policy")); + let view = ReferrerPolicy::view(&map).expect("valid token list").expect("policy present"); + let tokens: Result, _> = view.tokens().collect(); + let tokens = tokens.expect("valid policy tokens"); + assert_eq!(tokens[0].policy(), Some(ReferrerPolicyValue::StrictOrigin)); + assert_eq!(tokens[1].as_str(), "future-policy"); + assert_eq!(tokens[1].policy(), None); + assert_eq!(view.preferred(), Ok(ReferrerPolicyValue::StrictOrigin)); + map.insert("referrer-policy", HeaderValue::from_static("not a token")); + ReferrerPolicy::view(&map).expect_err("invalid policy tokens must fail"); + } +} + +mod websocket { + use super::*; + + #[test] + fn validates_accept_canonical_base64_and_length() { + let digest = [ + 0xb3, 0x7a, 0x4f, 0x2c, 0xc0, 0x62, 0x4f, 0x16, 0x90, 0xf6, 0x46, 0x06, 0xcf, 0x38, 0x59, 0x45, 0xb2, 0xbe, 0xc4, 0xea, + ]; + let accept = SecWebSocketAcceptOwned::from_digest(digest).expect("a digest has a valid encoding"); + assert_eq!(accept.encoded(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + SecWebSocketAcceptOwned::try_from("s3pPLMBiTxaQ9kYGzzhZRbK+xOo=").expect("canonical accept value"); + SecWebSocketAcceptOwned::try_from("s3pPLMBiTxaQ9kYGzzhZRbK+xOp=").expect_err("noncanonical tail bits must fail"); + } + + #[test] + fn version_lists_cross_field_lines_and_reject_noncanonical_values() { + let mut map = HeaderMap::new(); + map.append("sec-websocket-version", HeaderValue::from_static("13, 8")); + map.append("sec-websocket-version", HeaderValue::from_static("7")); + let view = SecWebSocketVersion::view(&map) + .expect("valid version list") + .expect("version present"); + assert_eq!(view.versions().collect::>(), vec![7, 8, 13]); + view.requested().expect_err("multiple versions are not one request version"); + map.insert("sec-websocket-version", HeaderValue::from_static("013")); + SecWebSocketVersion::view(&map).expect_err("leading zero is noncanonical"); + map.insert("sec-websocket-version", HeaderValue::from_static("256")); + SecWebSocketVersion::view(&map).expect_err("version exceeds u8"); + assert_eq!( + SecWebSocketVersionOwned::try_from("13").and_then(|version| version.requested()), + Ok(13) + ); + } + + #[test] + fn extension_builder_emits_quoted_token_values() { + let extensions = SecWebSocketExtensionsOwned::builder() + .extension("permessage-deflate") + .parameter_flag("client_no_context_takeover") + .quoted_parameter("mode", "fast") + .extension("x-test") + .build() + .expect("valid extension construction"); + let wire = extensions.extensions().next().expect("first extension").expect("valid extension"); + let mode = wire.parameters().nth(1).expect("mode parameter").expect("valid mode parameter"); + assert_eq!(mode.value(), Some(b"\"fast\"".as_slice())); + } + + #[test] + fn protocol_lists_are_tokens_and_case_sensitive() { + let protocols = SecWebSocketProtocolOwned::new("chat") + .and_then(|protocols| protocols.with_protocol("superchat")) + .expect("valid protocol tokens"); + let collected: Result, _> = protocols.protocols().collect(); + assert_eq!(collected, Ok(vec!["chat", "superchat"])); + protocols.selected().expect_err("multiple protocols are not one selected protocol"); + assert_eq!( + SecWebSocketProtocolOwned::try_from("chat").and_then(|protocol| { protocol.selected().map(std::string::ToString::to_string) }), + Ok(String::from("chat")) + ); + SecWebSocketProtocolOwned::new("not a token").expect_err("protocol must be a token"); + } + + #[test] + fn derives_accept_from_client_key() { + let key = SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ==").expect("valid key"); + let accept = SecWebSocketAcceptOwned::from_key(&key).expect("valid derived response"); + assert_eq!(accept.encoded(), b"s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); + } + + #[test] + fn key_and_accept_are_singleton_headers() { + let mut map = HeaderMap::new(); + map.append("sec-websocket-key", HeaderValue::from_static("dGhlIHNhbXBsZSBub25jZQ==")); + map.append("sec-websocket-key", HeaderValue::from_static("dGhlIHNhbXBsZSBub25jZQ==")); + SecWebSocketKey::view(&map).expect_err("WebSocket key is a singleton"); + } + + #[test] + fn validates_key_canonical_base64_and_length() { + let key = SecWebSocketKeyOwned::from_nonce(*b"the sample nonce").expect("a fixed nonce has a valid encoding"); + assert_eq!(key.encoded(), b"dGhlIHNhbXBsZSBub25jZQ=="); + SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ==").expect("canonical key"); + SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZR==").expect_err("noncanonical tail bits must fail"); + SecWebSocketKeyOwned::try_from("dGhlIHNhbXBsZSBub25jZQ=").expect_err("incorrect padding must fail"); + } + + #[test] + fn list_headers_preserve_field_lines_when_owned() { + let mut map = HeaderMap::new(); + map.append( + "sec-websocket-extensions", + HeaderValue::from_static("permessage-deflate; client_max_window_bits"), + ); + map.append("sec-websocket-extensions", HeaderValue::from_static("x-vendor; mode=\"fast\"")); + let owned = SecWebSocketExtensions::owned(&map) + .expect("valid extensions") + .expect("extensions present"); + let mut output = HeaderMap::new(); + SecWebSocketExtensions::insert(&mut output, owned).expect("an empty header map has capacity"); + let encoded: Vec<_> = output + .get_all("sec-websocket-extensions") + .iter() + .map(|value| value.as_bytes().to_vec()) + .collect(); + assert_eq!( + encoded, + vec![ + b"permessage-deflate; client_max_window_bits".to_vec(), + b"x-vendor; mode=\"fast\"".to_vec(), + ] + ); + } + + #[test] + fn rejects_malformed_extension_parameters() { + SecWebSocketExtensionsOwned::try_from("permessage-deflate; mode=\"open").expect_err("unterminated quote must fail"); + SecWebSocketExtensionsOwned::try_from("permessage-deflate; mode=\"not token\"").expect_err("decoded quoted value must be a token"); + SecWebSocketExtensionsOwned::try_from("permessage-deflate; =value").expect_err("parameter name is required"); + SecWebSocketExtensionsOwned::try_from("permessage-deflate; mode=").expect_err("parameter value is required"); + SecWebSocketExtensionsOwned::try_from("permessage-deflate; mode =fast").expect_err("whitespace before equals must fail"); + SecWebSocketExtensionsOwned::try_from("permessage-deflate; mode= fast").expect_err("whitespace after equals must fail"); + } + + #[test] + fn parses_quoted_extension_parameters_and_lists() { + let mut map = HeaderMap::new(); + map.insert( + "sec-websocket-extensions", + HeaderValue::from_static("permessage-deflate; mode=\"fa\\st\"; client_max_window_bits, x-test"), + ); + let view = SecWebSocketExtensions::view(&map) + .expect("valid extension list") + .expect("extensions present"); + let mut extensions = view.extensions(); + let first = extensions.next().expect("first extension").expect("valid first extension"); + assert_eq!(first.name(), "permessage-deflate"); + let parameters: Result, _> = first.parameters().collect(); + let parameters = parameters.expect("valid parameters"); + assert_eq!(parameters.len(), 2); + assert_eq!(parameters[0].name(), "mode"); + assert_eq!(parameters[0].value(), Some(b"\"fa\\st\"".as_slice())); + assert!(parameters[0].is_quoted()); + assert_eq!(parameters[1].name(), "client_max_window_bits"); + assert_eq!(parameters[1].value(), None); + assert_eq!( + extensions.next().expect("second extension").expect("valid second extension").name(), + "x-test" + ); + } + + #[test] + fn borrowed_views_preserve_original_wire() { + let mut map = HeaderMap::new(); + map.insert("sec-websocket-key", HeaderValue::from_static("dGhlIHNhbXBsZSBub25jZQ==")); + let view = SecWebSocketKey::view(&map).expect("valid key").expect("key present"); + assert_eq!(view.as_field_value().as_bytes(), view.encoded()); + assert_eq!( + SecWebSocketKey::owned(&map).expect("valid key").expect("key present").encoded(), + b"dGhlIHNhbXBsZSBub25jZQ==" + ); + } +} + +mod reusable_credentials { + use super::*; + + #[test] + fn debug_does_not_expose_credentials() { + let authorization = AuthorizationOwned::::basic(b"user", b"secret").expect("valid credentials"); + let mut map = HeaderMap::new(); + Authorization::::insert(&mut map, authorization).expect("an empty header map has capacity"); + let authorization = Authorization::::view(&map) + .expect("valid basic authorization") + .expect("authorization present"); + let mut credentials = BasicCredentials::new(); + authorization.extract(&mut credentials).expect("valid basic credentials"); + let debug = format!("{credentials:?}"); + assert!(!debug.contains("secret")); + assert!(debug.contains("sensitive")); + } +} + +mod security_additional { + use super::*; + + #[test] + fn referrer_policy_fast_path_agrees_with_the_general_parser() { + let inputs = [ + "no-referrer", + "no-referrer-when-downgrade", + "origin", + "origin-when-cross-origin", + "same-origin", + "strict-origin", + "strict-origin-when-cross-origin", + "unsafe-url", + " strict-origin ", + "strict-origin, future-policy", + "future-policy", + "no-referrer,", + ",no-referrer", + "", + "not a token", + "NO-REFERRER", + ]; + for input in inputs { + let mut map = HeaderMap::new(); + map.insert("referrer-policy", HeaderValue::from_static(input)); + let borrowed = ReferrerPolicy::view(&map).map(|view| { + view.map(|view| { + view.tokens() + .map(|token| token.map(|token| token.as_str().to_owned())) + .collect::>() + }) + }); + let owned = ReferrerPolicyOwned::try_from(input).map(|policy| { + policy + .tokens() + .map(|token| token.map(|token| token.as_str().to_owned())) + .collect::>() + }); + assert_eq!(borrowed.is_ok(), owned.is_ok(), "acceptance disagreed for {input:?}"); + if let (Ok(Some(borrowed)), Ok(owned)) = (borrowed, owned) { + assert_eq!(borrowed, owned, "tokens disagreed for {input:?}"); + } + } + } +} + +mod websocket_additional { + use super::*; + + #[test] + fn list_fast_paths_agree_with_general_delimiter_handling() { + let mut map = HeaderMap::new(); + map.append("sec-websocket-version", HeaderValue::from_static(" 13 ,, 8\t")); + let view = SecWebSocketVersion::view(&map) + .expect("optional whitespace and empty members are tolerated") + .expect("version present"); + assert_eq!(view.versions().collect::>(), vec![8, 13]); + + map.insert("sec-websocket-version", HeaderValue::from_static(",")); + let error = SecWebSocketVersion::view(&map).expect_err("a line of only delimiters holds no member"); + assert_eq!(error.kind(), DecodeErrorKind::MissingValue); + + let mut map = HeaderMap::new(); + map.append("sec-websocket-protocol", HeaderValue::from_static("chat")); + map.append("sec-websocket-protocol", HeaderValue::from_static("\"open")); + let error = SecWebSocketProtocol::view(&map).expect_err("an unterminated quote is rejected"); + assert_eq!(error.kind(), DecodeErrorKind::UnterminatedQuote); + assert_eq!(error.value_index(), Some(1)); + } +} diff --git a/crates/http_headers/tests/negotiation_semantics.rs b/crates/http_headers/tests/negotiation_semantics.rs new file mode 100644 index 000000000..2baa1bf0b --- /dev/null +++ b/crates/http_headers/tests/negotiation_semantics.rs @@ -0,0 +1,692 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Semantic negotiation coverage independent of Content-Type and HTTP adapters. + +#![cfg(feature = "headers-negotiation")] +#![expect(clippy::unwrap_used, reason = "test failures provide sufficient context")] +#![expect(clippy::assertions_on_result_states, reason = "invalid-input tables assert rejection")] + +use std::cmp::Ordering; +use std::collections::{BTreeSet, HashSet}; +use std::error::Error; +use std::hash::{DefaultHasher, Hash, Hasher}; +use std::iter; + +use http_headers::headers::{ + Accept, AcceptEncoding, AcceptEncodingEntry, AcceptEncodingOwned, AcceptEntry, AcceptLanguage, AcceptLanguageEntry, + AcceptLanguageOwned, AcceptOwned, ContentCoding, ContentCodingKind, LanguageRange, MediaRange, MediaRangeKind, NegotiationParameter, + NegotiationParameterValue, NegotiationToken, Quality, QualityView, +}; +use http_headers::sink::{EncodedValues, FieldSink, InsertError}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldSensitivity, FieldValue, FieldValueRef}; + +struct Source<'a> { + name: &'static FieldName, + values: Vec>, +} + +impl<'a> Source<'a> { + fn new(name: &'static FieldName, values: impl IntoIterator) -> Self { + Self { + name, + values: values.into_iter().map(FieldValueRef::new).collect(), + } + } +} + +impl FieldSource for Source<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + (self.name == name).then(|| FieldLines::from_borrowed(name, &self.values)).flatten() + } +} + +struct SingleSource<'a> { + name: &'static FieldName, + bytes: &'a [u8], +} + +impl FieldSource for SingleSource<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::single(name, self.bytes)) + } +} + +#[derive(Default)] +struct Stored(Vec); + +impl FieldSource for Stored { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_slice(name, &self.0) + } +} + +impl FieldSink for Stored { + fn set_values(&mut self, _: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.0 = values.into_iter().collect(); + Ok(()) + } + + fn append_values(&mut self, _name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.0.extend(values); + Ok(()) + } + + fn remove_values(&mut self, _name: &'static FieldName) { + self.0.clear(); + } +} + +fn hash(value: impl Hash) -> u64 { + let mut hasher = DefaultHasher::new(); + value.hash(&mut hasher); + hasher.finish() +} + +fn quality(value: &[u8]) -> QualityView<'_> { + QualityView::parse(value, DecodeMode::Relaxed).unwrap() +} + +#[test] +fn quality_expected_values_and_exact_conversion() { + for (wire, thousandths, canonical) in [ + ("0", 0, "0"), + ("0.", 0, "0"), + ("0.000", 0, "0"), + ("0.001", 1, "0.001"), + ("0.010", 10, "0.01"), + ("0.100", 100, "0.1"), + ("0.501", 501, "0.501"), + ("1", 1000, "1"), + ("1.", 1000, "1"), + ("1.000", 1000, "1"), + ] { + let parsed = QualityView::parse(wire.as_bytes(), DecodeMode::Strict).unwrap(); + let compact = Quality::from_thousandths(thousandths).unwrap(); + assert_eq!(parsed.to_quality().unwrap(), compact); + assert_eq!(Quality::try_from(parsed).unwrap().thousandths(), thousandths); + assert_eq!(parsed.to_string(), canonical); + assert_eq!(compact.to_string(), canonical); + assert_eq!(parsed.is_zero(), thousandths == 0); + assert_eq!(parsed.is_one(), thousandths == 1000); + } + assert_eq!(quality(b".5").to_quality().unwrap().thousandths(), 500); + assert_eq!(quality(b" \t.50000\t ").to_quality().unwrap().thousandths(), 500); + assert!(quality(b"0.0001").to_quality().is_err()); + assert_eq!(quality(b"0.0001").to_string(), "0.0001"); + assert!(Quality::from_thousandths(1001).is_err()); + assert!(Quality::from_thousandths(u16::MAX).is_err()); +} + +#[test] +fn quality_errors_have_distinct_messages_and_no_sources() { + let invalid = Quality::from_thousandths(1001).unwrap_err(); + assert_eq!(invalid.to_string(), "quality must be an accepted decimal between zero and one"); + assert!(invalid.source().is_none()); + let inexact = quality(b"0.0001").to_quality().unwrap_err(); + assert_eq!(inexact.to_string(), "quality is not an exact number of thousandths"); + assert!(inexact.source().is_none()); +} + +#[test] +fn typed_encoding_entries_serialize_zero_one_and_fractional_weights() { + for (wire, expected) in [ + (b"0".as_slice(), b"gzip;q=0".as_slice()), + (b"1", b"gzip;q=1"), + (b"0.001", b"gzip;q=0.001"), + (b"0.010", b"gzip;q=0.01"), + (b"0.100", b"gzip;q=0.1"), + (b".0001000", b"gzip;q=0.0001"), + ] { + let entry = AcceptEncodingEntry::new(ContentCoding::parse("gzip").unwrap(), Some(quality(wire))).unwrap(); + let header = AcceptEncodingOwned::from_entries([entry]).unwrap(); + assert_eq!(header.values().next().unwrap().as_bytes(), expected); + assert_eq!(header.entries().next().unwrap(), entry); + } +} + +#[test] +fn decimal_equality_ordering_and_hashing_are_exact() { + let equal = [ + quality(b"0.5"), + quality(b".50"), + quality(b"0.5000"), + Quality::from_thousandths(500).unwrap().into(), + ]; + assert_eq!(equal.into_iter().collect::>().len(), 1); + assert_eq!(equal.into_iter().collect::>().len(), 1); + assert!(equal.windows(2).all(|pair| hash(pair[0]) == hash(pair[1]))); + assert_eq!(quality(b".0001000"), quality(b"0.0001")); + assert_eq!(hash(quality(b".0001000")), hash(quality(b"0.0001"))); + assert_eq!(quality(b"1.00000000"), QualityView::ONE); + assert_eq!(quality(b".00000000"), QualityView::ZERO); + let ascending = [ + quality(b"0"), + quality(b".00000000000000000000001"), + quality(b".0001"), + quality(b".00010000001"), + quality(b".001"), + quality(b".00100000001"), + quality(b"0.49999999999999999999999999999999999"), + quality(b".5"), + quality(b"0.50000000000000000000000000000000001"), + quality(b"0.99999999999999999999999999999999999"), + quality(b"1"), + ]; + for pair in ascending.windows(2) { + assert!(pair[0] < pair[1]); + assert!(pair[1] > pair[0]); + } + let mut long = b"0.".to_vec(); + long.extend(iter::repeat_n(b'0', 4096)); + long.push(b'1'); + let mut larger = long.clone(); + *larger.last_mut().unwrap() = b'2'; + assert!(quality(&long) > QualityView::ZERO); + assert!(quality(&long) < quality(&larger)); + assert!(quality(&long).to_quality().is_err()); + assert_eq!(quality(&long).to_string().as_bytes(), long); +} + +#[test] +fn quality_rejection_preserves_strict_and_relaxed_grammars() { + for invalid in [ + "", + ".", + "00", + "01", + "2", + "-0", + "+0.5", + "NaN", + "0.1e0", + "1.001", + "1.0000001", + "\"0.5\"", + "0..5", + ".5x", + ] { + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert!(QualityView::parse(invalid.as_bytes(), mode).is_err(), "{invalid:?}"); + } + } + for relaxed_only in [".5", "0.0001", "0.5000", "1.0000", " 0.5 "] { + assert!(QualityView::parse(relaxed_only.as_bytes(), DecodeMode::Strict).is_err()); + assert!(QualityView::parse(relaxed_only.as_bytes(), DecodeMode::Relaxed).is_ok()); + } +} + +#[test] +fn media_ranges_and_codings_have_case_insensitive_semantics() { + for (wire, kind) in [ + ("*/*", MediaRangeKind::Any), + ("TEXT/*", MediaRangeKind::TypeWildcard), + ("Text/HTML", MediaRangeKind::Exact), + ("te*xt/ht*ml", MediaRangeKind::Exact), + ("**/*", MediaRangeKind::TypeWildcard), + ] { + let range = MediaRange::parse(wire).unwrap(); + assert_eq!(range.kind(), kind); + assert_eq!(range.to_string(), wire); + } + assert_eq!(MediaRange::parse("Text/HTML").unwrap(), MediaRange::parse("text/html").unwrap()); + assert_eq!( + hash(MediaRange::parse("Text/HTML").unwrap()), + hash(MediaRange::parse("text/html").unwrap()) + ); + for invalid in ["*/html", "text", "text/html/extra", "text/ html", "/html", "text/"] { + assert!(MediaRange::parse(invalid).is_err()); + } + assert!(MediaRange::new(NegotiationToken::new("*").unwrap(), NegotiationToken::new("html").unwrap()).is_err()); + for (wire, kind) in [ + ("*", ContentCodingKind::Wildcard), + ("IDENTITY", ContentCodingKind::Identity), + ("GZip", ContentCodingKind::Gzip), + ("compress", ContentCodingKind::Compress), + ("deflate", ContentCodingKind::Deflate), + ("br", ContentCodingKind::Br), + ("zstd", ContentCodingKind::Zstd), + ("dcb", ContentCodingKind::Dcb), + ("dcz", ContentCodingKind::Dcz), + ("x-private", ContentCodingKind::Extension), + ("g*zip", ContentCodingKind::Extension), + ] { + let coding = ContentCoding::parse(wire).unwrap(); + assert_eq!(coding.kind(), kind); + assert_eq!(coding.as_str(), wire); + assert_eq!(coding.to_string(), wire); + assert!(coding.token().eq_ignore_ascii_case(wire)); + } + assert_eq!(ContentCoding::parse("GZIP").unwrap(), ContentCoding::parse("gzip").unwrap()); + assert_eq!( + hash(ContentCoding::parse("GZIP").unwrap()), + hash(ContentCoding::parse("gzip").unwrap()) + ); + assert!(ContentCoding::parse("g zip").is_err()); + assert!(NegotiationToken::new("").is_err()); + assert!(NegotiationToken::new("ümlaut").is_err()); +} + +#[test] +fn negotiation_token_order_is_ascii_case_insensitive_and_prefix_aware() { + for (left, right, expected) in [ + ("GZip", "gzip", Ordering::Equal), + ("BR", "gzip", Ordering::Less), + ("gzip", "BR", Ordering::Greater), + ("A", "aa", Ordering::Less), + ("aa", "A", Ordering::Greater), + ("gZiP", "GZIQ", Ordering::Less), + ("Z", "a", Ordering::Greater), + ] { + let left = NegotiationToken::new(left).unwrap(); + let right = NegotiationToken::new(right).unwrap(); + assert_eq!(left.cmp(&right), expected); + assert_eq!(left.partial_cmp(&right), Some(expected)); + assert_eq!(right.cmp(&left), expected.reverse()); + assert_eq!(left == right, expected == Ordering::Equal); + } +} + +#[test] +fn basic_language_ranges_retain_primary_and_subtags() { + for (wire, primary, rest) in [ + ("*", None, vec![]), + ("e", Some("e"), vec![]), + ("abcdefgh", Some("abcdefgh"), vec![]), + ("ZH-Hans-CN-12345678", Some("ZH"), vec!["Hans", "CN", "12345678"]), + ("abcdefgh-1-a", Some("abcdefgh"), vec!["1", "a"]), + ] { + let range = LanguageRange::parse(wire).unwrap(); + assert_eq!(range.as_str(), wire); + assert_eq!(range.to_string(), wire); + assert_eq!(range.primary().map(NegotiationToken::as_str), primary); + assert_eq!(range.subtags().map(NegotiationToken::as_str).collect::>(), rest); + assert_eq!(range.is_wildcard(), primary.is_none()); + } + for invalid in [ + "", + "abcdefghi", + "1en", + "en-123456789", + "en-", + "-en", + "en--US", + "en_US", + "en-*", + "*-en", + "é", + ] { + assert!(LanguageRange::parse(invalid).is_err(), "{invalid}"); + } + let mixed = LanguageRange::parse("eN-uS").unwrap(); + assert!(mixed.eq_ignore_ascii_case("EN-US")); + assert_eq!(mixed, LanguageRange::parse("en-US").unwrap()); + assert_eq!(hash(mixed), hash(LanguageRange::parse("EN-us").unwrap())); +} + +fn check_accept_entries<'a>(mut entries: impl Iterator>) { + let first = entries.next().unwrap(); + assert_eq!(first.range().type_().as_str(), "TEXT"); + assert_eq!(first.range().subtype().as_str(), "HTML"); + assert_eq!(first.range().kind(), MediaRangeKind::Exact); + assert_eq!(first.quality().to_quality().unwrap().thousandths(), 500); + assert_eq!(first.explicit_quality(), Some(quality(b".5"))); + let parameters = first.parameters().collect::>(); + assert_eq!(parameters.len(), 2); + assert_eq!(parameters[0].name(), NegotiationToken::new("level").unwrap()); + let value = parameters[0].value().unwrap(); + assert!(value.is_quoted()); + assert_eq!(value.raw_bytes(), b"\"a,b;c\\\"\\\\\xff\""); + assert_eq!(value.decoded_bytes().collect::>(), b"a,b;c\"\\\xff"); + assert_eq!(value.to_decoded_bytes(), b"a,b;c\"\\\xff"); + assert_eq!(parameters[1].value().unwrap().raw_bytes(), b"two"); + assert!(!parameters[1].value().unwrap().is_quoted()); + let extensions = first.extensions().collect::>(); + assert_eq!(extensions.len(), 2); + assert_eq!(extensions[0].name().as_str(), "Flag"); + assert_eq!(extensions[0].value(), None); + assert_eq!(extensions[1].name().as_str(), "empty"); + assert_eq!(extensions[1].value().unwrap().to_decoded_bytes(), b""); + assert!(extensions[1].value().unwrap().is_quoted()); + let wildcard = entries.next().unwrap(); + assert_eq!(wildcard.range().kind(), MediaRangeKind::TypeWildcard); + assert_eq!(wildcard.quality(), QualityView::ZERO); + assert_eq!(wildcard.parameters().count(), 0); + assert_eq!(wildcard.extensions().count(), 0); + let any = entries.next().unwrap(); + assert_eq!(any.range().kind(), MediaRangeKind::Any); + assert_eq!(any.quality(), QualityView::ONE); + assert_eq!(any.explicit_quality(), None); + let repeated = entries.next().unwrap(); + assert_eq!(repeated.range().type_().as_str(), "text"); + assert_eq!(repeated.quality(), QualityView::ONE); + assert_eq!(repeated.explicit_quality(), Some(QualityView::ONE)); + assert_eq!(repeated.parameters().count(), 0); + assert!(entries.next().is_none()); +} + +#[test] +fn accept_parameters_extensions_quoting_order_and_wire_are_independent() { + let lines = [ + b" , TEXT/HTML;Level=\"a,b;c\\\"\\\\\xff\";level=two;q=0.5;Flag;empty=\"\", text/*;q=0.000 , */*, ".as_slice(), + b"text/html;q=1", + ]; + let source = Source::new(&FieldName::Accept, lines); + let view = Accept::view(&source).unwrap().unwrap(); + let owned = Accept::owned(&source).unwrap().unwrap(); + for _ in 0..8 { + check_accept_entries(view.entries()); + check_accept_entries(owned.entries()); + } + assert_eq!(view.values().map(FieldValueRef::as_bytes).collect::>(), lines); + assert_eq!(owned.values().map(FieldValueRef::as_bytes).collect::>(), lines); + assert_eq!(view.items().collect::>(), owned.items().collect::>()); + let mut sink = Stored::default(); + Accept::insert(&mut sink, owned).unwrap(); + assert_eq!(sink.0.iter().map(FieldValue::as_bytes).collect::>(), lines); + check_accept_entries(Accept::view(&sink).unwrap().unwrap().entries()); +} + +#[test] +fn encoding_and_language_entries_preserve_duplicates_default_and_explicit_quality() { + let source = Source::new( + &FieldName::AcceptEncoding, + [b" , GZIP, gzip;q=1,, identity;q=0, *;q=0.5".as_slice(), b"x-private"], + ); + let view = AcceptEncoding::view(&source).unwrap().unwrap(); + let owned = AcceptEncoding::owned(&source).unwrap().unwrap(); + let entries = view.entries().collect::>(); + assert_eq!(entries, owned.entries().collect::>()); + assert_eq!( + entries.iter().map(|entry| entry.coding().as_str()).collect::>(), + ["GZIP", "gzip", "identity", "*", "x-private"] + ); + assert_eq!(entries[0].quality(), QualityView::ONE); + assert_eq!(entries[0].explicit_quality(), None); + assert_eq!(entries[1].explicit_quality(), Some(QualityView::ONE)); + assert_ne!(entries[0], entries[1]); + assert!(entries[2].quality().is_zero()); + assert_eq!(entries[2].coding().kind(), ContentCodingKind::Identity); + assert_eq!(entries[3].coding().kind(), ContentCodingKind::Wildcard); + assert_eq!(entries[4].coding().kind(), ContentCodingKind::Extension); + let source = Source::new( + &FieldName::AcceptLanguage, + [b"en-US, EN-us;q=1, *;q=0".as_slice(), b"zh-Hans-CN;q=0.125"], + ); + let view = AcceptLanguage::view(&source).unwrap().unwrap(); + let owned = AcceptLanguage::owned(&source).unwrap().unwrap(); + let entries = view.entries().collect::>(); + assert_eq!(entries, owned.entries().collect::>()); + assert_eq!(entries[0].range(), entries[1].range()); + assert_eq!(entries[0].explicit_quality(), None); + assert_eq!(entries[1].explicit_quality(), Some(QualityView::ONE)); + assert!(entries[2].range().is_wildcard()); + assert!(entries[2].quality().is_zero()); + assert_eq!(entries[3].quality().to_quality().unwrap().thousandths(), 125); +} + +#[test] +fn relaxed_precision_is_shared_across_all_three_headers() { + let source = Source::new(&FieldName::Accept, [b"text/plain;p=x; Q \t= .0001000 ;flag".as_slice()]); + assert!(Accept::view(&source).is_err()); + let view = Accept::view_with(&source, DecodeMode::Relaxed).unwrap().unwrap(); + let owned = Accept::owned_with(&source, DecodeMode::Relaxed).unwrap().unwrap(); + assert_eq!(view.entries().next().unwrap().quality(), quality(b".0001")); + assert_eq!(owned.entries().next().unwrap().quality(), quality(b".0001")); + assert_eq!(view.entries().next().unwrap().parameters().next().unwrap().name().as_str(), "p"); + assert_eq!(view.entries().next().unwrap().extensions().next().unwrap().name().as_str(), "flag"); + let source = Source::new(&FieldName::AcceptEncoding, [b"br; Q \t= .0001000".as_slice()]); + assert!(AcceptEncoding::view(&source).is_err()); + let view = AcceptEncoding::view_with(&source, DecodeMode::Relaxed).unwrap().unwrap(); + let owned = AcceptEncoding::owned_with(&source, DecodeMode::Relaxed).unwrap().unwrap(); + assert_eq!(view.entries().next().unwrap().quality(), quality(b".0001")); + assert_eq!(view.entries().collect::>(), owned.entries().collect::>()); + let source = Source::new(&FieldName::AcceptLanguage, [b"en; Q \t= .0001000".as_slice()]); + assert!(AcceptLanguage::view(&source).is_err()); + let view = AcceptLanguage::view_with(&source, DecodeMode::Relaxed).unwrap().unwrap(); + let owned = AcceptLanguage::owned_with(&source, DecodeMode::Relaxed).unwrap().unwrap(); + assert_eq!(view.entries().next().unwrap().quality(), quality(b".0001")); + assert_eq!(view.entries().collect::>(), owned.entries().collect::>()); +} + +#[test] +fn typed_constructors_validate_cross_component_rules_and_retain_exact_values() { + let range = MediaRange::parse("Text/HTML").unwrap(); + let parameter = NegotiationParameter::new( + NegotiationToken::new("Level").unwrap(), + Some(NegotiationParameterValue::parse(b"\"a\\\"b,\xff\"").unwrap()), + ); + let flag = NegotiationParameter::new(NegotiationToken::new("preview").unwrap(), None); + let reserved = NegotiationParameter::new( + NegotiationToken::new("Q").unwrap(), + Some(NegotiationParameterValue::parse(b"1").unwrap()), + ); + assert!(AcceptEntry::new(range, &[flag], None, &[]).is_err()); + assert!(AcceptEntry::new(range, &[reserved], None, &[]).is_err()); + assert!(AcceptEntry::new(range, &[], Some(QualityView::ONE), &[reserved]).is_err()); + assert!(AcceptEntry::new(range, &[], None, &[flag]).is_err()); + let parameters = [parameter, parameter]; + let extensions = [flag]; + let entry = AcceptEntry::new(range, ¶meters, Some(quality(b".0001000")), &extensions).unwrap(); + let owned = AcceptOwned::from_entries([entry]).unwrap(); + assert_eq!( + owned.values().next().unwrap().as_bytes(), + b"Text/HTML;Level=\"a\\\"b,\xff\";Level=\"a\\\"b,\xff\";q=0.0001;preview" + ); + let decoded = owned.entries().next().unwrap(); + assert_eq!(decoded.range(), range); + assert_eq!(decoded.quality(), quality(b".0001")); + assert_eq!(decoded.parameters().collect::>(), parameters); + assert_eq!(decoded.extensions().collect::>(), extensions); + + let entry = AcceptEncodingEntry::new(ContentCoding::parse("GZIP").unwrap(), Some(quality(b".0001000"))).unwrap(); + let header = AcceptEncodingOwned::from_entries([entry, entry]).unwrap(); + assert_eq!(header.values().next().unwrap().as_bytes(), b"GZIP;q=0.0001, GZIP;q=0.0001"); + assert_eq!(header.entries().collect::>(), [entry, entry]); + let entry = AcceptLanguageEntry::new( + LanguageRange::parse("en-US").unwrap(), + Some(Quality::from_thousandths(900).unwrap().into()), + ) + .unwrap(); + let header = AcceptLanguageOwned::from_entries([entry]).unwrap(); + assert_eq!(header.values().next().unwrap().as_bytes(), b"en-US;q=0.9"); + assert_eq!(header.entries().next().unwrap(), entry); + for invalid in [b"".as_slice(), b"a b", b"\"unclosed", b"\"line\n\"", b"\"trailing\\\"", b"\"a\"b\""] { + assert!(NegotiationParameterValue::parse(invalid).is_err(), "{invalid:?}"); + } + assert_eq!(NegotiationParameterValue::parse(b"\"\\\xff\"").unwrap().to_decoded_bytes(), [0xff]); +} + +#[test] +fn absence_empty_lists_and_constructor_budgets_are_distinct() { + for name in [&FieldName::Accept, &FieldName::AcceptEncoding, &FieldName::AcceptLanguage] { + let absent = Source::new(name, []); + assert!(Accept::view(&absent).unwrap().is_none()); + assert!(AcceptEncoding::view(&absent).unwrap().is_none()); + assert!(AcceptLanguage::view(&absent).unwrap().is_none()); + } + assert_eq!(AcceptOwned::from_entries([]).unwrap().entries().count(), 0); + assert_eq!(AcceptEncodingOwned::from_entries([]).unwrap().entries().count(), 0); + assert_eq!(AcceptLanguageOwned::from_entries([]).unwrap().entries().count(), 0); + assert_eq!(AcceptOwned::try_from(" , , ").unwrap().entries().count(), 0); + assert_eq!(AcceptEncodingOwned::try_from(" , , ").unwrap().entries().count(), 0); + assert_eq!(AcceptLanguageOwned::try_from(" , , ").unwrap().entries().count(), 0); + let entry = AcceptEncodingEntry::new(ContentCoding::parse("br").unwrap(), None).unwrap(); + assert_eq!( + AcceptEncodingOwned::from_entries(iter::repeat_n(entry, 1_024)) + .unwrap() + .entries() + .count(), + 1_024 + ); + assert_eq!( + AcceptEncodingOwned::from_entries(iter::repeat_n(entry, 1_025)).unwrap_err().kind(), + DecodeErrorKind::SourceLimitExceeded + ); + let maximum = "x".repeat(65_536); + let maximum_coding = ContentCoding::parse(&maximum).unwrap(); + let maximum_entry = AcceptEncodingEntry::new(maximum_coding, None).unwrap(); + assert_eq!( + AcceptEncodingOwned::from_entries([maximum_entry]) + .unwrap() + .values() + .next() + .unwrap() + .as_bytes() + .len(), + 65_536 + ); + assert_eq!( + AcceptEncodingEntry::new(maximum_coding, Some(QualityView::ONE)).unwrap_err().kind(), + DecodeErrorKind::SourceLimitExceeded + ); + assert_eq!( + AcceptEncodingOwned::from_entries([maximum_entry, entry]).unwrap_err().kind(), + DecodeErrorKind::SourceLimitExceeded + ); + assert!( + AcceptEntry::new( + MediaRange::new(NegotiationToken::new(&maximum).unwrap(), NegotiationToken::new("x").unwrap()).unwrap(), + &[], + None, + &[], + ) + .is_err() + ); +} + +fn rejects_in_all_positions(valid: &'static [u8], invalid: &'static [u8], kind: DecodeErrorKind) { + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + for position in 0..3 { + let mut lines = [valid; 3]; + lines[position] = invalid; + let source = Source::new(F::name(), lines); + let owned = F::owned_with(&source, mode).err().unwrap(); + let view = F::view_with(&source, mode).err().unwrap(); + assert_eq!(owned.kind(), kind); + assert_eq!(view.kind(), kind); + assert_eq!(owned.value_index(), None); + assert_eq!(view.value_index(), None); + let joined = lines.join(&b','); + let source = SingleSource { + name: F::name(), + bytes: &joined, + }; + let owned = F::owned_with(&source, mode).err().unwrap(); + let view = F::view_with(&source, mode).err().unwrap(); + assert_eq!(owned.kind(), kind); + assert_eq!(view.kind(), kind); + assert_eq!(owned.value_index(), None); + assert_eq!(view.value_index(), None); + } + } +} + +#[test] +fn streaming_projection_handles_many_members_and_late_quoted_parameters() { + let mut wire = iter::repeat_n("text/plain;p=x;q=0.125", 12).collect::>().join(", "); + wire.push_str(", application/json;p=\"a;b,c\";q=0.75;flag"); + let source = SingleSource { + name: &FieldName::Accept, + bytes: wire.as_bytes(), + }; + let view = Accept::view(&source).unwrap().unwrap(); + let owned = Accept::owned(&source).unwrap().unwrap(); + for entries in [view.entries().collect::>(), owned.entries().collect::>()] { + assert_eq!(entries.len(), 13); + for entry in &entries[..12] { + assert_eq!(entry.quality().to_quality().unwrap().thousandths(), 125); + assert_eq!(entry.parameters().next().unwrap().value().unwrap().raw_bytes(), b"x"); + } + let last = entries[12]; + assert_eq!(last.range().subtype().as_str(), "json"); + assert_eq!(last.quality().to_quality().unwrap().thousandths(), 750); + assert_eq!(last.parameters().next().unwrap().value().unwrap().to_decoded_bytes(), b"a;b,c"); + assert_eq!(last.extensions().next().unwrap().name().as_str(), "flag"); + } +} + +#[test] +fn invalid_members_and_source_budgets_never_become_partial_typed_lists() { + rejects_in_all_positions::(b"text/html", b"*/plain", DecodeErrorKind::InvalidSyntax); + rejects_in_all_positions::(b"text/html", b"text/html;q=0.5;Q=0.4", DecodeErrorKind::InvalidSyntax); + rejects_in_all_positions::(b"text/html", b"text/html;x=\"bad", DecodeErrorKind::UnterminatedQuote); + rejects_in_all_positions::(b"text/html", b"text/html;x", DecodeErrorKind::InvalidSyntax); + rejects_in_all_positions::(b"gzip", b"bad coding", DecodeErrorKind::InvalidToken); + rejects_in_all_positions::(b"gzip", b"gzip;q=1;x=y", DecodeErrorKind::InvalidSyntax); + rejects_in_all_positions::(b"en", b"en-*", DecodeErrorKind::InvalidToken); + rejects_in_all_positions::(b"en", b"en;q=1.0001", DecodeErrorKind::InvalidSyntax); + let entries = iter::repeat_n("en", 1_024).collect::>().join(","); + let source = Source::new(&FieldName::AcceptLanguage, [entries.as_bytes()]); + assert_eq!(AcceptLanguage::view(&source).unwrap().unwrap().entries().count(), 1_024); + let too_many = format!("{entries},en"); + let source = Source::new(&FieldName::AcceptLanguage, [too_many.as_bytes()]); + assert!(AcceptLanguage::view(&source).is_err()); + assert!(AcceptLanguage::owned(&source).is_err()); + let too_large = vec![b'a'; 65_537]; + let source = Source::new(&FieldName::AcceptEncoding, [too_large.as_slice()]); + assert!(AcceptEncoding::view(&source).is_err()); + assert!(AcceptEncoding::owned(&source).is_err()); + let source = Source::new(&FieldName::Accept, iter::repeat_n(b"*/*".as_slice(), 129)); + assert!(Accept::view(&source).is_err()); + assert!(Accept::owned(&source).is_err()); +} + +#[test] +fn raw_owned_slice_forwarding_retains_sensitivity_and_spelling() { + let source = Stored(vec![ + FieldValue::from_static("GZIP;q=0.500").with_sensitivity(FieldSensitivity::Sensitive), + FieldValue::from_static("identity;q=0"), + ]); + let view = AcceptEncoding::view(&source).unwrap().unwrap(); + assert_eq!(view.entries().next().unwrap().quality(), quality(b".5")); + assert!(view.values().next().unwrap().is_sensitive()); + let owned = AcceptEncoding::owned(&source).unwrap().unwrap(); + let mut sink = Stored::default(); + AcceptEncoding::insert(&mut sink, owned).unwrap(); + assert_eq!(sink.0, source.0); + assert!(sink.0[0].is_sensitive()); + assert!(!sink.0[1].is_sensitive()); +} + +#[cfg(feature = "http")] +#[test] +fn http_sources_have_the_same_semantics_and_preserve_wire_forwarding() { + use http::{HeaderMap, HeaderValue, header}; + + let mut map = HeaderMap::new(); + let mut first = HeaderValue::from_static("GZIP;q=0.500, *;q=0"); + first.set_sensitive(true); + map.append(header::ACCEPT_ENCODING, first); + map.append(header::ACCEPT_ENCODING, HeaderValue::from_static("identity;q=1")); + let view = AcceptEncoding::view(&map).unwrap().unwrap(); + let owned = AcceptEncoding::owned(&map).unwrap().unwrap(); + assert_eq!(view.entries().collect::>(), owned.entries().collect::>()); + assert_eq!(view.entries().next().unwrap().quality(), quality(b".5")); + let mut forwarded = HeaderMap::new(); + AcceptEncoding::insert(&mut forwarded, owned).unwrap(); + assert_eq!(map, forwarded); + assert!(forwarded[header::ACCEPT_ENCODING].is_sensitive()); + map.insert(header::ACCEPT, HeaderValue::from_static("text/plain;p=\"a,b\";q=0.125;flag")); + let accept = Accept::view(&map).unwrap().unwrap(); + let entry = accept.entries().next().unwrap(); + assert_eq!(entry.parameters().next().unwrap().value().unwrap().to_decoded_bytes(), b"a,b"); + assert_eq!(entry.extensions().next().unwrap().value(), None); + map.insert(header::ACCEPT_LANGUAGE, HeaderValue::from_static("fr-CH;q=0.75, *;q=0")); + let language = AcceptLanguage::view(&map).unwrap().unwrap(); + assert_eq!(language.entries().next().unwrap().range().primary().unwrap().as_str(), "fr"); + assert_eq!( + language.entries().next().unwrap().quality().to_quality().unwrap().thousandths(), + 750 + ); + let mut invalid = HeaderMap::new(); + invalid.append(header::ACCEPT, HeaderValue::from_static("text/plain")); + invalid.append(header::ACCEPT, HeaderValue::from_static("text/html;p=\"unterminated")); + let view_error = Accept::view(&invalid).unwrap_err(); + let owned_error = Accept::owned(&invalid).unwrap_err(); + assert_eq!(view_error.kind(), DecodeErrorKind::UnterminatedQuote); + assert_eq!(owned_error.kind(), DecodeErrorKind::UnterminatedQuote); + assert_eq!(view_error.value_index(), None); + assert_eq!(owned_error.value_index(), Some(1)); +} diff --git a/crates/http_headers/tests/negotiation_source_stability.rs b/crates/http_headers/tests/negotiation_source_stability.rs new file mode 100644 index 000000000..83afda9f4 --- /dev/null +++ b/crates/http_headers/tests/negotiation_source_stability.rs @@ -0,0 +1,181 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Retained negotiation views must not re-enter a stateful field source. + +#![cfg(feature = "headers-negotiation")] +#![expect(clippy::unwrap_used, reason = "test failures provide sufficient context")] + +use std::cell::Cell; +use std::iter; + +use http_headers::headers::{Accept, AcceptEncoding, AcceptLanguage, QualityView}; +use http_headers::source::{FieldLines, FieldSource, MAX_CUSTOM_FIELD_BYTES, MAX_CUSTOM_FIELD_LINES, MAX_CUSTOM_LIST_ITEMS}; +use http_headers::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldName, FieldSensitivity, FieldValueRef}; + +struct SwitchingSource { + name: &'static FieldName, + lines: [FieldValueRef<'static>; 2], + changed: Cell, + calls: Cell, +} + +impl FieldSource for SwitchingSource { + fn lines(&self, name: &'static FieldName) -> Option> { + assert_eq!(name, self.name); + self.calls.set(self.calls.get() + 1); + if self.changed.get() { + Some(FieldLines::single(name, b"\r")) + } else { + FieldLines::from_borrowed(name, &self.lines) + } + } +} + +struct BorrowedSource<'a>(&'a [FieldValueRef<'a>]); + +impl FieldSource for BorrowedSource<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_borrowed(name, self.0) + } +} + +#[test] +fn retained_views_do_not_reenter_stateful_custom_sources() { + macro_rules! retained { + ($header:ty, $first:expr, $strict:expr, $relaxed:expr) => { + for (mode, second, quality_wire) in [ + (DecodeMode::Strict, $strict.as_slice(), b"0.125".as_slice()), + ( + DecodeMode::Relaxed, + $relaxed.as_slice(), + b".12500000000000000000001".as_slice(), + ), + ] { + let source = SwitchingSource { + name: <$header>::name(), + lines: [ + FieldValueRef::new($first).with_sensitivity(FieldSensitivity::Sensitive), + FieldValueRef::new(second), + ], + changed: Cell::new(false), + calls: Cell::new(0), + }; + let checked = FieldLines::from_borrowed(source.name, &source.lines).unwrap(); + let expected_members = checked.comma_items().collect::, _>>().unwrap(); + let view = <$header>::view_with(&source, mode).unwrap().unwrap(); + assert_eq!(source.calls.get(), 1); + source.changed.set(true); + + let quality = QualityView::parse(quality_wire, mode).unwrap(); + let expected = [(QualityView::ONE, None), (QualityView::ONE, None), (quality, Some(quality))]; + for _ in 0..8 { + assert_eq!( + view.entries() + .map(|entry| (entry.quality(), entry.explicit_quality())) + .collect::>(), + expected + ); + assert_eq!(view.items().collect::>(), expected_members); + assert!(view.values().next().unwrap().is_sensitive()); + assert_eq!(source.calls.get(), 1); + } + let error = <$header>::view_with(&source, mode).unwrap_err(); + assert_eq!(error, DecodeError::new(source.name, DecodeErrorKind::InvalidSyntax)); + assert_eq!(source.calls.get(), 2); + assert_eq!(view.entries().count(), 3); + assert_eq!(source.calls.get(), 2); + } + }; + } + + retained!( + Accept, + b"text/plain, , text/plain,", + b"text/html;p=\"a,b;\\\"\xff\";q=0.125;flag", + b"text/html;p=\"a,b;\\\"\xff\"; q = .12500000000000000000001 ;flag" + ); + retained!( + AcceptEncoding, + b"GZIP, , GZIP,", + b"identity;q=0.125", + b"identity; q = .12500000000000000000001" + ); + retained!(AcceptLanguage, b"en-US, , en-US,", b"*;q=0.125", b"*; q = .12500000000000000000001"); +} + +fn assert_error(lines: &[FieldValueRef<'_>], expected: DecodeErrorKind) { + let source = BorrowedSource(lines); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let expected = DecodeError::new(F::name(), expected); + assert_eq!(F::view_with(&source, mode).err().unwrap(), expected); + assert_eq!(F::owned_with(&source, mode).err().unwrap(), expected); + } +} + +fn assert_precedence(good: &str, invalid_token: &[u8]) { + assert_error::(&[FieldValueRef::new(invalid_token)], DecodeErrorKind::InvalidToken); + assert_error::( + &[FieldValueRef::new(invalid_token), FieldValueRef::new(b"\r")], + DecodeErrorKind::InvalidSyntax, + ); + assert_error::( + &[FieldValueRef::new(b"\"unterminated"), FieldValueRef::new(b"\r")], + DecodeErrorKind::InvalidSyntax, + ); + let too_many_lines = vec![FieldValueRef::new(b"\"unterminated"); MAX_CUSTOM_FIELD_LINES + 1]; + assert_error::(&too_many_lines, DecodeErrorKind::SourceLimitExceeded); + let too_many_bytes = vec![b'a'; MAX_CUSTOM_FIELD_BYTES]; + assert_error::( + &[FieldValueRef::new(b"\"unterminated"), FieldValueRef::new(&too_many_bytes)], + DecodeErrorKind::SourceLimitExceeded, + ); + for (count, expected) in [ + (MAX_CUSTOM_LIST_ITEMS - 1, DecodeErrorKind::UnterminatedQuote), + (MAX_CUSTOM_LIST_ITEMS, DecodeErrorKind::SourceLimitExceeded), + ] { + let mut wire = iter::repeat_n(good, count).collect::>().join(","); + wire.push_str(",\"unterminated"); + assert_error::(&[FieldValueRef::new(wire.as_bytes())], expected); + } +} + +#[test] +fn source_and_member_error_precedence_is_unchanged() { + assert_precedence::("text/plain", b"text/pl ain"); + assert_precedence::("gzip", b"bad coding"); + assert_precedence::("en", b"bad_range"); +} + +#[cfg(feature = "http")] +#[test] +fn http_views_still_exempt_the_retained_adapter_from_all_custom_budgets() { + macro_rules! unlimited { + ($header:ty, $http_name:expr, $member:expr) => {{ + let line_count = MAX_CUSTOM_FIELD_LINES + 1; + let per_line = if cfg!(miri) { + // Joining removes the final comma; still exceed the real byte budget. + (MAX_CUSTOM_FIELD_BYTES / line_count + 2).div_ceil($member.len() + 1) + } else { + 128 + }; + let wire = iter::repeat_n($member, per_line).collect::>().join(","); + assert!(wire.len() * line_count > MAX_CUSTOM_FIELD_BYTES); + assert!(per_line * line_count > MAX_CUSTOM_LIST_ITEMS); + let mut map = http::HeaderMap::new(); + for _ in 0..line_count { + map.append($http_name, http::HeaderValue::from_str(&wire).unwrap()); + } + let view = <$header>::view(&map).unwrap().unwrap(); + for _ in 0..2 { + assert_eq!(view.entries().count(), per_line * line_count); + assert_eq!(view.items().count(), per_line * line_count); + } + let custom = vec![FieldValueRef::new(wire.as_bytes()); line_count]; + assert_error::<$header>(&custom, DecodeErrorKind::SourceLimitExceeded); + }}; + } + unlimited!(Accept, http::header::ACCEPT, "text/plain;q=1"); + unlimited!(AcceptEncoding, http::header::ACCEPT_ENCODING, "gzip;q=1"); + unlimited!(AcceptLanguage, http::header::ACCEPT_LANGUAGE, "en;q=1"); +} diff --git a/crates/http_headers/tests/owned_token_lists.rs b/crates/http_headers/tests/owned_token_lists.rs new file mode 100644 index 000000000..8705892a2 --- /dev/null +++ b/crates/http_headers/tests/owned_token_lists.rs @@ -0,0 +1,68 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Typed list construction across inline and shared field-value sizes. + +#![cfg(feature = "headers-negotiation")] + +use http_headers::headers::{AllowOwned, FieldNameView, MethodView, VaryEntryView, VaryOwned}; + +#[test] +fn allow_construction_preserves_methods_across_storage_boundaries() { + for length in [63, 64, 65, 129] { + let extension = "m".repeat(length - 10); + let methods = [MethodView::GET, MethodView::new(&extension).unwrap(), MethodView::GET]; + let expected = format!("GET, {extension}, GET"); + assert_eq!(expected.len(), length); + + let value = AllowOwned::from_methods(methods); + assert_eq!(value.values().len(), 1); + let line = value.values().next().unwrap(); + assert_eq!(line.as_bytes(), expected.as_bytes()); + assert!(!line.is_sensitive()); + assert_eq!(value.methods().collect::>(), methods); + + let retained = value.clone(); + drop(value); + assert_eq!(retained.values().next().unwrap().as_bytes(), expected.as_bytes()); + assert_eq!(retained.methods().collect::>(), methods); + } +} + +#[test] +fn vary_construction_preserves_wildcards_names_and_spelling_across_storage_boundaries() { + for length in [63, 64, 65, 129] { + let custom_name = format!("X-{}", "a".repeat(length - 21)); + let custom = FieldNameView::new(&custom_name).unwrap(); + let origin = FieldNameView::new("Origin").unwrap(); + let entries = [ + VaryEntryView::from_field_name(origin), + VaryEntryView::WILDCARD, + VaryEntryView::from_field_name(custom), + VaryEntryView::from_field_name(origin), + ]; + let expected = format!("Origin, *, {custom_name}, Origin"); + assert_eq!(expected.len(), length); + + let value = VaryOwned::from_entries(entries); + assert_eq!(value.values().len(), 1); + let line = value.values().next().unwrap(); + assert_eq!(line.as_bytes(), expected.as_bytes()); + assert!(!line.is_sensitive()); + assert!(value.contains_wildcard()); + assert_eq!(value.entries().collect::>(), entries); + + let retained = value.clone(); + drop(value); + assert_eq!(retained.values().next().unwrap().as_bytes(), expected.as_bytes()); + assert_eq!(retained.entries().collect::>(), entries); + assert!(retained.contains_wildcard()); + + let names = VaryOwned::from_field_names([custom, origin, custom]); + assert_eq!( + names.values().next().unwrap().as_bytes(), + format!("{custom_name}, Origin, {custom_name}").as_bytes() + ); + assert!(!names.contains_wildcard()); + } +} diff --git a/crates/http_headers/tests/range_invariants.rs b/crates/http_headers/tests/range_invariants.rs new file mode 100644 index 000000000..8ed8b170e --- /dev/null +++ b/crates/http_headers/tests/range_invariants.rs @@ -0,0 +1,331 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Field-byte validation and constructor/parser identity for range headers. + +#![cfg(feature = "headers-range")] +#![expect(clippy::unwrap_used, reason = "test fixture and round-trip failures provide sufficient context")] + +use std::str; + +use http_headers::headers::{ByteContentRange, ByteRangeSpec, CompleteLength, ContentRange, ContentRangeOwned, RangeOwned}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef, InvalidFieldValue, SingleValueField}; + +struct Source<'a> { + values: [FieldValueRef<'a>; 1], + borrowed: bool, +} + +impl FieldSource for Source<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + if self.borrowed { + FieldLines::from_borrowed(name, &self.values) + } else { + Some(FieldLines::single(name, self.values[0].as_bytes())) + } + } +} + +fn assert_forbidden_field_bytes_are_rejected(wire: &[u8]) { + assert_eq!(FieldValue::from_bytes(wire).unwrap_err(), InvalidFieldValue); + let strict = ::decode_view(FieldValueRef::new(wire)).unwrap_err(); + assert_eq!(strict.header(), &FieldName::ContentRange); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let direct = ::decode_view_with(FieldValueRef::new(wire), mode).unwrap_err(); + assert_eq!(direct.header(), &FieldName::ContentRange); + if mode == DecodeMode::Strict { + assert_eq!(direct, strict); + } else { + assert_eq!(direct.kind(), DecodeErrorKind::InvalidSyntax); + } + for borrowed in [false, true] { + let source = Source { + values: [FieldValueRef::new(wire)], + borrowed, + }; + let view = ::view_with(&source, mode).unwrap_err(); + let owned = ::owned_with(&source, mode).unwrap_err(); + assert_eq!(view.kind(), DecodeErrorKind::InvalidSyntax); + assert_eq!(view.header(), &FieldName::ContentRange); + assert_eq!(owned, view); + } + } + let wire = str::from_utf8(wire).unwrap(); + assert_eq!( + ContentRangeOwned::try_from(wire).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + ContentRangeOwned::try_from(wire.to_owned()).unwrap_err().kind(), + DecodeErrorKind::InvalidSyntax + ); +} + +fn assert_content_range_decoding(wire: &str, mode: DecodeMode, expected: ByteContentRange) { + let field = FieldValue::from_str(wire).unwrap(); + let direct = ::decode_view_with(field.as_field_value_ref(), mode).unwrap(); + assert_eq!(direct.byte_range(), Some(expected)); + assert_eq!(direct.extension_payload(), None); + assert_eq!(direct.as_field_value().as_bytes(), wire.as_bytes()); + let owned = ::decode_owned_with(field.clone(), mode).unwrap(); + assert_eq!(owned.byte_range(), Some(expected)); + assert_eq!(owned.extension_payload(), None); + assert_eq!(owned.as_field_value().as_bytes(), wire.as_bytes()); + if mode == DecodeMode::Strict { + let view = ::decode_view(field.as_field_value_ref()).unwrap(); + assert_eq!(view.byte_range(), direct.byte_range()); + assert_eq!(::decode_owned(field).unwrap(), owned); + } + for borrowed in [false, true] { + let source = Source { + values: [FieldValueRef::new(wire.as_bytes())], + borrowed, + }; + let view = ::view_with(&source, mode).unwrap().unwrap(); + assert_eq!(view.byte_range(), Some(expected)); + assert_eq!(view.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!(::owned_with(&source, mode).unwrap().unwrap(), owned); + if mode == DecodeMode::Strict { + assert_eq!( + ::view(&source).unwrap().unwrap().byte_range(), + Some(expected) + ); + assert_eq!(::owned(&source).unwrap().unwrap(), owned); + } + } +} + +#[test] +fn control_bytes_at_every_content_range_position_are_rejected_by_all_entry_points() { + for base in [b"bytes 0-1/2".as_slice(), b"bytes 0-1/*", b"bytes */2", b"items opaque"] { + for byte in (0_u8..=0x1f).filter(|byte| *byte != b'\t').chain([0x7f]) { + for position in 0..=base.len() { + let mut wire = base.to_vec(); + wire.insert(position, byte); + assert_forbidden_field_bytes_are_rejected(&wire); + } + } + } +} + +#[test] +fn relaxed_content_range_delimiters_accept_only_sp_and_htab_without_changing_wire() { + let known = ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: Some(2), + }; + let unknown = ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: None, + }; + let unsatisfied = ByteContentRange::Unsatisfied { complete_length: 2 }; + for (pattern, expected) in [ + ("bytes 0|-1/2", known), + ("bytes 0-|1/2", known), + ("bytes 0-1|/2", known), + ("bytes 0-1/|2", known), + ("bytes 0|-|1|/|2", known), + ("Bytes 0|-|1|/|*", unknown), + ("bytes *|/|2", unsatisfied), + ] { + for whitespace in [" ", "\t", " \t", "\t "] { + let wire = pattern.replace('|', whitespace); + assert_content_range_decoding(&wire, DecodeMode::Relaxed, expected); + let value = FieldValue::from_str(&wire).unwrap(); + ::decode_view(value.as_field_value_ref()).unwrap_err(); + ::decode_view_with(value.as_field_value_ref(), DecodeMode::Strict).unwrap_err(); + ::decode_owned(value.clone()).unwrap_err(); + ::decode_owned_with(value, DecodeMode::Strict).unwrap_err(); + for borrowed in [false, true] { + let source = Source { + values: [FieldValueRef::new(wire.as_bytes())], + borrowed, + }; + ::view(&source).unwrap_err(); + ::owned(&source).unwrap_err(); + } + } + } + for (wire, expected) in [("bytes 0-1/2", known), ("bytes 0-1/*", unknown), ("bytes */2", unsatisfied)] { + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_content_range_decoding(wire, mode, expected); + } + } +} + +#[test] +fn relaxed_content_ranges_still_reject_outer_or_missing_component_whitespace() { + for wire in [ + " bytes 0-1/2", + "bytes\t0-1/2", + "bytes 0-1/2", + "bytes \t0-1/2", + "bytes 0-1/2 ", + "bytes 0-1/2\t", + "bytes 0-1/* ", + "bytes 0-1/ \t", + "bytes -1 /2", + "bytes 0- \t/2", + ] { + let field = FieldValue::from_str(wire).unwrap(); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + ::decode_view_with(field.as_field_value_ref(), mode).unwrap_err(); + ::decode_owned_with(field.clone(), mode).unwrap_err(); + let source = Source { + values: [field.as_field_value_ref()], + borrowed: false, + }; + ::view_with(&source, mode).unwrap_err(); + ::owned_with(&source, mode).unwrap_err(); + } + } +} + +fn assert_range_round_trip(value: &RangeOwned) { + let reparsed = RangeOwned::try_from(value.clone().into_field_value()).unwrap(); + assert_eq!(&reparsed, value); + assert_eq!(reparsed.is_bytes(), value.is_bytes()); + assert_eq!(reparsed.extension_range_set(), value.extension_range_set()); + #[cfg(feature = "serde")] + { + let json = serde_json::to_string(value).unwrap(); + let decoded: RangeOwned = serde_json::from_str(&json).unwrap(); + assert_eq!(&decoded, value); + assert_eq!(decoded.is_bytes(), value.is_bytes()); + assert_eq!(decoded.as_field_value().as_bytes(), value.as_field_value().as_bytes()); + } +} + +fn assert_content_range_round_trip(value: &ContentRangeOwned) { + let reparsed = ContentRangeOwned::try_from(value.clone().into_field_value()).unwrap(); + assert_eq!(&reparsed, value); + assert_eq!(reparsed.byte_range(), value.byte_range()); + assert_eq!(reparsed.extension_payload(), value.extension_payload()); + #[cfg(feature = "serde")] + { + let json = serde_json::to_string(value).unwrap(); + let decoded: ContentRangeOwned = serde_json::from_str(&json).unwrap(); + assert_eq!(&decoded, value); + assert_eq!(decoded.byte_range(), value.byte_range()); + assert_eq!(decoded.as_field_value().as_bytes(), value.as_field_value().as_bytes()); + } +} + +#[test] +fn extension_constructors_reject_every_case_of_the_reserved_bytes_unit() { + for mask in 0..32 { + let unit: String = b"bytes" + .iter() + .enumerate() + .map(|(index, byte)| { + char::from(if mask & (1 << index) == 0 { + *byte + } else { + byte.to_ascii_uppercase() + }) + }) + .collect(); + for payload in ["0-1", "opaque"] { + let error = RangeOwned::extension(&unit, payload).unwrap_err(); + assert_eq!(error.header(), &FieldName::Range); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + } + for payload in ["0-1/2", "*/2", "opaque"] { + let error = ContentRangeOwned::extension(&unit, payload).unwrap_err(); + assert_eq!(error.header(), &FieldName::ContentRange); + assert_eq!(error.kind(), DecodeErrorKind::InvalidSyntax); + } + let range_wire = format!("{unit}=0-1"); + let range = RangeOwned::try_from(range_wire.as_str()).unwrap(); + assert_eq!(range.unit().unwrap(), unit); + assert!(range.is_bytes()); + assert_eq!( + range.byte_ranges().unwrap().collect::>(), + [ByteRangeSpec::FromTo { first: 0, last: 1 }] + ); + assert_eq!(range.as_field_value().as_bytes(), range_wire.as_bytes()); + assert_range_round_trip(&range); + let content_wire = format!("{unit} 0-1/2"); + let content = ContentRangeOwned::try_from(content_wire.as_str()).unwrap(); + assert_eq!(content.unit().unwrap(), unit); + assert_eq!(content.as_field_value().as_bytes(), content_wire.as_bytes()); + assert_eq!( + content.byte_range(), + Some(ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: Some(2) + }) + ); + assert_content_range_round_trip(&content); + } +} + +#[test] +fn byte_specific_constructors_keep_canonical_wire_and_semantics() { + let range = RangeOwned::bytes([ + ByteRangeSpec::from_to(0, 1).unwrap(), + ByteRangeSpec::starting_at(3), + ByteRangeSpec::suffix(2), + ]) + .unwrap(); + assert_eq!(range.as_field_value().as_bytes(), b"bytes=0-1, 3-, -2"); + assert!(range.is_bytes()); + assert_eq!(range.extension_range_set(), None); + assert_range_round_trip(&range); + for (content, wire, expected) in [ + ( + ContentRangeOwned::bytes(0, 1, CompleteLength::Known(2)).unwrap(), + "bytes 0-1/2", + ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: Some(2), + }, + ), + ( + ContentRangeOwned::bytes_range(0..2, CompleteLength::Unknown).unwrap(), + "bytes 0-1/*", + ByteContentRange::Satisfied { + first: 0, + last: 1, + complete_length: None, + }, + ), + ( + ContentRangeOwned::unsatisfied_bytes(2).unwrap(), + "bytes */2", + ByteContentRange::Unsatisfied { complete_length: 2 }, + ), + ] { + assert_eq!(content.as_field_value().as_bytes(), wire.as_bytes()); + assert_eq!(content.byte_range(), Some(expected)); + assert_eq!(content.extension_payload(), None); + assert_content_range_round_trip(&content); + } +} + +#[test] +fn nonreserved_extension_units_preserve_spelling_and_opaque_semantics() { + for unit in ["items", "Items", "bytesx", "xbytes", "ByTeS-x"] { + let range = RangeOwned::extension(unit, "opaque").unwrap(); + assert_eq!(range.unit().unwrap(), unit); + assert!(!range.is_bytes()); + assert!(range.byte_ranges().is_none()); + assert_eq!(range.extension_range_set(), Some(b"opaque".as_slice())); + assert_eq!(range.as_field_value().as_bytes(), format!("{unit}=opaque").as_bytes()); + assert_range_round_trip(&range); + let content = ContentRangeOwned::extension(unit, "\topaque / payload\t").unwrap(); + assert_eq!(content.unit().unwrap(), unit); + assert_eq!(content.byte_range(), None); + assert_eq!(content.extension_payload(), Some(b"\topaque / payload\t".as_slice())); + assert_eq!( + content.as_field_value().as_bytes(), + format!("{unit} \topaque / payload\t").as_bytes() + ); + assert_content_range_round_trip(&content); + } +} diff --git a/crates/http_headers/tests/raw_source_validation.rs b/crates/http_headers/tests/raw_source_validation.rs new file mode 100644 index 000000000..09ab28a07 --- /dev/null +++ b/crates/http_headers/tests/raw_source_validation.rs @@ -0,0 +1,198 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Regression coverage for raw custom-source field lines. + +#![cfg(feature = "headers-all")] + +use http_headers::headers::{ + Accept, AcceptEncoding, AcceptLanguage, AcceptRanges, AccessControlAllowHeaders, AccessControlExposeHeaders, + AccessControlRequestHeaders, Allow, Authorization, Basic, CacheControl, ContentSecurityPolicy, ETag, IfMatch, IfNoneMatch, + ReferrerPolicy, SecWebSocketExtensions, SecWebSocketProtocol, Server, SetCookie, UserAgent, Vary, +}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeError, DecodeErrorKind, Field, FieldName, FieldValueRef}; + +#[derive(Clone, Copy)] +enum RawRepr { + Single, + Borrowed, +} + +struct RawSource { + name: &'static FieldName, + bytes: &'static [u8], + borrowed: [FieldValueRef<'static>; 1], + repr: RawRepr, +} + +impl RawSource { + fn new(name: &'static FieldName, bytes: &'static [u8], repr: RawRepr) -> Self { + Self { + name, + bytes, + borrowed: [FieldValueRef::new(bytes)], + repr, + } + } +} + +impl FieldSource for RawSource { + fn lines(&self, name: &'static FieldName) -> Option> { + if name != self.name { + return None; + } + match self.repr { + RawRepr::Single => Some(FieldLines::single(name, self.bytes)), + RawRepr::Borrowed => FieldLines::from_borrowed(name, &self.borrowed), + } + } +} + +fn decode_owned(source: &RawSource) -> Result<(), DecodeError> { + F::owned(source).map(|_| ()) +} + +fn decode_view(source: &RawSource) -> Result<(), DecodeError> { + F::view(source).map(|_| ()) +} + +#[test] +fn affected_owned_decoders_reject_invalid_raw_field_lines() { + type Decoder = fn(&RawSource) -> Result<(), DecodeError>; + + let decoders: &[(&FieldName, Decoder)] = &[ + (&FieldName::CacheControl, decode_owned::), + (&FieldName::IfMatch, decode_owned::), + (&FieldName::IfNoneMatch, decode_owned::), + (&FieldName::Accept, decode_owned::), + (&FieldName::AcceptEncoding, decode_owned::), + (&FieldName::AcceptLanguage, decode_owned::), + (&FieldName::Allow, decode_owned::), + (&FieldName::Vary, decode_owned::), + (&FieldName::AccessControlAllowHeaders, decode_owned::), + (&FieldName::AccessControlExposeHeaders, decode_owned::), + (&FieldName::AccessControlRequestHeaders, decode_owned::), + (&FieldName::ContentSecurityPolicy, decode_owned::), + (&FieldName::ReferrerPolicy, decode_owned::), + (&FieldName::AcceptRanges, decode_owned::), + (&FieldName::SetCookie, decode_owned::), + (&FieldName::SecWebSocketExtensions, decode_owned::), + (&FieldName::SecWebSocketProtocol, decode_owned::), + ]; + let invalid = [ + b"\r".as_slice(), + b"\n".as_slice(), + b"\0".as_slice(), + b"\x1f".as_slice(), + b"\x7f".as_slice(), + ]; + + for &(name, decode) in decoders { + for &bytes in &invalid { + for repr in [RawRepr::Single, RawRepr::Borrowed] { + let source = RawSource::new(name, bytes, repr); + decode(&source).expect_err(name.as_str()); + } + } + } +} + +#[test] +fn affected_borrowed_decoders_reject_invalid_raw_field_lines() { + type Decoder = fn(&RawSource) -> Result<(), DecodeError>; + + let decoders: &[(&FieldName, Decoder)] = &[ + (&FieldName::UserAgent, decode_view::), + (&FieldName::Server, decode_view::), + (&FieldName::SetCookie, decode_view::), + (&FieldName::ContentSecurityPolicy, decode_view::), + ]; + + for bytes in [ + b"\r".as_slice(), + b"\n".as_slice(), + b"\0".as_slice(), + b"\x1f".as_slice(), + b"\x7f".as_slice(), + ] { + for &(name, decode) in decoders { + for repr in [RawRepr::Single, RawRepr::Borrowed] { + let source = RawSource::new(name, bytes, repr); + assert_eq!(decode(&source).expect_err(name.as_str()).kind(), DecodeErrorKind::InvalidSyntax); + } + } + } +} + +#[test] +fn entity_tags_reject_del_from_raw_sources() { + for bytes in [b"\"\x7f\"".as_slice(), b"\"abc\x7fdefgh\""] { + for repr in [RawRepr::Single, RawRepr::Borrowed] { + let source = RawSource::new(&FieldName::Etag, bytes, repr); + assert_eq!( + decode_view::(&source).expect_err("DEL is not etagc").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert_eq!( + decode_owned::(&source).expect_err("DEL is not etagc").kind(), + DecodeErrorKind::InvalidSyntax + ); + } + } +} + +#[test] +fn delimited_iteration_reports_preflight_errors_once() { + let mut items = FieldLines::single(&FieldName::Vary, b"valid,\ninvalid").comma_items(); + assert_eq!( + items.next().expect("preflight error").expect_err("invalid raw bytes").kind(), + DecodeErrorKind::InvalidSyntax + ); + assert!(items.next().is_none()); +} + +#[test] +fn sensitive_borrowed_headers_override_source_classification() { + let authorization_source = RawSource::new(&FieldName::Authorization, b"Basic dXNlcjpwYXNz", RawRepr::Borrowed); + let authorization = Authorization::::view(&authorization_source) + .expect("valid authorization") + .expect("authorization is present"); + let authorization_value = authorization.as_field_value(); + assert!(authorization_value.is_sensitive()); + assert!(authorization_value.try_to_field_value().unwrap().is_sensitive()); + + let cookie_source = RawSource::new(&FieldName::SetCookie, b"session=secret", RawRepr::Borrowed); + let cookies = SetCookie::view(&cookie_source) + .expect("valid Set-Cookie") + .expect("Set-Cookie is present"); + let cookie_value = cookies.iter().next().unwrap(); + assert!(cookie_value.is_sensitive()); + assert!(cookie_value.try_to_field_value().unwrap().is_sensitive()); + + #[cfg(feature = "http")] + { + let authorization_value = http::HeaderValue::try_from(authorization_value).unwrap(); + let cookie_value = http::HeaderValue::try_from(cookie_value).unwrap(); + assert!(authorization_value.is_sensitive()); + assert!(cookie_value.is_sensitive()); + } +} + +#[cfg(feature = "http")] +#[test] +fn sensitive_borrowed_headers_override_http_map_classification() { + let mut map = http::HeaderMap::new(); + map.insert(http::header::AUTHORIZATION, http::HeaderValue::from_static("Basic dXNlcjpwYXNz")); + map.insert(http::header::SET_COOKIE, http::HeaderValue::from_static("session=secret")); + + let authorization = Authorization::::view(&map).unwrap().unwrap().as_field_value(); + let cookie = SetCookie::view(&map).unwrap().unwrap().iter().next().unwrap(); + for (value, secret) in [(authorization, "dXNlcjpwYXNz"), (cookie, "session=secret")] { + assert!(value.is_sensitive()); + let owned = value.try_to_field_value().unwrap(); + assert!(owned.is_sensitive()); + assert!(!format!("{owned:?}").contains(secret)); + assert!(http::HeaderValue::try_from(value).unwrap().is_sensitive()); + } +} diff --git a/crates/http_headers/tests/response_failures.rs b/crates/http_headers/tests/response_failures.rs new file mode 100644 index 000000000..0b9ac9ae1 --- /dev/null +++ b/crates/http_headers/tests/response_failures.rs @@ -0,0 +1,117 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Downstream checks for response method success and storage-error propagation. + +#![cfg(feature = "headers-all")] + +use std::time::Duration; + +use http_headers::headers::{CacheControl, SetCookieOwned, UserAgentOwned}; +use http_headers::sink::{EncodedValues, FieldSink, FieldSinkExt, InsertError, InsertErrorKind}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{FieldName, FieldValue}; + +const STORAGE_ERROR: InsertError = InsertError::new(InsertErrorKind::CapacityExceeded); + +#[derive(Debug, Default)] +struct ToggleSink { + reject: bool, + name: Option<&'static FieldName>, + values: Vec, +} + +impl FieldSource for ToggleSink { + fn lines(&self, name: &'static FieldName) -> Option> { + if self.name == Some(name) { + FieldLines::from_slice(name, &self.values) + } else { + None + } + } +} + +impl FieldSink for ToggleSink { + fn set_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + if self.reject { + Err(STORAGE_ERROR) + } else { + self.name = Some(name); + self.values = values.into_iter().collect(); + Ok(()) + } + } + + fn append_values(&mut self, name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + if self.reject { + return Err(STORAGE_ERROR); + } + if self.name == Some(name) { + self.values.extend(values); + } else { + self.name = Some(name); + self.values = values.into_iter().collect(); + } + Ok(()) + } + + fn remove_values(&mut self, name: &'static FieldName) { + if self.name == Some(name) { + self.name = None; + self.values.clear(); + } + } +} + +#[test] +fn response_methods_cover_success_and_failure_in_the_same_downstream_sink() { + let mut sink = ToggleSink::default(); + + sink.set_user_agent(UserAgentOwned::try_from_static("client/1").expect("valid")) + .expect("enabled storage succeeds"); + assert_eq!(sink.values[0], "client/1"); + sink.reject = true; + assert_eq!( + sink.set_user_agent(UserAgentOwned::try_from_static("client/2").expect("valid")) + .err(), + Some(STORAGE_ERROR) + ); + assert_eq!(sink.values[0], "client/1"); + + sink.reject = false; + sink.set_content_length(42).expect("enabled storage succeeds"); + assert_eq!(sink.values[0], "42"); + sink.reject = true; + assert_eq!(sink.set_content_length(42).err(), Some(STORAGE_ERROR)); + assert_eq!(sink.values[0], "42"); + + sink.reject = false; + sink.set_cache_control(CacheControl::private()).expect("enabled storage succeeds"); + assert_eq!(sink.values[0], "private"); + sink.reject = true; + assert_eq!(sink.set_cache_control(CacheControl::no_cache()).err(), Some(STORAGE_ERROR)); + assert_eq!(sink.values[0], "private"); + + sink.reject = false; + sink.set_access_control_max_age(600).expect("enabled storage succeeds"); + assert_eq!(sink.values[0], "600"); + sink.reject = true; + assert_eq!(sink.set_access_control_max_age(300).err(), Some(STORAGE_ERROR)); + assert_eq!( + sink.set_access_control_max_age_duration(Duration::from_mins(5)).err(), + Some(STORAGE_ERROR) + ); + assert_eq!(sink.values[0], "600"); + + sink.reject = false; + let mut first = SetCookieOwned::new(); + first.push_str("a=1").expect("valid cookie"); + sink.append_set_cookie(first).expect("enabled storage succeeds"); + assert_eq!(sink.values[0], "a=1"); + sink.reject = true; + let mut second = SetCookieOwned::new(); + second.push_str("b=2").expect("valid cookie"); + assert_eq!(sink.append_set_cookie(second).err(), Some(STORAGE_ERROR)); + assert_eq!(sink.values.len(), 1); + assert_eq!(sink.values[0], "a=1"); +} diff --git a/crates/http_headers/tests/sec_web_socket_version.rs b/crates/http_headers/tests/sec_web_socket_version.rs new file mode 100644 index 000000000..771d9eb07 --- /dev/null +++ b/crates/http_headers/tests/sec_web_socket_version.rs @@ -0,0 +1,93 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Version 13 decoding retains complete version-set and source validation. + +#![cfg(all(feature = "http", feature = "headers-websocket"))] +#![expect(clippy::unwrap_used, reason = "test failures provide sufficient context")] + +use http::{HeaderMap, HeaderValue}; +use http_headers::headers::SecWebSocketVersion; +use http_headers::source::{FieldLines, FieldSource, MAX_CUSTOM_FIELD_LINES}; +use http_headers::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue}; + +fn assert_versions(source: &impl FieldSource, expected: &[u8]) { + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let view = SecWebSocketVersion::view_with(source, mode).unwrap().unwrap(); + let owned = SecWebSocketVersion::owned_with(source, mode).unwrap().unwrap(); + assert_eq!(view.versions().collect::>(), expected); + assert_eq!(owned, view); + if let [only] = expected { + assert_eq!(view.requested(), Ok(*only)); + } else { + assert_eq!(view.requested().unwrap_err().kind(), DecodeErrorKind::UnexpectedMultipleValues); + } + } +} + +fn assert_error(source: &impl FieldSource, expected: DecodeErrorKind) { + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!(SecWebSocketVersion::view_with(source, mode).unwrap_err().kind(), expected); + assert_eq!(SecWebSocketVersion::owned_with(source, mode).unwrap_err().kind(), expected); + } +} + +fn version_map(values: &[&str]) -> HeaderMap { + let mut map = HeaderMap::new(); + for value in values { + map.append( + http::header::SEC_WEBSOCKET_VERSION, + HeaderValue::from_bytes(value.as_bytes()).unwrap(), + ); + } + map +} + +#[test] +fn owned_and_borrowed_decode_every_version() { + for version in 0..=u8::MAX { + assert_versions(&version_map(&[&version.to_string()]), &[version]); + } + assert_versions(&version_map(&[" \t13\t "]), &[13]); +} + +#[test] +fn version_thirteen_does_not_hide_later_versions_or_errors() { + assert_versions(&version_map(&["13", "8"]), &[8, 13]); + assert_versions(&version_map(&["13", "13"]), &[13]); + assert_versions(&version_map(&["13, 8, 13"]), &[8, 13]); + assert_versions(&version_map(&["13", ", 7, 13, ,"]), &[7, 13]); + assert_versions(&version_map(&["13", "255"]), &[13, 255]); + + for invalid in ["013", "13x", "1300", "256"] { + assert_error(&version_map(&[invalid]), DecodeErrorKind::InvalidSyntax); + assert_error(&version_map(&["13", invalid]), DecodeErrorKind::InvalidSyntax); + } + assert_error(&version_map(&["13", "\"13"]), DecodeErrorKind::UnterminatedQuote); + assert_error(&version_map(&[""]), DecodeErrorKind::MissingValue); + + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!(SecWebSocketVersion::view_with(&HeaderMap::new(), mode).unwrap(), None); + assert_eq!(SecWebSocketVersion::owned_with(&HeaderMap::new(), mode).unwrap(), None); + } +} + +struct VersionsSource(Vec); + +impl FieldSource for VersionsSource { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::SecWebSocketVersion) + .then(|| FieldLines::from_slice(name, &self.0)) + .flatten() + } +} + +#[test] +fn version_thirteen_preserves_custom_source_limits() { + let mut source = VersionsSource(vec![FieldValue::from_static("13")]); + assert_versions(&source, &[13]); + source.0.resize(MAX_CUSTOM_FIELD_LINES, FieldValue::from_static("13")); + assert_versions(&source, &[13]); + source.0.push(FieldValue::from_static("13")); + assert_error(&source, DecodeErrorKind::SourceLimitExceeded); +} diff --git a/crates/http_headers/tests/serde.rs b/crates/http_headers/tests/serde.rs new file mode 100644 index 000000000..e70137e6b --- /dev/null +++ b/crates/http_headers/tests/serde.rs @@ -0,0 +1,574 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Serde round-trip and validation coverage for owned header data. + +#![cfg(feature = "headers-all")] +#![cfg(feature = "serde")] +#![expect(clippy::assertions_on_result_states, reason = "rejection is the complete contract under test")] +#![expect( + clippy::string_lit_as_bytes, + reason = "string literals keep the exhaustive byte-oriented table readable" +)] +#![expect(clippy::unwrap_used, reason = "test failures provide sufficient context")] + +use http_headers::headers::*; +use http_headers::sink::EncodedValues; +use http_headers::source::{FieldLines, FieldSource, MAX_CUSTOM_FIELD_BYTES, MAX_CUSTOM_FIELD_LINES, MAX_CUSTOM_LIST_ITEMS}; +use http_headers::{Field, FieldName, FieldSensitivity, FieldValue}; +use serde::de::value::{BorrowedBytesDeserializer, Error as ValueError}; +use serde::de::{DeserializeOwned, DeserializeSeed, IntoDeserializer, MapAccess, SeqAccess, Visitor}; +use serde::{Deserialize, Deserializer, Serialize, forward_to_deserialize_any}; + +struct Values(Vec); + +impl FieldSource for Values { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_slice(name, &self.0) + } +} + +fn assert_owned_round_trip(lines: &[&[u8]]) +where + H: Field, + H::Owned: Clone + DeserializeOwned + Serialize, +{ + let source = Values(lines.iter().map(|line| FieldValue::from_bytes(line).unwrap()).collect()); + let original = H::owned(&source).unwrap().unwrap(); + let json = serde_json::to_string(&original).unwrap(); + let decoded: H::Owned = serde_json::from_str(&json).unwrap(); + assert_eq!(serde_json::to_value(decoded).unwrap(), serde_json::to_value(original).unwrap()); +} + +#[test] +fn every_owned_header_round_trips_through_validated_field_values() { + macro_rules! cases { + ($(($header:ty, $($line:expr),+)),+ $(,)?) => { + $(assert_owned_round_trip::<$header>(&[$($line.as_bytes()),+]));+ + }; + } + + cases!( + (CacheControl, "public, max-age=60"), + (Accept, "text/html"), + (AcceptEncoding, "gzip"), + (AcceptLanguage, "en-US"), + (Allow, "GET, POST"), + (Host, "example.com:443"), + (Server, "example/1.0"), + (Vary, "accept-encoding"), + (AcceptRanges, "bytes"), + (ContentRange, "bytes 0-9/10"), + (Range, "bytes=0-9"), + (ETag, "\"tag\""), + (Location, "/next"), + (UserAgent, "example-client/1.0"), + (ContentType, "text/plain; charset=utf-8"), + (IfMatch, "\"tag\""), + (IfNoneMatch, "*"), + (IfModifiedSince, "Sun, 06 Nov 1994 08:49:37 GMT"), + (IfUnmodifiedSince, "Sun, 06 Nov 1994 08:49:37 GMT"), + (IfRange, "\"tag\""), + (LastModified, "Sun, 06 Nov 1994 08:49:37 GMT"), + (AccessControlAllowCredentials, "true"), + (AccessControlAllowHeaders, "content-type, x-request-id"), + (AccessControlAllowMethods, "GET, POST"), + (AccessControlAllowOrigin, "https://example.com"), + (AccessControlExposeHeaders, "etag, x-request-id"), + (AccessControlMaxAge, "600"), + (AccessControlRequestHeaders, "content-type"), + (AccessControlRequestMethod, "POST"), + (ContentSecurityPolicy, "default-src 'self'"), + (ReferrerPolicy, "no-referrer"), + (StrictTransportSecurity, "max-age=31536000; includeSubDomains"), + (XContentTypeOptions, "nosniff"), + (SecWebSocketAccept, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="), + (SecWebSocketExtensions, "permessage-deflate"), + (SecWebSocketKey, "dGhlIHNhbXBsZSBub25jZQ=="), + (SecWebSocketProtocol, "chat"), + (SecWebSocketVersion, "13"), + (Authorization, "Basic dXNlcjpwYXNz"), + (Authorization, "Bearer token"), + (ContentLength, "42"), + (SetCookie, "session=abc; Path=/", "theme=dark; Path=/"), + ); +} + +#[test] +fn core_types_preserve_names_bytes_sensitivity_and_line_boundaries() { + let name = FieldName::try_from_bytes("X-Example").unwrap(); + let decoded_name: FieldName = serde_json::from_str(&serde_json::to_string(&name).unwrap()).unwrap(); + assert_eq!(decoded_name.as_str(), "x-example"); + + let binary = FieldValue::from_bytes(b"text\x80") + .unwrap() + .with_sensitivity(FieldSensitivity::Sensitive); + let decoded_binary: FieldValue = serde_json::from_str(&serde_json::to_string(&binary).unwrap()).unwrap(); + assert_eq!(decoded_binary.as_bytes(), b"text\x80"); + assert!(decoded_binary.is_sensitive()); + + let values = EncodedValues::from_vec(vec![ + FieldValue::from_static("first"), + FieldValue::from_static("second").with_sensitivity(FieldSensitivity::Sensitive), + ]); + let decoded_values: EncodedValues = serde_json::from_str(&serde_json::to_string(&values).unwrap()).unwrap(); + let decoded_values = decoded_values.into_iter().collect::>(); + assert_eq!(decoded_values.len(), 2); + assert_eq!(decoded_values[0].as_bytes(), b"first"); + assert_eq!(decoded_values[1].as_bytes(), b"second"); + assert!(decoded_values[1].is_sensitive()); +} + +#[cfg(feature = "http")] +#[test] +fn http_field_names_round_trip_through_serde() { + let name = FieldName::from(http::HeaderName::from_lowercase(b"x\"y").unwrap()); + let decoded: FieldName = serde_json::from_str(&serde_json::to_string(&name).unwrap()).unwrap(); + assert_eq!(decoded, name); +} + +#[test] +fn repeated_cache_and_range_lines_round_trip_without_collapsing() { + for (name, values) in [ + ( + &FieldName::CacheControl, + vec![FieldValue::from_static("public"), FieldValue::from_static("max-age=60")], + ), + ( + &FieldName::AcceptRanges, + vec![FieldValue::from_static("bytes"), FieldValue::from_static("items")], + ), + ] { + let expected = serde_json::to_value(&values).unwrap(); + let source = Values(values); + if name == &FieldName::CacheControl { + let original = CacheControl::owned(&source).unwrap().unwrap(); + assert_eq!(serde_json::to_value(&original).unwrap(), expected); + let decoded: CacheControlOwned = serde_json::from_value(expected).unwrap(); + assert_eq!(serde_json::to_value(decoded).unwrap(), serde_json::to_value(original).unwrap()); + } else { + let original = AcceptRanges::owned(&source).unwrap().unwrap(); + assert_eq!(serde_json::to_value(&original).unwrap(), expected); + let decoded: AcceptRangesOwned = serde_json::from_value(expected).unwrap(); + assert_eq!(serde_json::to_value(decoded).unwrap(), serde_json::to_value(original).unwrap()); + } + } + + let none = AcceptRangesOwned::none(); + let expected = serde_json::to_value([FieldValue::from_static("none")]).unwrap(); + assert_eq!(serde_json::to_value(none).unwrap(), expected); +} + +#[test] +fn relaxed_owned_values_round_trip_through_serde() { + let source = Values(vec![FieldValue::from_static("/a\\b")]); + let original = Location::owned_with(&source, http_headers::DecodeMode::Relaxed).unwrap().unwrap(); + let decoded: LocationOwned = serde_json::from_str(&serde_json::to_string(&original).unwrap()).unwrap(); + assert_eq!(decoded.as_bytes(), original.as_bytes()); +} + +#[test] +fn typed_deserialization_reuses_header_validation() { + let invalid_value = r#"[{"bytes":[10],"sensitivity":"NonSensitive"}]"#; + assert!(serde_json::from_str::(invalid_value).is_err()); + + let multiple_values = r#"[{"bytes":[116,101,120,116,47,112,108,97,105,110],"sensitivity":"NonSensitive"},{"bytes":[116,101,120,116,47,104,116,109,108],"sensitivity":"NonSensitive"}]"#; + assert!(serde_json::from_str::(multiple_values).is_err()); + + let invalid_name = serde_json::to_string("bad\nname").unwrap(); + assert!(serde_json::from_str::(&invalid_name).is_err()); + + let extension_method = AccessControlRequestMethodOwned::try_from("CUSTOM").unwrap(); + let extension_json = serde_json::to_string(&extension_method).unwrap(); + let decoded_extension: AccessControlRequestMethodOwned = serde_json::from_str(&extension_json).unwrap(); + assert_eq!(decoded_extension.method().unwrap().as_str(), "CUSTOM"); +} + +#[test] +fn empty_set_cookie_round_trips_without_becoming_absent() { + let json = serde_json::to_string(&SetCookieOwned::new()).unwrap(); + let decoded: SetCookieOwned = serde_json::from_str(&json).unwrap(); + assert!(decoded.is_empty()); +} + +#[test] +fn fixed_payloads_pin_the_successful_serde_representation() { + const NAME: &str = r#""x-example""#; + const SENSITIVE_BINARY: &str = r#"{"bytes":[116,101,120,116,128],"sensitivity":"Sensitive"}"#; + const ENCODED_VALUES: &str = + r#"[{"bytes":[102,105,114,115,116],"sensitivity":"NonSensitive"},{"bytes":[115,101,99,111,110,100],"sensitivity":"Sensitive"}]"#; + const USER_AGENT: &str = + r#"[{"bytes":[101,120,97,109,112,108,101,45,99,108,105,101,110,116,47,49,46,48],"sensitivity":"NonSensitive"}]"#; + const ACCEPT: &str = r#"[{"bytes":[116,101,120,116,47,104,116,109,108],"sensitivity":"NonSensitive"},{"bytes":[97,112,112,108,105,99,97,116,105,111,110,47,106,115,111,110],"sensitivity":"NonSensitive"}]"#; + const SET_COOKIE: &str = r#"[{"bytes":[115,101,115,115,105,111,110,61,97,98,99],"sensitivity":"Sensitive"},{"bytes":[116,104,101,109,101,61,100,97,114,107],"sensitivity":"Sensitive"}]"#; + + let name: FieldName = serde_json::from_str(NAME).unwrap(); + assert_eq!(name.as_str(), "x-example"); + assert_eq!(serde_json::to_string(&name).unwrap(), NAME); + + let binary: FieldValue = serde_json::from_str(SENSITIVE_BINARY).unwrap(); + assert_eq!(binary.as_bytes(), b"text\x80"); + assert!(binary.is_sensitive()); + assert_eq!(serde_json::to_string(&binary).unwrap(), SENSITIVE_BINARY); + + let values: EncodedValues = serde_json::from_str(ENCODED_VALUES).unwrap(); + let lines = values.iter().collect::>(); + assert_eq!(lines[0].as_bytes(), b"first"); + assert_eq!(lines[1].as_bytes(), b"second"); + assert!(!lines[0].is_sensitive()); + assert!(lines[1].is_sensitive()); + assert_eq!(serde_json::to_string(&values).unwrap(), ENCODED_VALUES); + + let user_agent: UserAgentOwned = serde_json::from_str(USER_AGENT).unwrap(); + assert_eq!(user_agent.as_bytes(), b"example-client/1.0"); + assert_eq!(serde_json::to_string(&user_agent).unwrap(), USER_AGENT); + + let accept: AcceptOwned = serde_json::from_str(ACCEPT).unwrap(); + assert_eq!( + accept.values().map(http_headers::FieldValueRef::as_bytes).collect::>(), + [b"text/html".as_slice(), b"application/json".as_slice()] + ); + assert_eq!(serde_json::to_string(&accept).unwrap(), ACCEPT); + + let cookies: SetCookieOwned = serde_json::from_str(SET_COOKIE).unwrap(); + assert_eq!( + cookies.iter().map(FieldValue::as_bytes).collect::>(), + [b"session=abc".as_slice(), b"theme=dark".as_slice()] + ); + assert!(cookies.iter().all(FieldValue::is_sensitive)); + assert_eq!(serde_json::to_string(&cookies).unwrap(), SET_COOKIE); +} + +struct HostileBytesDeserializer { + remaining: usize, + reported: usize, +} + +struct HostileBytesAccess { + remaining: usize, + reported: usize, +} + +impl<'de> SeqAccess<'de> for HostileBytesAccess { + type Error = ValueError; + + fn next_element_seed>(&mut self, seed: T) -> Result, Self::Error> { + if self.remaining == 0 { + return Ok(None); + } + self.remaining -= 1; + seed.deserialize(b'a'.into_deserializer()).map(Some) + } + + fn size_hint(&self) -> Option { + Some(self.reported) + } +} + +impl<'de> Deserializer<'de> for HostileBytesDeserializer { + type Error = ValueError; + + fn deserialize_any>(self, visitor: V) -> Result { + self.deserialize_seq(visitor) + } + + fn deserialize_seq>(self, visitor: V) -> Result { + visitor.visit_seq(HostileBytesAccess { + remaining: self.remaining, + reported: self.reported, + }) + } + + forward_to_deserialize_any! { + bool i8 i16 i32 i64 u8 u16 u32 u64 f32 f64 char str string bytes byte_buf + option unit unit_struct newtype_struct tuple tuple_struct map struct enum identifier + ignored_any + } +} + +struct HostileFieldValueDeserializer { + byte_count: usize, + byte_hint: usize, +} + +struct HostileFieldValueMap { + state: u8, + byte_count: usize, + byte_hint: usize, +} + +impl<'de> MapAccess<'de> for HostileFieldValueMap { + type Error = ValueError; + + fn next_key_seed>(&mut self, seed: K) -> Result, Self::Error> { + let key = match self.state { + 0 => "ignored", + 1 => "bytes", + 2 => "sensitivity", + _ => return Ok(None), + }; + self.state += 1; + seed.deserialize(BorrowedBytesDeserializer::new(key.as_bytes())).map(Some) + } + + fn next_value_seed>(&mut self, seed: V) -> Result { + match self.state { + 1 => seed.deserialize(0_u8.into_deserializer()), + 2 => seed.deserialize(HostileBytesDeserializer { + remaining: self.byte_count, + reported: self.byte_hint, + }), + 3 => seed.deserialize("NonSensitive".into_deserializer()), + _ => Err(serde::de::Error::custom("value requested without a key")), + } + } +} + +impl<'de> Deserializer<'de> for HostileFieldValueDeserializer { + type Error = ValueError; + + fn deserialize_any>(self, visitor: V) -> Result { + self.deserialize_map(visitor) + } + + fn deserialize_map>(self, visitor: V) -> Result { + visitor.visit_map(HostileFieldValueMap { + state: 0, + byte_count: self.byte_count, + byte_hint: self.byte_hint, + }) + } + + fn deserialize_struct>( + self, + _name: &'static str, + _fields: &'static [&'static str], + visitor: V, + ) -> Result { + self.deserialize_map(visitor) + } + + forward_to_deserialize_any! { + bool i8 i16 i32 i64 u8 u16 u32 u64 f32 f64 char str string bytes byte_buf + option unit unit_struct newtype_struct seq tuple tuple_struct enum identifier ignored_any + } +} + +struct HostileValuesDeserializer { + remaining: usize, + reported: usize, + byte_count: usize, + byte_hint: usize, +} + +impl<'de> SeqAccess<'de> for HostileValuesDeserializer { + type Error = ValueError; + + fn next_element_seed>(&mut self, seed: T) -> Result, Self::Error> { + if self.remaining == 0 { + return Ok(None); + } + self.remaining -= 1; + seed.deserialize(HostileFieldValueDeserializer { + byte_count: self.byte_count, + byte_hint: self.byte_hint, + }) + .map(Some) + } + + fn size_hint(&self) -> Option { + Some(self.reported) + } +} + +impl<'de> Deserializer<'de> for HostileValuesDeserializer { + type Error = ValueError; + + fn deserialize_any>(self, visitor: V) -> Result { + self.deserialize_seq(visitor) + } + + fn deserialize_seq>(self, visitor: V) -> Result { + visitor.visit_seq(self) + } + + forward_to_deserialize_any! { + bool i8 i16 i32 i64 u8 u16 u32 u64 f32 f64 char str string bytes byte_buf + option unit unit_struct newtype_struct tuple tuple_struct map struct enum identifier + ignored_any + } +} + +#[test] +fn hostile_collection_hints_cannot_drive_unbounded_deserialization_allocations() { + let valid = EncodedValues::deserialize(HostileValuesDeserializer { + remaining: 1, + reported: usize::MAX, + byte_count: 1, + byte_hint: usize::MAX, + }) + .unwrap(); + assert_eq!(valid.iter().next().unwrap().as_bytes(), b"a"); + + for (byte_limit, extra_byte, line_limit, extra_line) in [ + (65_536, 65_537, 128, 129), + ( + MAX_CUSTOM_FIELD_BYTES, + MAX_CUSTOM_FIELD_BYTES + 1, + MAX_CUSTOM_FIELD_LINES, + MAX_CUSTOM_FIELD_LINES + 1, + ), + ] { + let byte_boundary = FieldValue::deserialize(HostileFieldValueDeserializer { + byte_count: byte_limit, + byte_hint: usize::MAX, + }) + .unwrap(); + assert_eq!(byte_boundary.as_bytes().len(), byte_limit); + + let oversized_bytes = FieldValue::deserialize(HostileFieldValueDeserializer { + byte_count: extra_byte, + byte_hint: usize::MAX, + }) + .unwrap(); + assert_eq!(oversized_bytes.as_bytes().len(), extra_byte); + + let line_boundary = EncodedValues::deserialize(HostileValuesDeserializer { + remaining: line_limit, + reported: usize::MAX, + byte_count: 0, + byte_hint: usize::MAX, + }) + .unwrap(); + assert_eq!(line_boundary.len(), line_limit); + + let oversized_lines = EncodedValues::deserialize(HostileValuesDeserializer { + remaining: extra_line, + reported: usize::MAX, + byte_count: 0, + byte_hint: usize::MAX, + }) + .unwrap(); + assert_eq!(oversized_lines.len(), extra_line); + + let oversized_typed_bytes = UserAgentOwned::deserialize(HostileValuesDeserializer { + remaining: 1, + reported: usize::MAX, + byte_count: extra_byte, + byte_hint: usize::MAX, + }); + assert_eq!( + oversized_typed_bytes.unwrap_err().to_string(), + "source limit exceeded: field-value byte budget exceeded" + ); + + let line_boundary = SetCookieOwned::deserialize(HostileValuesDeserializer { + remaining: line_limit, + reported: usize::MAX, + byte_count: 1, + byte_hint: usize::MAX, + }) + .unwrap(); + assert_eq!(line_boundary.len(), line_limit); + + let oversized_typed_lines = SetCookieOwned::deserialize(HostileValuesDeserializer { + remaining: extra_line, + reported: usize::MAX, + byte_count: 1, + byte_hint: usize::MAX, + }); + assert_eq!( + oversized_typed_lines.unwrap_err().to_string(), + "source limit exceeded: field-value line budget exceeded" + ); + } +} + +#[test] +fn typed_serde_accumulates_the_byte_budget_across_field_lines() { + let boundary = SetCookieOwned::deserialize(HostileValuesDeserializer { + remaining: 2, + reported: usize::MAX, + byte_count: 32_768, + byte_hint: usize::MAX, + }) + .unwrap(); + assert_eq!(boundary.len(), 2); + assert_eq!(boundary.iter().map(|value| value.as_bytes().len()).sum::(), 65_536); + assert!(boundary.iter().all(|value| value.as_bytes().len() == 32_768)); + + let oversized = SetCookieOwned::deserialize(HostileValuesDeserializer { + remaining: 2, + reported: usize::MAX, + byte_count: 32_769, + byte_hint: usize::MAX, + }) + .unwrap_err(); + assert_eq!(oversized.to_string(), "source limit exceeded: field-value byte budget exceeded"); +} + +#[test] +fn typed_serde_enforces_list_item_budgets_at_the_collection_boundary() { + for (item_limit, extra_item) in [(1_024, 1_025), (MAX_CUSTOM_LIST_ITEMS, MAX_CUSTOM_LIST_ITEMS + 1)] { + let boundary = std::iter::repeat_n("*/*", item_limit).collect::>().join(","); + let boundary = EncodedValues::single(FieldValue::try_from(boundary).unwrap()); + let json = serde_json::to_string(&boundary).unwrap(); + assert_eq!(serde_json::from_str::(&json).unwrap().entries().count(), item_limit); + + let over_limit = std::iter::repeat_n("*/*", extra_item).collect::>().join(","); + let over_limit = EncodedValues::single(FieldValue::try_from(over_limit).unwrap()); + let json = serde_json::to_string(&over_limit).unwrap(); + assert!( + serde_json::from_str::(&json) + .unwrap_err() + .to_string() + .contains("invalid accept header: source limit exceeded") + ); + } +} + +#[test] +fn typed_serde_preserves_grammar_error_kinds() { + for (wire, expected) in [ + ("text/html;x", "invalid accept header: invalid syntax"), + ("text/pl ain", "invalid accept header: invalid token"), + ("text/html;x=\"bad", "invalid accept header: unterminated quoted string"), + ] { + let encoded = EncodedValues::single(FieldValue::from_str(wire).unwrap()); + let json = serde_json::to_string(&encoded).unwrap(); + let error = serde_json::from_str::(&json).unwrap_err().to_string(); + assert!(error.starts_with(expected), "{error}"); + } +} + +#[test] +fn field_value_deserialization_covers_sequence_metadata_and_diagnostics() { + let sequence: FieldValue = serde_json::from_str(r#"[[97],"NonSensitive"]"#).unwrap(); + assert_eq!(sequence.as_bytes(), b"a"); + + let with_unknown: FieldValue = serde_json::from_str(r#"{"ignored":0,"bytes":[97],"sensitivity":"NonSensitive"}"#).unwrap(); + assert_eq!(with_unknown.as_bytes(), b"a"); + + for invalid in [ + r#"{"bytes":[97],"bytes":[98],"sensitivity":"NonSensitive"}"#, + r#"{"bytes":[97],"sensitivity":"NonSensitive","sensitivity":"Sensitive"}"#, + r#"{"sensitivity":"NonSensitive"}"#, + r#"{"bytes":[97]}"#, + "null", + ] { + assert!(serde_json::from_str::(invalid).is_err(), "{invalid}"); + } + assert!(serde_json::from_str::(r#"{"bytes":"not bytes","sensitivity":"NonSensitive"}"#).is_err()); + assert!(serde_json::from_str::("null").is_err()); + + for byte_count in [65_536, MAX_CUSTOM_FIELD_BYTES] { + let boundary = UserAgentOwned::deserialize(HostileValuesDeserializer { + remaining: 1, + reported: usize::MAX, + byte_count, + byte_hint: usize::MAX, + }) + .unwrap(); + assert_eq!(boundary.as_bytes().len(), byte_count); + } +} diff --git a/crates/http_headers/tests/source_limits.rs b/crates/http_headers/tests/source_limits.rs new file mode 100644 index 000000000..908d6b788 --- /dev/null +++ b/crates/http_headers/tests/source_limits.rs @@ -0,0 +1,521 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Regression coverage for custom-source decode budgets. + +#![cfg(feature = "headers-all")] + +use std::iter; + +#[cfg(feature = "http")] +use http_headers::headers::UserAgent; +use http_headers::headers::{ + Accept, AcceptRanges, AccessControlAllowCredentials, AccessControlAllowOrigin, AccessControlMaxAge, AccessControlRequestMethod, + ContentSecurityPolicy, IfMatch, Location, Range, ReferrerPolicy, SetCookie, StrictTransportSecurity, +}; +use http_headers::sink::{EncodedValues, FieldSink, InsertError, InsertErrorKind}; +use http_headers::source::{FieldLines, FieldSource, MAX_CUSTOM_FIELD_BYTES, MAX_CUSTOM_FIELD_LINES, MAX_CUSTOM_LIST_ITEMS}; +use http_headers::{DecodeErrorKind, DecodeMode, Field, FieldName, FieldValue, FieldValueRef}; + +#[test] +fn published_source_limits_have_stable_values() { + assert_eq!(MAX_CUSTOM_FIELD_BYTES, 65_536); + assert_eq!(MAX_CUSTOM_FIELD_LINES, 128); + assert_eq!(MAX_CUSTOM_LIST_ITEMS, 1_024); +} + +fn assert_source_limit(source: &impl FieldSource) { + assert_decode_error::(source, DecodeErrorKind::SourceLimitExceeded); +} + +fn assert_decode_error(source: &impl FieldSource, expected: DecodeErrorKind) { + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!(F::view_with(source, mode).map(|_| ()).map_err(|error| error.kind()), Err(expected)); + assert_eq!(F::owned_with(source, mode).map(|_| ()).map_err(|error| error.kind()), Err(expected)); + } +} + +struct SingleSource { + name: &'static FieldName, + bytes: Vec, +} + +impl FieldSource for SingleSource { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::single(name, &self.bytes)) + } +} + +struct ValuesSource { + name: &'static FieldName, + values: Vec, +} + +impl FieldSource for ValuesSource { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::from_slice(name, &self.values)).flatten() + } +} + +struct BorrowedSource<'a>(&'a [FieldValueRef<'a>]); + +impl FieldSource for BorrowedSource<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + FieldLines::from_borrowed(name, self.0) + } +} + +fn repeated_value(value: &'static str, count: usize) -> Vec { + (0..count).map(|_| FieldValue::from_static(value)).collect() +} + +fn comma_list(item: &str, count: usize) -> Vec { + iter::repeat_n(item, count).collect::>().join(",").into_bytes() +} + +fn range_list(count: usize) -> Vec { + let mut value = b"bytes=".to_vec(); + value.extend_from_slice(&comma_list("0-0", count)); + value +} + +fn hsts_directives(count: usize) -> Vec { + let mut value = b"max-age=1".to_vec(); + for _ in 1..count { + value.extend_from_slice(b";x"); + } + value +} + +fn padded_value(value: &[u8]) -> Vec { + let mut padded = vec![b' '; 65_536]; + padded.extend_from_slice(value); + padded +} + +#[test] +fn custom_source_total_byte_limit_accepts_boundary_and_rejects_overflow() { + let half = 65_536 / 2; + let boundary = ValuesSource { + name: &FieldName::SetCookie, + values: vec![ + FieldValue::try_from(vec![b'a'; half]).expect("valid field value"), + FieldValue::try_from(vec![b'b'; half]).expect("valid field value"), + ], + }; + let decoded = SetCookie::owned(&boundary) + .expect("aggregate boundary is accepted") + .expect("field is present"); + assert_eq!(decoded.len(), 2); + assert_eq!( + SetCookie::view(&boundary) + .expect("aggregate boundary view is accepted") + .expect("field is present") + .len(), + 2 + ); + + let over_limit = ValuesSource { + name: &FieldName::SetCookie, + values: vec![ + FieldValue::try_from(vec![b'a'; half]).expect("valid field value"), + FieldValue::try_from(vec![b'b'; half + 1]).expect("valid field value"), + ], + }; + assert_source_limit::(&over_limit); + + let list_over_limit = ValuesSource { + name: &FieldName::AcceptRanges, + values: repeated_value("bytes", 129), + }; + assert_source_limit::(&list_over_limit); + + let mut negotiation_boundary_bytes = b"text/plain".to_vec(); + negotiation_boundary_bytes.resize(65_536, b' '); + let negotiation_boundary = SingleSource { + name: &FieldName::Accept, + bytes: negotiation_boundary_bytes, + }; + assert!(Accept::view(&negotiation_boundary).expect("byte boundary is accepted").is_some()); + assert!(Accept::owned(&negotiation_boundary).expect("byte boundary is accepted").is_some()); + + let negotiation_over_limit = SingleSource { + name: &FieldName::Accept, + bytes: vec![b'a'; 65_537], + }; + assert_source_limit::(&negotiation_over_limit); +} + +#[test] +fn custom_source_line_limit_accepts_boundary_and_rejects_overflow() { + let boundary = ValuesSource { + name: &FieldName::SetCookie, + values: repeated_value("a=1", 128), + }; + let decoded = SetCookie::owned(&boundary) + .expect("boundary is accepted") + .expect("field is present"); + assert_eq!(decoded.len(), 128); + assert_eq!( + SetCookie::view(&boundary) + .expect("boundary view is accepted") + .expect("field is present") + .len(), + 128 + ); + + let over_limit = ValuesSource { + name: &FieldName::SetCookie, + values: repeated_value("a=1", 129), + }; + assert_source_limit::(&over_limit); + + let negotiation_boundary = ValuesSource { + name: &FieldName::Accept, + values: repeated_value("text/plain", 128), + }; + assert!(Accept::view(&negotiation_boundary).expect("line boundary is accepted").is_some()); + assert!(Accept::owned(&negotiation_boundary).expect("line boundary is accepted").is_some()); + + let negotiation_over_limit = ValuesSource { + name: &FieldName::Accept, + values: repeated_value("text/plain", 129), + }; + assert_source_limit::(&negotiation_over_limit); +} + +#[test] +fn custom_source_list_item_limit_accepts_boundary_and_rejects_overflow() { + let boundary = SingleSource { + name: &FieldName::ReferrerPolicy, + bytes: comma_list("origin", 1_024), + }; + let decoded = ReferrerPolicy::owned(&boundary) + .expect("boundary is accepted") + .expect("field is present"); + assert_eq!(decoded.tokens().count(), 1_024); + assert_eq!( + ReferrerPolicy::view(&boundary) + .expect("boundary view is accepted") + .expect("field is present") + .tokens() + .count(), + 1_024 + ); + + let over_limit = SingleSource { + name: &FieldName::ReferrerPolicy, + bytes: comma_list("origin", 1_025), + }; + assert_source_limit::(&over_limit); + + let over_limit_before_final_item = SingleSource { + name: &FieldName::ReferrerPolicy, + bytes: comma_list("origin", 1_026), + }; + assert_source_limit::(&over_limit_before_final_item); + + let negotiation_boundary = SingleSource { + name: &FieldName::Accept, + bytes: comma_list("text/plain", 1_024), + }; + assert!(Accept::view(&negotiation_boundary).expect("list boundary is accepted").is_some()); + assert!(Accept::owned(&negotiation_boundary).expect("list boundary is accepted").is_some()); + + let negotiation_over_limit = SingleSource { + name: &FieldName::Accept, + bytes: comma_list("text/plain", 1_025), + }; + assert_source_limit::(&negotiation_over_limit); + + let boundary = SingleSource { + name: &FieldName::IfMatch, + bytes: comma_list("\"tag\"", 1_024), + }; + assert!(IfMatch::owned(&boundary).expect("boundary is accepted").is_some()); + + let over_limit = SingleSource { + name: &FieldName::IfMatch, + bytes: comma_list("\"tag\"", 1_025), + }; + assert_source_limit::(&over_limit); + + let boundary = SingleSource { + name: &FieldName::IfMatch, + bytes: comma_list("\"tag\\\"", 1_024), + }; + assert!(IfMatch::owned(&boundary).expect("backslash boundary is accepted").is_some()); + + let over_limit = SingleSource { + name: &FieldName::IfMatch, + bytes: comma_list("\"tag\\\"", 1_025), + }; + assert_source_limit::(&over_limit); +} + +#[test] +fn raw_and_validated_list_sources_share_the_aggregate_item_budget() { + let first = comma_list("origin", 512); + for count in [512, 513] { + let second = comma_list("origin", count); + let values = [FieldValueRef::new(&first), FieldValueRef::new(&second)]; + let raw = BorrowedSource(&values); + let stored = ValuesSource { + name: &FieldName::ReferrerPolicy, + values: vec![FieldValue::from_bytes(&first).unwrap(), FieldValue::from_bytes(&second).unwrap()], + }; + if count == 512 { + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!(ReferrerPolicy::view_with(&raw, mode).unwrap().unwrap().tokens().count(), 1_024); + assert_eq!(ReferrerPolicy::owned_with(&raw, mode).unwrap().unwrap().tokens().count(), 1_024); + assert_eq!(ReferrerPolicy::view_with(&stored, mode).unwrap().unwrap().tokens().count(), 1_024); + assert_eq!(ReferrerPolicy::owned_with(&stored, mode).unwrap().unwrap().tokens().count(), 1_024); + } + } else { + assert_source_limit::(&raw); + assert_source_limit::(&stored); + } + } +} + +#[test] +fn singleton_list_item_limit_accepts_boundary_and_rejects_overflow() { + let range_boundary = SingleSource { + name: &FieldName::Range, + bytes: range_list(1_024), + }; + assert!(Range::view(&range_boundary).expect("Range boundary is accepted").is_some()); + assert!(Range::owned(&range_boundary).expect("Range boundary is accepted").is_some()); + + let range_over_limit = SingleSource { + name: &FieldName::Range, + bytes: range_list(1_025), + }; + assert_source_limit::(&range_over_limit); + + let hsts_boundary = SingleSource { + name: &FieldName::StrictTransportSecurity, + bytes: hsts_directives(1_024), + }; + assert!( + StrictTransportSecurity::view(&hsts_boundary) + .expect("HSTS boundary is accepted") + .is_some() + ); + assert!( + StrictTransportSecurity::owned(&hsts_boundary) + .expect("HSTS boundary is accepted") + .is_some() + ); + + let hsts_over_limit = SingleSource { + name: &FieldName::StrictTransportSecurity, + bytes: hsts_directives(1_025), + }; + assert_source_limit::(&hsts_over_limit); +} + +#[test] +fn bounded_opaque_single_value_does_not_treat_delimiters_as_list_items() { + let mut bytes = b"https://example.com/".to_vec(); + bytes.extend(iter::repeat_n(b',', 1_025)); + bytes.extend(iter::repeat_n(b';', 1_025)); + let source = SingleSource { + name: &FieldName::Location, + bytes, + }; + + assert!(Location::owned(&source).expect("valid URI decodes").is_some()); +} + +#[test] +fn raw_borrowed_sources_obey_aggregate_byte_and_line_boundaries() { + let first = vec![b'a'; 32_768]; + let mut second = vec![b'b'; 32_768]; + { + let values = [FieldValueRef::new(&first), FieldValueRef::new(&second)]; + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!(SetCookie::view_with(&BorrowedSource(&values), mode).unwrap().unwrap().len(), 2); + assert_eq!(SetCookie::owned_with(&BorrowedSource(&values), mode).unwrap().unwrap().len(), 2); + } + } + second.push(b'b'); + let values = [FieldValueRef::new(&first), FieldValueRef::new(&second)]; + assert_source_limit::(&BorrowedSource(&values)); + + let mut values = vec![FieldValueRef::new(b"a=1"); 128]; + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!(SetCookie::view_with(&BorrowedSource(&values), mode).unwrap().unwrap().len(), 128); + assert_eq!(SetCookie::owned_with(&BorrowedSource(&values), mode).unwrap().unwrap().len(), 128); + } + values.push(FieldValueRef::new(b"a=1")); + assert_source_limit::(&BorrowedSource(&values)); +} + +#[test] +fn admission_limits_do_not_reclassify_invalid_field_bytes() { + let oversized = vec![b'a'; 65_537]; + for invalid in [b"\r".as_slice(), b"\n", b"\0", b"\x1f", b"\x7f"] { + let values = [FieldValueRef::new(invalid)]; + assert_decode_error::(&BorrowedSource(&values), DecodeErrorKind::InvalidSyntax); + assert_decode_error::(&BorrowedSource(&values), DecodeErrorKind::InvalidSyntax); + + let values = [FieldValueRef::new(invalid), FieldValueRef::new(&oversized)]; + assert_decode_error::(&BorrowedSource(&values), DecodeErrorKind::InvalidSyntax); + assert_decode_error::(&BorrowedSource(&values), DecodeErrorKind::InvalidSyntax); + + let values = [FieldValueRef::new(&oversized), FieldValueRef::new(invalid)]; + assert_source_limit::(&BorrowedSource(&values)); + assert_source_limit::(&BorrowedSource(&values)); + } + + let values = vec![FieldValueRef::new(b"\r"); 129]; + assert_source_limit::(&BorrowedSource(&values)); + assert_source_limit::(&BorrowedSource(&values)); + + let mut source = SingleSource { + name: &FieldName::SetCookie, + bytes: vec![b'a'; 65_536], + }; + source.bytes[0] = b'\r'; + assert_decode_error::(&source, DecodeErrorKind::InvalidSyntax); + source.bytes.push(b'a'); + assert_source_limit::(&source); + + source.name = &FieldName::AcceptRanges; + assert_source_limit::(&source); + source.bytes.truncate(65_536); + assert_decode_error::(&source, DecodeErrorKind::InvalidSyntax); +} + +#[test] +fn delimited_iteration_distinguishes_item_admission_from_quote_errors() { + for (tail, kind, index) in [ + ("tail", DecodeErrorKind::SourceLimitExceeded, None), + ("\"unterminated", DecodeErrorKind::UnterminatedQuote, Some(0)), + ] { + let mut bytes = comma_list("origin", 1_024); + bytes.push(b','); + bytes.extend_from_slice(tail.as_bytes()); + let lines = FieldLines::single(&FieldName::ReferrerPolicy, &bytes); + let mut items = lines.comma_items(); + for _ in 0..1_024 { + assert_eq!(items.next(), Some(Ok(b"origin".as_slice()))); + } + let error = items.next().unwrap().unwrap_err(); + assert_eq!(error.kind(), kind); + assert_eq!(error.header(), &FieldName::ReferrerPolicy); + assert_eq!(error.value_index(), index); + assert_eq!(items.next(), None); + assert_eq!(items.next(), None); + } +} + +#[test] +fn handwritten_direct_decoders_preflight_custom_source_budgets() { + let credentials = SingleSource { + name: &FieldName::AccessControlAllowCredentials, + bytes: padded_value(b"true"), + }; + assert_source_limit::(&credentials); + + let max_age = SingleSource { + name: &FieldName::AccessControlMaxAge, + bytes: padded_value(b"1"), + }; + assert_source_limit::(&max_age); + + let request_method = SingleSource { + name: &FieldName::AccessControlRequestMethod, + bytes: padded_value(b"GET"), + }; + assert_source_limit::(&request_method); + + let origin = SingleSource { + name: &FieldName::AccessControlAllowOrigin, + bytes: padded_value(b"https://example.com"), + }; + assert_source_limit::(&origin); + + let policy = SingleSource { + name: &FieldName::ContentSecurityPolicy, + bytes: vec![b'a'; 65_537], + }; + assert_source_limit::(&policy); + + let recognized = ValuesSource { + name: &FieldName::ReferrerPolicy, + values: repeated_value("origin", 129), + }; + assert_source_limit::(&recognized); +} + +struct LimitedSink { + values: Vec, +} + +impl FieldSource for LimitedSink { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == &FieldName::SetCookie) + .then(|| FieldLines::from_slice(name, &self.values)) + .flatten() + } +} + +impl FieldSink for LimitedSink { + fn set_values(&mut self, _name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + self.values = values.into_iter().collect(); + Ok(()) + } + + fn append_values(&mut self, _name: &'static FieldName, values: EncodedValues) -> Result<(), InsertError> { + if self.values.len().checked_add(values.len()).is_none_or(|length| length > 128) { + return Err(InsertError::new(InsertErrorKind::CapacityExceeded)); + } + self.values.extend(values); + Ok(()) + } + + fn remove_values(&mut self, _name: &'static FieldName) { + self.values.clear(); + } +} + +#[test] +fn appending_to_an_over_limit_custom_sink_returns_insert_error() { + let mut sink = LimitedSink { + values: repeated_value("a=1", 129), + }; + let original_len = sink.values.len(); + + assert_eq!( + sink.append_encoded(&FieldName::SetCookie, FieldValue::from_static("b=2")), + Err(InsertError::new(InsertErrorKind::CapacityExceeded)) + ); + assert_eq!(sink.values.len(), original_len); +} + +#[cfg(feature = "http")] +#[test] +fn validated_http_map_values_are_not_subject_to_custom_source_limits() { + let mut map = http::HeaderMap::new(); + map.insert( + http::header::USER_AGENT, + http::HeaderValue::from_bytes(&vec![b'a'; 65_537]).expect("valid long field value"), + ); + for _ in 0..129 { + map.append(http::header::SET_COOKIE, http::HeaderValue::from_static("a=1")); + } + map.insert( + http::header::REFERRER_POLICY, + http::HeaderValue::from_bytes(&comma_list("origin", 1_025)).expect("valid long list"), + ); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + assert_eq!(UserAgent::view_with(&map, mode).unwrap().unwrap().as_bytes().len(), 65_537); + assert_eq!(UserAgent::owned_with(&map, mode).unwrap().unwrap().as_bytes().len(), 65_537); + assert_eq!(SetCookie::view_with(&map, mode).unwrap().unwrap().len(), 129); + assert_eq!(SetCookie::owned_with(&map, mode).unwrap().unwrap().len(), 129); + assert_eq!(ReferrerPolicy::view_with(&map, mode).unwrap().unwrap().tokens().count(), 1_025); + assert_eq!(ReferrerPolicy::owned_with(&map, mode).unwrap().unwrap().tokens().count(), 1_025); + } +} diff --git a/crates/http_headers/tests/support_api.rs b/crates/http_headers/tests/support_api.rs new file mode 100644 index 000000000..4ada57cd1 --- /dev/null +++ b/crates/http_headers/tests/support_api.rs @@ -0,0 +1,80 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Public support API integration coverage. + +#![cfg(feature = "headers-all")] + +#[cfg(feature = "http")] +use http_headers::headers::{LocationOwned, UserAgent, UserAgentOwned}; +use http_headers::sink::{EncodedValues, InsertError, InsertErrorKind}; +use http_headers::{DecodeError, DecodeErrorKind, FieldName, FieldValue}; + +#[test] +fn encoded_values_and_errors_expose_stable_public_behavior() { + let adopted = EncodedValues::from_vec(vec![FieldValue::from_static("a")]); + assert_eq!(adopted.len(), 1); + assert!(!adopted.is_empty()); + + let encoded: EncodedValues = [FieldValue::from_static("first"), FieldValue::from_static("second")] + .into_iter() + .collect(); + assert_eq!(encoded.len(), 2); + assert!(!format!("{encoded:?}").contains("first")); + assert_eq!( + encoded.into_iter().collect::>(), + [FieldValue::from_static("first"), FieldValue::from_static("second"),] + ); + + let cases = [ + (DecodeErrorKind::MissingValue, "missing value"), + (DecodeErrorKind::UnexpectedMultipleValues, "unexpected multiple values"), + (DecodeErrorKind::InvalidSyntax, "invalid syntax"), + (DecodeErrorKind::InvalidUtf8, "invalid UTF-8"), + (DecodeErrorKind::InvalidToken, "invalid token"), + (DecodeErrorKind::InvalidNumber, "invalid number"), + (DecodeErrorKind::UnterminatedQuote, "unterminated quoted string"), + ]; + for (kind, expected) in cases { + assert_eq!(kind.to_string(), expected); + } + let error = DecodeError::new(&FieldName::ContentType, DecodeErrorKind::InvalidSyntax).at_value(2); + assert_eq!(error.value_index(), Some(2)); + assert_eq!(error.to_string(), "invalid content-type header: invalid syntax at value 2"); + assert_eq!(InsertError::new(InsertErrorKind::InvalidValue).to_string(), "invalid field value"); +} + +#[test] +#[cfg(feature = "http")] +fn location_and_user_agent_use_public_construction_and_decode() { + for value in [ + "https://example.com/people", + "../people/tim?tab=1#profile", + "#profile", + "", + "https://[v1.address]/", + ] { + LocationOwned::try_from(value).expect("valid URI-reference"); + } + for value in ["/bad%2", "/café", "http://[::1", "/a[b]", "1abc:def"] { + assert_eq!( + LocationOwned::try_from(value).expect_err("invalid URI-reference must fail").kind(), + DecodeErrorKind::InvalidSyntax + ); + } + let signed = LocationOwned::try_from("/download?signature=secret").expect("valid URI-reference"); + let debug = format!("{signed:?}"); + assert!(!debug.contains("secret")); + assert!(debug.contains("redacted")); + + let mut map = http::HeaderMap::new(); + UserAgent::insert(&mut map, UserAgentOwned::try_from("example-client/1.0").expect("valid user agent")).expect("map has capacity"); + assert_eq!( + UserAgent::view(&map).expect("valid header").expect("present").as_str(), + Ok("example-client/1.0") + ); + assert_eq!( + UserAgentOwned::try_from(" \t").expect_err("blank user agent must fail").kind(), + DecodeErrorKind::InvalidSyntax + ); +} diff --git a/crates/http_headers/tests/typed_tokens.rs b/crates/http_headers/tests/typed_tokens.rs new file mode 100644 index 000000000..fa444c50d --- /dev/null +++ b/crates/http_headers/tests/typed_tokens.rs @@ -0,0 +1,222 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Semantic method and selection-field APIs across owned and borrowed headers. + +#![cfg(feature = "headers-negotiation")] + +use std::collections::HashSet; + +use http_headers::headers::{Allow, AllowOwned, FieldNameView, MethodView, Vary, VaryEntryView, VaryOwned}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeErrorKind, FieldName, FieldSensitivity, FieldValue, FieldValueRef, InvalidFieldName}; + +struct Source<'a> { + name: &'static FieldName, + values: &'a [FieldValue], +} + +impl FieldSource for Source<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + (name == self.name).then(|| FieldLines::from_slice(name, self.values)).flatten() + } +} + +#[test] +fn allow_preserves_case_extensions_order_and_empty_members() { + let mut first = FieldValue::from_static("GET, , get, *"); + first.set_sensitivity(FieldSensitivity::Sensitive); + let values = [first, FieldValue::from_static("CUSTOM, GET")]; + let source = Source { + name: &FieldName::Allow, + values: &values, + }; + let view = Allow::view(&source).unwrap().unwrap(); + let owned = Allow::owned(&source).unwrap().unwrap(); + let expected = ["GET", "get", "*", "CUSTOM", "GET"]; + assert_eq!(view.methods().map(MethodView::as_str).collect::>(), expected); + assert_eq!(owned.methods().map(MethodView::as_str).collect::>(), expected); + assert_eq!( + owned.values().map(FieldValueRef::as_bytes).collect::>(), + [b"GET, , get, *".as_slice(), b"CUSTOM, GET"] + ); + assert!(owned.values().next().unwrap().is_sensitive()); + let methods: HashSet<_> = view.methods().collect(); + assert_eq!(methods.len(), 4); + assert!(methods.contains(&MethodView::GET)); + assert_ne!(MethodView::GET, MethodView::new("get").unwrap()); + assert_eq!(MethodView::new("*").unwrap().as_str(), "*"); +} + +#[test] +fn method_construction_and_typed_allow_construction() { + let standards = [ + MethodView::GET, + MethodView::HEAD, + MethodView::POST, + MethodView::PUT, + MethodView::DELETE, + MethodView::CONNECT, + MethodView::OPTIONS, + MethodView::TRACE, + MethodView::PATCH, + ]; + let owned = AllowOwned::from_methods(standards); + assert_eq!(owned.methods().collect::>(), standards); + assert_eq!( + owned.values().next().unwrap().as_bytes(), + b"GET, HEAD, POST, PUT, DELETE, CONNECT, OPTIONS, TRACE, PATCH" + ); + let extension = MethodView::try_from("custom-method").unwrap(); + assert_eq!(extension.as_bytes(), b"custom-method"); + assert_eq!(extension.as_ref(), "custom-method"); + assert_eq!(extension.to_string(), "custom-method"); + assert_eq!(format!("{extension:?}"), "MethodView(\"custom-method\")"); + let empty = AllowOwned::from_methods([]); + assert_eq!(empty.methods().count(), 0); + assert_eq!(empty.values().len(), 1); + assert_eq!(empty.values().next().unwrap().as_bytes(), b""); + for invalid in ["", "bad method", "GET,POST", "méthode", "GET\n"] { + let error = MethodView::new(invalid).unwrap_err(); + assert_eq!(error.to_string(), "invalid HTTP method"); + } +} + +#[test] +fn vary_names_compare_and_hash_without_losing_wire_spelling() { + let values = [ + FieldValue::from_static("Accept-Encoding, , Origin"), + FieldValue::from_static("accept-encoding, X-Custom"), + ]; + let source = Source { + name: &FieldName::Vary, + values: &values, + }; + let view = Vary::view(&source).unwrap().unwrap(); + let owned = Vary::owned(&source).unwrap().unwrap(); + assert!(!view.contains_wildcard()); + assert!(!owned.contains_wildcard()); + assert_eq!(view.entries().collect::>(), owned.entries().collect::>()); + assert_eq!( + view.entries().map(VaryEntryView::as_str).collect::>(), + ["Accept-Encoding", "Origin", "accept-encoding", "X-Custom"] + ); + let names: HashSet<_> = view.entries().map(|entry| entry.field_name().unwrap()).collect(); + assert_eq!(names.len(), 3); + assert!(names.contains(&FieldNameView::new("ACCEPT-ENCODING").unwrap())); + assert_eq!( + FieldNameView::new("X-Custom").unwrap().try_to_field_name().unwrap().as_str(), + "x-custom" + ); + assert_eq!( + FieldNameView::new("ACCEPT").unwrap().try_to_field_name().unwrap(), + FieldName::Accept + ); + assert_ne!(VaryOwned::try_from("Origin").unwrap(), VaryOwned::try_from("origin").unwrap()); +} + +#[test] +fn vary_wildcards_remain_visible_across_field_lines() { + for wire in ["*", "*, Origin", "Origin, *", "Origin, *, ACCEPT"] { + let values = [FieldValue::from_static("X-First"), FieldValue::from_str(wire).unwrap()]; + let source = Source { + name: &FieldName::Vary, + values: &values, + }; + let view = Vary::view(&source).unwrap().unwrap(); + let owned = Vary::owned(&source).unwrap().unwrap(); + assert!(view.contains_wildcard()); + assert!(owned.contains_wildcard()); + assert_eq!(view.entries().filter(|entry| entry.is_wildcard()).count(), 1); + assert_eq!(view.entries().find(|entry| entry.is_wildcard()).unwrap().field_name(), None); + assert_eq!(owned.values().nth(1).unwrap().as_bytes(), wire.as_bytes()); + } + let name = FieldNameView::new("Origin").unwrap(); + let rebuilt = VaryOwned::from_entries([VaryEntryView::from_field_name(name), VaryEntryView::WILDCARD]); + assert_eq!(rebuilt.values().next().unwrap().as_bytes(), b"Origin, *"); + assert_eq!(VaryEntryView::WILDCARD.to_string(), "*"); + assert_eq!(VaryEntryView::from_field_name(name).to_string(), "Origin"); + assert!(VaryEntryView::from_field_name(FieldNameView::new("*").unwrap()).is_wildcard()); + assert!(VaryOwned::wildcard().contains_wildcard()); + let names = VaryOwned::from_field_names([name, name]); + assert_eq!(names.entries().count(), 2); + let empty = VaryOwned::from_field_names([]); + assert_eq!(empty.entries().count(), 0); + assert!(!empty.contains_wildcard()); +} + +#[test] +fn vary_contains_wildcard_requires_a_complete_wildcard_member() { + for wire in ["Origin", "Origin, X-*", "**", " ,\t,"] { + let values = [FieldValue::from_str(wire).unwrap()]; + let source = Source { + name: &FieldName::Vary, + values: &values, + }; + assert!(!Vary::view(&source).unwrap().unwrap().contains_wildcard()); + assert!(!Vary::owned(&source).unwrap().unwrap().contains_wildcard()); + assert_eq!(values[0].as_bytes(), wire.as_bytes()); + } +} + +#[test] +fn field_name_views_keep_validation_and_materialization_separate() { + let name = FieldNameView::try_from("X-Foo").unwrap(); + assert!(name.eq_ignore_ascii_case("x-foo")); + assert!(!name.eq_ignore_ascii_case("x-bar")); + assert_eq!(name.as_ref(), "X-Foo"); + assert_eq!(name.as_bytes(), b"X-Foo"); + assert_eq!(name.to_string(), "X-Foo"); + assert_eq!(format!("{name:?}"), "FieldNameView(\"X-Foo\")"); + for invalid in ["", "bad name", "name:", "naïve", "x\n"] { + assert_eq!(FieldNameView::new(invalid).unwrap_err(), InvalidFieldName); + } + let long = "x".repeat(65_536); + let name = FieldNameView::new(&long).unwrap(); + assert_eq!(name.as_str(), long); + assert_eq!(name.try_to_field_name().unwrap_err(), InvalidFieldName); + let value = VaryOwned::try_from(long.as_str()).unwrap(); + assert_eq!(value.entries().next().unwrap().field_name().unwrap(), name); +} + +#[test] +fn malformed_members_fail_decoding_instead_of_disappearing() { + for name in [&FieldName::Allow, &FieldName::Vary] { + let values = [FieldValue::from_static("GET"), FieldValue::from_static("bad token")]; + let source = Source { name, values: &values }; + let error = if name == &FieldName::Allow { + Allow::view(&source).unwrap_err() + } else { + Vary::view(&source).unwrap_err() + }; + assert_eq!(error.kind(), DecodeErrorKind::InvalidToken); + } +} + +#[cfg(feature = "http")] +#[test] +fn http_conversions_and_round_trips_preserve_typed_members() { + use http::HeaderMap; + + assert_eq!(MethodView::GET.try_to_method().unwrap(), http::Method::GET); + assert_eq!(MethodView::new("CUSTOM").unwrap().try_to_method().unwrap().as_str(), "CUSTOM"); + assert_eq!( + FieldNameView::new("X-Custom").unwrap().try_to_http_header_name().unwrap().as_str(), + "x-custom" + ); + let mut map = HeaderMap::new(); + Allow::insert(&mut map, AllowOwned::from_methods([MethodView::GET, MethodView::HEAD])).unwrap(); + assert_eq!( + Allow::view(&map).unwrap().unwrap().methods().collect::>(), + [MethodView::GET, MethodView::HEAD] + ); + Vary::insert( + &mut map, + VaryOwned::from_field_names([FieldNameView::new("Accept-Encoding").unwrap()]), + ) + .unwrap(); + assert_eq!( + Vary::view(&map).unwrap().unwrap().entries().next().unwrap().as_str(), + "Accept-Encoding" + ); +} diff --git a/crates/http_headers/tests/websocket_source_limits.rs b/crates/http_headers/tests/websocket_source_limits.rs new file mode 100644 index 000000000..8a22aa53b --- /dev/null +++ b/crates/http_headers/tests/websocket_source_limits.rs @@ -0,0 +1,157 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Source limits and error precedence across WebSocket fallback paths. + +#![cfg(feature = "headers-websocket")] + +use std::iter; + +use http_headers::headers::{SecWebSocketExtensions, SecWebSocketProtocol, SecWebSocketVersion}; +use http_headers::source::{FieldLines, FieldSource}; +use http_headers::{DecodeError, DecodeErrorKind, DecodeMode, Field, FieldName, FieldValueRef}; + +struct RawSource<'a> { + name: &'static FieldName, + values: Vec>, + single: bool, +} + +impl FieldSource for RawSource<'_> { + fn lines(&self, name: &'static FieldName) -> Option> { + if name != self.name { + return None; + } + if self.single { + Some(FieldLines::single(name, self.values[0].as_bytes())) + } else { + FieldLines::from_borrowed(name, &self.values) + } + } +} + +fn check(values: &[&[u8]], view: Result, owned: Result) { + // Miri checks every real budget with both ownerships, without repeating the + // large inputs across source representations and mode-independent decoders. + let budget_sample = cfg!(miri) && (values.len() >= 128 || values.iter().any(|value| value.len() >= 1_024)); + let mut source = RawSource { + name: F::name(), + values: values.iter().map(|bytes| FieldValueRef::new(bytes)).collect(), + single: false, + }; + for single in [false, true] { + if single && (values.len() != 1 || budget_sample) { + continue; + } + source.single = single; + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + if budget_sample && mode == DecodeMode::Relaxed { + continue; + } + assert_eq!( + F::view_with(&source, mode).map(|value| value.is_some()), + view, + "{}: borrowed, single={single}, mode={mode:?}", + F::name() + ); + assert_eq!( + F::owned_with(&source, mode).map(|value| value.is_some()), + owned, + "{}: owned, single={single}, mode={mode:?}", + F::name() + ); + } + } +} + +fn check_limits(item: &[u8], members_at_limit: usize) { + let invalid = Err(DecodeError::new(F::name(), DecodeErrorKind::InvalidSyntax)); + let limit = Err(DecodeError::new(F::name(), DecodeErrorKind::SourceLimitExceeded)); + let mut bytes = item.to_vec(); + bytes.resize(65_536, b' '); + check::(&[&bytes], Ok(true), Ok(true)); + bytes.push(b' '); + check::(&[&bytes], limit, limit); + + let mut lines = vec![item; 128]; + check::(&lines, Ok(true), Ok(true)); + lines.push(item); + check::(&lines, limit, limit); + + let mut members = iter::repeat_n(item, members_at_limit).collect::>().join(&b','); + check::(&[&members], Ok(true), Ok(true)); + members.push(b','); + members.extend_from_slice(item); + check::(&[&members], limit, limit); + + // Source validation precedes even a malformed first grammar item. + check::(&[b"/", &bytes], limit, limit); + check::(&[b"/", &members], limit, limit); + for invalid_bytes in [b"\r".as_slice(), b"\n", b"\0", b"\x1f", b"\x7f"] { + check::(&[invalid_bytes], invalid, invalid); + check::(&[b"\"unterminated", invalid_bytes], invalid, invalid); + } +} + +#[test] +fn websocket_custom_source_limits_cover_bare_and_quoted_paths() { + check_limits::(b"13", 1_024); + check_limits::(b"chat", 1_024); + check_limits::(b"permessage-deflate", 1_024); + // Each parameter adds a semicolon; the initial extension also consumes an item. + check_limits::(b"x; p=\"ab\"", 1_023); +} + +#[test] +fn websocket_extension_parameter_budget_covers_quoted_values() { + let mut bytes = b"x".to_vec(); + for _ in 1..1_024 { + bytes.extend_from_slice(b";p=\"a\\b\""); + } + check::(&[&bytes], Ok(true), Ok(true)); + bytes.extend_from_slice(b";p=\"a\\b\""); + let limit = Err(DecodeError::new( + &FieldName::SecWebSocketExtensions, + DecodeErrorKind::SourceLimitExceeded, + )); + check::(&[&bytes], limit, limit); +} + +#[test] +fn websocket_fallback_errors_preserve_kind_and_physical_line_index() { + let version_syntax = Err(DecodeError::new(&FieldName::SecWebSocketVersion, DecodeErrorKind::InvalidSyntax)); + let version_quote = Err(DecodeError::new(&FieldName::SecWebSocketVersion, DecodeErrorKind::UnterminatedQuote).at_value(1)); + check::(&[b"256", b"\"13"], version_syntax, version_syntax); + check::(&[b"13", b"\"13"], version_quote, version_quote); + check::(&[b"13", b"\"13", b"13\""], version_quote, version_quote); + check::(&[b"13", b"\"13\""], version_syntax, version_syntax); + + let protocol_token = Err(DecodeError::new(&FieldName::SecWebSocketProtocol, DecodeErrorKind::InvalidToken)); + let protocol_quote = Err(DecodeError::new(&FieldName::SecWebSocketProtocol, DecodeErrorKind::UnterminatedQuote).at_value(1)); + check::(&[b"/", b"\"chat"], protocol_token, protocol_token); + check::(&[b"chat", b"\"chat"], protocol_quote, protocol_token); + check::(&[b"chat", b"\"chat", b"chat\""], protocol_quote, protocol_token); + + let extension_token = Err(DecodeError::new(&FieldName::SecWebSocketExtensions, DecodeErrorKind::InvalidToken)); + let extension_quote = Err(DecodeError::new(&FieldName::SecWebSocketExtensions, DecodeErrorKind::UnterminatedQuote).at_value(1)); + check::(&[b"/", b"x;p=\"a"], extension_token, extension_token); + check::(&[b"x", b"x;p=\"a"], extension_quote, extension_quote); + check::(&[b"x;p=\"ab\"", b"x;p=\"a"], extension_quote, extension_quote); + check::(&[b"x", b"\"x"], extension_quote, extension_token); +} + +#[cfg(feature = "http")] +#[test] +fn http_websocket_versions_are_not_subject_to_custom_source_limits() { + let mut map = http::HeaderMap::new(); + for _ in 0..128 { + map.append(http::header::SEC_WEBSOCKET_VERSION, http::HeaderValue::from_static("13")); + } + map.append(http::header::SEC_WEBSOCKET_VERSION, http::HeaderValue::from_static("255")); + for mode in [DecodeMode::Strict, DecodeMode::Relaxed] { + let borrowed = SecWebSocketVersion::view_with(&map, mode).unwrap().unwrap(); + let owned = SecWebSocketVersion::owned_with(&map, mode).unwrap().unwrap(); + assert_eq!(borrowed.versions().collect::>(), [13, 255]); + assert_eq!(owned.versions().collect::>(), [13, 255]); + } +} diff --git a/crates/http_headers_simd/CHANGELOG.md b/crates/http_headers_simd/CHANGELOG.md new file mode 100644 index 000000000..6a8522106 --- /dev/null +++ b/crates/http_headers_simd/CHANGELOG.md @@ -0,0 +1,5 @@ +# Changelog + +## [0.1.0] + +- Initial integration into the Oxidizer workspace. diff --git a/crates/http_headers_simd/Cargo.toml b/crates/http_headers_simd/Cargo.toml new file mode 100644 index 000000000..bf45be4b0 --- /dev/null +++ b/crates/http_headers_simd/Cargo.toml @@ -0,0 +1,43 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +[package] +name = "http_headers_simd" +description = "SIMD implementation details for http_headers. Do not depend on this crate directly; use http_headers." +version = "0.1.0" +readme = "README.md" +edition.workspace = true +rust-version.workspace = true +authors.workspace = true +license.workspace = true +homepage.workspace = true +include.workspace = true +repository = "https://github.com/microsoft/oxidizer/tree/main/crates/http_headers_simd" +keywords = ["http", "headers", "simd", "parsing"] +categories = ["web-programming"] + +[package.metadata.cargo_check_external_types] +allowed_external_types = [] + +[package.metadata.docs.rs] +all-features = true + +[features] +default = ["std"] +std = [] +benchmarking = [] +test-util = [] + +[dev-dependencies] +bolero = { workspace = true, features = ["std"] } +criterion = { workspace = true } +metabench = { workspace = true } + +[[bench]] +name = "http_headers_simd_no_std_dispatch" +harness = false + +# >>> anvil-managed: anvil-lints +[lints] +workspace = true +# <<< anvil-managed: anvil-lints diff --git a/crates/http_headers_simd/README.md b/crates/http_headers_simd/README.md new file mode 100644 index 000000000..3876bbd82 --- /dev/null +++ b/crates/http_headers_simd/README.md @@ -0,0 +1,37 @@ +
+ Http Headers Simd Logo + +# Http Headers Simd + +[![crate.io](https://img.shields.io/crates/v/http_headers_simd.svg)](https://crates.io/crates/http_headers_simd) +[![docs.rs](https://docs.rs/http_headers_simd/badge.svg)](https://docs.rs/http_headers_simd) +[![MSRV](https://img.shields.io/crates/msrv/http_headers_simd)](https://crates.io/crates/http_headers_simd) +[![CI](https://github.com/microsoft/oxidizer/actions/workflows/anvil-pr.yml/badge.svg)](https://github.com/microsoft/oxidizer/actions/workflows/anvil-pr.yml) +[![Coverage](https://codecov.io/gh/microsoft/oxidizer/graph/badge.svg?token=FCUG0EL5TI)](https://codecov.io/gh/microsoft/oxidizer) +[![License](https://img.shields.io/badge/license-MIT-blue.svg)](https://github.com/microsoft/oxidizer/blob/main/LICENSE) +This crate was developed as part of the Oxidizer project + +
+ +SIMD implementation details for the +[`http_headers`][__link0] crate. + +**Do not depend on this crate directly.** Use `http_headers` instead. + +The default `std` feature enables runtime CPU-feature detection. +With default features disabled, x86 and x86-64 use compile-time features +plus cached runtime `CPUID` detection for SSSE3 and SSE4.2 when those features +are not enabled at compile time. SSE2 is guaranteed on x86-64 and requires +compile-time support on x86. `AArch64` NEON availability follows compile-time +target features. Unavailable accelerated paths fall back to scalar code. + +The `benchmarking` and `test-util` features expose unstable repository +instrumentation only. + + +
+ +This crate was developed as part of The Oxidizer Project. Browse this crate's source code. + + + [__link0]: https://docs.rs/http_headers diff --git a/crates/http_headers_simd/benches/http_headers_simd_no_std_dispatch.rs b/crates/http_headers_simd/benches/http_headers_simd_no_std_dispatch.rs new file mode 100644 index 000000000..09f6eea41 --- /dev/null +++ b/crates/http_headers_simd/benches/http_headers_simd_no_std_dispatch.rs @@ -0,0 +1,238 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Scanner crossover benchmarks for the `no_std` dispatch path. +//! +//! Run with: +//! `cargo bench -p http_headers_simd --no-default-features --bench http_headers_simd_no_std_dispatch`. + +use std::hint::black_box; + +use criterion::Criterion; + +const NO_STD_DISPATCH: &str = "http_headers_simd_no_std_dispatch/dispatch"; + +macro_rules! unary_case { + ($identity:ident, $name:ident, $benchmark_name:literal, $length:expr, $scanner:path, $output:ty) => { + #[metabench::benchmark($identity, NO_STD_DISPATCH, $benchmark_name)] + fn $name() -> $output { + let bytes = black_box([b'a'; $length]); + black_box($scanner(&bytes)) + } + }; +} + +macro_rules! equality_case { + ($identity:ident, $name:ident, $benchmark_name:literal, $length:expr) => { + #[metabench::benchmark($identity, NO_STD_DISPATCH, $benchmark_name)] + fn $name() -> bool { + let left = black_box([b'a'; $length]); + let right = black_box([b'A'; $length]); + black_box(http_headers_simd::eq_ignore_ascii_case(&left, &right)) + } + }; +} + +macro_rules! scanner_cases { + ( + $macro_name:ident, + $scanner:path, + $output:ty, + [ + $(($identity:ident, $name:ident, $benchmark_name:literal, $length:expr)),+ $(,)? + ] + ) => { + $( + $macro_name!($identity, $name, $benchmark_name, $length, $scanner, $output); + )+ + }; +} + +scanner_cases!( + unary_case, + http_headers_simd::is_token, + bool, + [ + (TOKEN_15, token_15, "token_15", 15), + (TOKEN_16, token_16, "token_16", 16), + (TOKEN_17, token_17, "token_17", 17), + (TOKEN_31, token_31, "token_31", 31), + (TOKEN_32, token_32, "token_32", 32), + (TOKEN_33, token_33, "token_33", 33), + (TOKEN_47, token_47, "token_47", 47), + (TOKEN_48, token_48, "token_48", 48), + ] +); +scanner_cases!( + unary_case, + http_headers_simd::is_token68, + bool, + [ + (TOKEN68_15, token68_15, "token68_15", 15), + (TOKEN68_16, token68_16, "token68_16", 16), + (TOKEN68_17, token68_17, "token68_17", 17), + (TOKEN68_31, token68_31, "token68_31", 31), + (TOKEN68_32, token68_32, "token68_32", 32), + (TOKEN68_33, token68_33, "token68_33", 33), + (TOKEN68_47, token68_47, "token68_47", 47), + (TOKEN68_48, token68_48, "token68_48", 48), + ] +); +scanner_cases!( + unary_case, + http_headers_simd::is_field_value, + bool, + [ + (FIELD_VALUE_15, field_value_15, "field_value_15", 15), + (FIELD_VALUE_16, field_value_16, "field_value_16", 16), + (FIELD_VALUE_17, field_value_17, "field_value_17", 17), + (FIELD_VALUE_31, field_value_31, "field_value_31", 31), + (FIELD_VALUE_32, field_value_32, "field_value_32", 32), + (FIELD_VALUE_33, field_value_33, "field_value_33", 33), + (FIELD_VALUE_47, field_value_47, "field_value_47", 47), + (FIELD_VALUE_48, field_value_48, "field_value_48", 48), + ] +); +scanner_cases!( + unary_case, + http_headers_simd::find_interesting, + Option, + [ + (INTERESTING_15, interesting_15, "interesting_15", 15), + (INTERESTING_16, interesting_16, "interesting_16", 16), + (INTERESTING_17, interesting_17, "interesting_17", 17), + (INTERESTING_31, interesting_31, "interesting_31", 31), + (INTERESTING_32, interesting_32, "interesting_32", 32), + (INTERESTING_33, interesting_33, "interesting_33", 33), + (INTERESTING_47, interesting_47, "interesting_47", 47), + (INTERESTING_48, interesting_48, "interesting_48", 48), + ] +); + +equality_case!(EQUALITY_15, equality_15, "equality_15", 15); +equality_case!(EQUALITY_16, equality_16, "equality_16", 16); +equality_case!(EQUALITY_17, equality_17, "equality_17", 17); +equality_case!(EQUALITY_31, equality_31, "equality_31", 31); +equality_case!(EQUALITY_32, equality_32, "equality_32", 32); +equality_case!(EQUALITY_33, equality_33, "equality_33", 33); +equality_case!(EQUALITY_47, equality_47, "equality_47", 47); +equality_case!(EQUALITY_48, equality_48, "equality_48", 48); + +#[metabench::benchmark(DISPATCH_512, NO_STD_DISPATCH, "dispatch_512")] +fn dispatch_512() -> usize { + let left = black_box([b'a'; 512]); + let right = black_box([b'A'; 512]); + usize::from(http_headers_simd::is_token(&left)) + + usize::from(http_headers_simd::is_token68(&left)) + + usize::from(http_headers_simd::is_field_value(&left)) + + usize::from(http_headers_simd::eq_ignore_ascii_case(&left, &right)) + + http_headers_simd::find_interesting(&left).unwrap_or(left.len()) +} + +macro_rules! criterion_cases { + ($group:ident, [$(($identity:ident, $name:ident)),+ $(,)?]) => { + $( + $group.bench_function($identity.benchmark_name(), |bencher| { + bencher.iter($name); + }); + )+ + }; +} + +fn criterion_benchmarks(criterion: &mut Criterion) { + let mut dispatch = criterion.benchmark_group(NO_STD_DISPATCH); + criterion_cases!( + dispatch, + [ + (TOKEN_15, token_15), + (TOKEN_16, token_16), + (TOKEN_17, token_17), + (TOKEN_31, token_31), + (TOKEN_32, token_32), + (TOKEN_33, token_33), + (TOKEN_47, token_47), + (TOKEN_48, token_48), + (TOKEN68_15, token68_15), + (TOKEN68_16, token68_16), + (TOKEN68_17, token68_17), + (TOKEN68_31, token68_31), + (TOKEN68_32, token68_32), + (TOKEN68_33, token68_33), + (TOKEN68_47, token68_47), + (TOKEN68_48, token68_48), + (FIELD_VALUE_15, field_value_15), + (FIELD_VALUE_16, field_value_16), + (FIELD_VALUE_17, field_value_17), + (FIELD_VALUE_31, field_value_31), + (FIELD_VALUE_32, field_value_32), + (FIELD_VALUE_33, field_value_33), + (FIELD_VALUE_47, field_value_47), + (FIELD_VALUE_48, field_value_48), + (EQUALITY_15, equality_15), + (EQUALITY_16, equality_16), + (EQUALITY_17, equality_17), + (EQUALITY_31, equality_31), + (EQUALITY_32, equality_32), + (EQUALITY_33, equality_33), + (EQUALITY_47, equality_47), + (EQUALITY_48, equality_48), + (INTERESTING_15, interesting_15), + (INTERESTING_16, interesting_16), + (INTERESTING_17, interesting_17), + (INTERESTING_31, interesting_31), + (INTERESTING_32, interesting_32), + (INTERESTING_33, interesting_33), + (INTERESTING_47, interesting_47), + (INTERESTING_48, interesting_48), + (DISPATCH_512, dispatch_512), + ] + ); + dispatch.finish(); +} + +metabench::main!( + criterion = criterion_benchmarks, + benchmarks = [ + TOKEN_15, + TOKEN_16, + TOKEN_17, + TOKEN_31, + TOKEN_32, + TOKEN_33, + TOKEN_47, + TOKEN_48, + TOKEN68_15, + TOKEN68_16, + TOKEN68_17, + TOKEN68_31, + TOKEN68_32, + TOKEN68_33, + TOKEN68_47, + TOKEN68_48, + FIELD_VALUE_15, + FIELD_VALUE_16, + FIELD_VALUE_17, + FIELD_VALUE_31, + FIELD_VALUE_32, + FIELD_VALUE_33, + FIELD_VALUE_47, + FIELD_VALUE_48, + EQUALITY_15, + EQUALITY_16, + EQUALITY_17, + EQUALITY_31, + EQUALITY_32, + EQUALITY_33, + EQUALITY_47, + EQUALITY_48, + INTERESTING_15, + INTERESTING_16, + INTERESTING_17, + INTERESTING_31, + INTERESTING_32, + INTERESTING_33, + INTERESTING_47, + INTERESTING_48, + DISPATCH_512, + ], +); diff --git a/crates/http_headers_simd/favicon.ico b/crates/http_headers_simd/favicon.ico new file mode 100644 index 000000000..5d7bd15c9 --- /dev/null +++ b/crates/http_headers_simd/favicon.ico @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:3d2a6be9c244b42877fa86aaef97468175d8ecf64066e0bb6d48ef912e5261ba +size 432254 diff --git a/crates/http_headers_simd/logo.png b/crates/http_headers_simd/logo.png new file mode 100644 index 000000000..1523b6f74 --- /dev/null +++ b/crates/http_headers_simd/logo.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b8fb949b8c6797dfa3489b4b3306977ed6d98f8c4da42a25171709c9c1285497 +size 114089 diff --git a/crates/http_headers_simd/src/__fuzz__/campaign.toml b/crates/http_headers_simd/src/__fuzz__/campaign.toml new file mode 100644 index 000000000..2015cfa52 --- /dev/null +++ b/crates/http_headers_simd/src/__fuzz__/campaign.toml @@ -0,0 +1,24 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +[campaign] +package = "http_headers_simd" +default_engine = "libfuzzer" +default_time = "60s" +default_max_input_length = 1024 +corpus_root = "crates/http_headers_simd/src/__fuzz__" + +[[target]] +name = "differential_properties" +test = "unit" +description = "Compare every dispatched scanner with its scalar reference over arbitrary bytes." + +[[target]] +name = "sse2_matches_scalar" +test = "unit" +description = "Compare the x86 backends with scalar references over arbitrary bytes." + +[[target]] +name = "neon_matches_scalar" +test = "unit" +description = "Compare the NEON backend with scalar references over arbitrary bytes." diff --git a/crates/http_headers_simd/src/__fuzz__/differential_properties/corpus/boundaries b/crates/http_headers_simd/src/__fuzz__/differential_properties/corpus/boundaries new file mode 100644 index 000000000..17d4eb5d4 --- /dev/null +++ b/crates/http_headers_simd/src/__fuzz__/differential_properties/corpus/boundaries @@ -0,0 +1 @@ +aaaaaaaaaaaaaaa,aaaaaaaaaaaaaaa; diff --git a/crates/http_headers_simd/src/__fuzz__/differential_properties/crashes/.gitkeep b/crates/http_headers_simd/src/__fuzz__/differential_properties/crashes/.gitkeep new file mode 100644 index 000000000..e69de29bb diff --git a/crates/http_headers_simd/src/__fuzz__/neon_matches_scalar/corpus/boundaries b/crates/http_headers_simd/src/__fuzz__/neon_matches_scalar/corpus/boundaries new file mode 100644 index 000000000..1b9ce4618 --- /dev/null +++ b/crates/http_headers_simd/src/__fuzz__/neon_matches_scalar/corpus/boundaries @@ -0,0 +1 @@ +aaaaaaaaaaaaaaa#aaaaaaaaaaaaaaa? diff --git a/crates/http_headers_simd/src/__fuzz__/neon_matches_scalar/crashes/.gitkeep b/crates/http_headers_simd/src/__fuzz__/neon_matches_scalar/crashes/.gitkeep new file mode 100644 index 000000000..e69de29bb diff --git a/crates/http_headers_simd/src/__fuzz__/sse2_matches_scalar/corpus/boundaries b/crates/http_headers_simd/src/__fuzz__/sse2_matches_scalar/corpus/boundaries new file mode 100644 index 000000000..41e1adffe --- /dev/null +++ b/crates/http_headers_simd/src/__fuzz__/sse2_matches_scalar/corpus/boundaries @@ -0,0 +1 @@ +aaaaaaaaaaaaaaa/aaaaaaaaaaaaaaa= diff --git a/crates/http_headers_simd/src/__fuzz__/sse2_matches_scalar/crashes/.gitkeep b/crates/http_headers_simd/src/__fuzz__/sse2_matches_scalar/crashes/.gitkeep new file mode 100644 index 000000000..e69de29bb diff --git a/crates/http_headers_simd/src/api.rs b/crates/http_headers_simd/src/api.rs new file mode 100644 index 000000000..ce885034a --- /dev/null +++ b/crates/http_headers_simd/src/api.rs @@ -0,0 +1,947 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Public byte-scanning API and backend identification. + +use core::str; + +use crate::dispatch; +use crate::list::{EmptyMembers, TokenListScan}; + +/// Returns `bytes` as text when it holds nothing but ASCII. +/// +/// The check reduces the high bits of the whole slice with OR, which has no +/// data-dependent branch and vectorizes, so a short header value settles in a +/// handful of instructions. The general UTF-8 validator instead pays for +/// pointer alignment and for the multi-byte sequence state machine even when +/// the input turns out to be plain ASCII. +/// +/// `None` means the value holds bytes above `0x7f`. Callers that must accept +/// those need [`core::str::from_utf8`], which this function never replaces for +/// correctness — only for speed on the ASCII shape that headers almost always +/// take. +/// +/// # Examples +/// +/// ``` +/// assert_eq!( +/// http_headers_simd::ascii_str(b"Content-Type"), +/// Some("Content-Type") +/// ); +/// assert_eq!(http_headers_simd::ascii_str(&[0xff]), None); +/// ``` +#[must_use] +#[inline] +pub fn ascii_str(bytes: &[u8]) -> Option<&str> { + all_ascii(bytes).then(|| { + // SAFETY: every byte is below `0x80`, so each one is a single-byte + // UTF-8 sequence encoding the code point of the same value. + unsafe { str::from_utf8_unchecked(bytes) } + }) +} + +/// Returns whether every byte is below `0x80`. +/// +/// Chunking keeps the reduction inside a fixed-width accumulator so the +/// compiler emits one vector OR per block instead of a per-byte compare and +/// branch. +fn all_ascii(bytes: &[u8]) -> bool { + const BLOCK: usize = 16; + + let mut chunks = bytes.chunks_exact(BLOCK); + let mut high = 0_u8; + for chunk in &mut chunks { + let mut block = 0_u8; + for &byte in chunk { + block |= byte; + } + high |= block; + } + for &byte in chunks.remainder() { + high |= byte; + } + high < 0x80 +} + +/// Returns whether `bytes` is a non-empty RFC 9110 HTTP token. +/// +/// # Examples +/// +/// ``` +/// assert!(http_headers_simd::is_token(b"content-type")); +/// assert!(!http_headers_simd::is_token(b"content type")); +/// ``` +#[must_use] +#[inline] +pub fn is_token(bytes: &[u8]) -> bool { + dispatch::is_token(bytes) +} + +/// Returns whether `bytes` is an RFC 9110 `token68` value. +/// +/// A value contains one or more alphanumeric or `-._~+/` bytes followed only +/// by optional `=` padding. +/// +/// # Examples +/// +/// ``` +/// assert!(http_headers_simd::is_token68( +/// b"QWxhZGRpbjpvcGVuIHNlc2FtZQ==" +/// )); +/// assert!(!http_headers_simd::is_token68(b"=padding-first")); +/// ``` +#[must_use] +#[inline] +pub fn is_token68(bytes: &[u8]) -> bool { + dispatch::is_token68(bytes) +} + +/// Returns whether every byte is permitted in an HTTP field value. +/// +/// This accepts SP, HTAB, visible ASCII, and `obs-text` (`0x80..=0xff`). +/// Empty values are valid. +/// +/// # Examples +/// +/// ``` +/// assert!(http_headers_simd::is_field_value( +/// b"text/plain; charset=utf-8" +/// )); +/// assert!(!http_headers_simd::is_field_value(b"line\nbreak")); +/// ``` +#[must_use] +#[inline] +pub fn is_field_value(bytes: &[u8]) -> bool { + dispatch::is_field_value(bytes) +} + +/// Compares two byte strings using ASCII case-insensitive equality. +/// +/// Non-ASCII bytes compare exactly and are never case-folded. +/// +/// # Examples +/// +/// ``` +/// assert!(http_headers_simd::eq_ignore_ascii_case(b"gzip", b"GZIP")); +/// assert!(!http_headers_simd::eq_ignore_ascii_case(b"gzip", b"br")); +/// ``` +#[must_use] +#[inline] +pub fn eq_ignore_ascii_case(left: &[u8], right: &[u8]) -> bool { + dispatch::eq_ignore_ascii_case(left, right) +} + +/// Recognizes a simple origin-relative URI-reference. +/// +/// The accepted subset is a `path-absolute` optionally followed by a query and +/// a fragment, spelled entirely with bytes that carry no escaping or structure: +/// `unreserved`, `sub-delims`, `:`, `@`, `/`, `?`, and at most one `#`. A +/// leading `//` is rejected because it introduces an authority. +/// +/// A `false` result means "not known to be valid", not "invalid": percent +/// escapes, schemes, and authorities all leave the subset. Callers must run a +/// full parse in that case. +/// +/// # Examples +/// +/// ``` +/// assert!(http_headers_simd::is_simple_uri_path( +/// b"/docs?page=2#results" +/// )); +/// assert!(!http_headers_simd::is_simple_uri_path( +/// b"//example.com/docs" +/// )); +/// ``` +#[must_use] +#[inline] +pub fn is_simple_uri_path(bytes: &[u8]) -> bool { + dispatch::is_simple_uri_path(bytes) +} + +/// Returns a recognized simple URI-reference as text. +/// +/// The origin-relative subset is exactly the one [`is_simple_uri_path`] +/// accepts. The absolute subset adds a `scheme "://" host [":" port]` prefix +/// whose host is spelled only with `unreserved` and `sub-delims` bytes, +/// followed by a path, query, and fragment obeying the same rules as the +/// relative shape. Userinfo, IP-literals, an empty host, and every percent +/// escape leave the subset. +/// +/// The subset holds nothing but ASCII, so recognizing it settles UTF-8 +/// validity too and the caller needs no second pass over the bytes. +/// +/// `None` means "not known to be valid", not "invalid". Callers must run a +/// full parse in that case. +/// +/// # Examples +/// +/// ``` +/// assert_eq!( +/// http_headers_simd::as_simple_uri_reference(b"https://example.com/docs"), +/// Some("https://example.com/docs"), +/// ); +/// assert_eq!( +/// http_headers_simd::as_simple_uri_reference(b"https://[::1]/"), +/// None +/// ); +/// ``` +#[must_use] +#[inline] +pub fn as_simple_uri_reference(bytes: &[u8]) -> Option<&str> { + dispatch::is_simple_uri_reference(bytes).then(|| { + debug_assert!(bytes.is_ascii()); + // SAFETY: every byte the subset admits is drawn from `unreserved`, + // `sub-delims`, and a handful of ASCII delimiters, so the slice is + // ASCII and therefore already valid UTF-8. + unsafe { str::from_utf8_unchecked(bytes) } + }) +} + +/// Scans one field line for a comma-separated list of bare HTTP tokens. +/// +/// Members are separated by commas and may be surrounded by optional +/// whitespace, and `empty` decides whether a zero-length member is ignored or +/// rejects the line. Anything else — a quoted string, a parameter, a control +/// byte, or two members separated by whitespace alone — leaves the subset and +/// reports [`TokenListScan::Rejected`], which means "not a simple token list" +/// rather than "malformed": callers that need a diagnostic re-scan the line +/// with their own parser. +/// +/// # Examples +/// +/// ``` +/// use http_headers_simd::{EmptyMembers, TokenListScan, scan_token_list}; +/// +/// assert_eq!( +/// scan_token_list(b"gzip, br", EmptyMembers::Skip), +/// TokenListScan::Members, +/// ); +/// assert_eq!( +/// scan_token_list(b"gzip,,br", EmptyMembers::Reject), +/// TokenListScan::Rejected, +/// ); +/// ``` +#[must_use] +#[inline] +pub fn scan_token_list(bytes: &[u8], empty: EmptyMembers) -> TokenListScan { + dispatch::scan_token_list(bytes, empty) +} + +/// Scans a `bytes=` field line for a byte-range-set that needs no parsing. +/// +/// `start` is the offset of the payload, that is, the byte after the `=`. The +/// answer is `true` only when the payload is a comma separated list of well +/// formed `first-last`, `first-`, and `-suffix` specifications whose bounds +/// are ordered and carry no leading zeros, so callers may accept the line +/// outright. Empty members are accepted because the grammar's readers skip +/// them. Everything else — a payload longer than one sixteen byte window, a +/// leading zero, whitespace anywhere but directly after a comma, or any other +/// byte — reports `false`, which means "not provably valid" rather than +/// "malformed": callers fall back to the parser that produces the diagnostic. +/// +/// # Examples +/// +/// ``` +/// assert!(http_headers_simd::scan_byte_range_set( +/// b"bytes=0-499,-200", +/// 6 +/// )); +/// assert!(!http_headers_simd::scan_byte_range_set(b"bytes=500-100", 6)); +/// ``` +#[must_use] +#[inline] +pub fn scan_byte_range_set(bytes: &[u8], start: usize) -> bool { + dispatch::scan_byte_range_set(bytes, start) +} + +/// Returns whether every byte is a standard base64 alphabet character. +/// +/// Padding is excluded, so callers must check `=` positionally themselves. +/// +/// # Examples +/// +/// ``` +/// assert!(http_headers_simd::all_base64_alphabet(b"AZaz09+/")); +/// assert!(!http_headers_simd::all_base64_alphabet(b"YWJj=")); +/// ``` +#[must_use] +#[inline] +pub fn all_base64_alphabet(bytes: &[u8]) -> bool { + dispatch::all_base64_alphabet(bytes) +} + +/// Finds the first comma, semicolon, quote, backslash, SP, or HTAB. +/// +/// # Examples +/// +/// ``` +/// assert_eq!(http_headers_simd::find_interesting(b"gzip, br"), Some(4)); +/// assert_eq!(http_headers_simd::find_interesting(b"gzip"), None); +/// ``` +#[must_use] +#[inline] +pub fn find_interesting(bytes: &[u8]) -> Option { + dispatch::find_interesting(bytes) +} + +/// Finds the first occurrence of either requested byte. +/// +/// # Examples +/// +/// ``` +/// assert_eq!( +/// http_headers_simd::find_either(b"gzip, br", b',', b'"'), +/// Some(4) +/// ); +/// assert_eq!(http_headers_simd::find_either(b"gzip", b',', b'"'), None); +/// ``` +#[must_use] +#[inline] +pub fn find_either(bytes: &[u8], first: u8, second: u8) -> Option { + dispatch::find_either(bytes, first, second) +} + +/// Reports the available backend tier at [`simd_threshold`]. +/// +/// Equivalent to [`backend_for`] at that length, not an operation-specific +/// kernel selection. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(any(feature = "benchmarking", feature = "test-util"))] +/// # { +/// use http_headers_simd::{Backend, backend}; +/// +/// assert!(matches!( +/// backend(), +/// Backend::Scalar | Backend::Sse2 | Backend::Sse42 | Backend::Neon +/// )); +/// # } +/// ``` +#[cfg(any(feature = "benchmarking", feature = "test-util"))] +#[must_use] +pub fn backend() -> Backend { + dispatch::backend() +} + +/// Reports an available backend tier using the shared byte-scanner cutoff. +/// +/// Returns [`Backend::Scalar`] below [`simd_threshold`]. At or above that +/// cutoff, the reported tier describes available instruction sets, not the +/// exact kernel every scanner selects: for example, x86 token validation uses +/// SSE2 even when this helper reports [`Backend::Sse42`]. +/// +/// Equality, URI-tail, base64, and token-list scanning have separate cutoffs +/// documented by [`simd_threshold`], so this helper does not predict their +/// scalar/SIMD choice. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(any(feature = "benchmarking", feature = "test-util"))] +/// # { +/// use http_headers_simd::{Backend, backend_for}; +/// +/// assert_eq!(backend_for(0), Backend::Scalar); +/// # } +/// ``` +#[cfg(any(feature = "benchmarking", feature = "test-util"))] +#[must_use] +pub fn backend_for(len: usize) -> Backend { + dispatch::backend_for(len) +} + +/// Returns the shared byte-scanner cutoff, not a crate-wide SIMD minimum. +/// +/// This cutoff applies to [`is_token`], [`is_token68`], [`is_field_value`], +/// [`find_interesting`], and [`find_either`]: 16 bytes on x86/x86-64 and 32 bytes +/// on other architectures, subject to instruction-set availability. +/// +/// [`eq_ignore_ascii_case`] uses a separate 32-byte cutoff. URI-tail scanning, +/// [`all_base64_alphabet`], and [`scan_token_list`] use 16-byte cutoffs. +/// For [`as_simple_uri_reference`], the URI cutoff applies to the tail after +/// any absolute authority prefix, not to the length of the whole reference. +/// Range-window classification in [`scan_byte_range_set`] and ASCII conversion +/// in [`ascii_str`] do not use this shared cutoff. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(any(feature = "benchmarking", feature = "test-util"))] +/// # { +/// use http_headers_simd::{Backend, backend_for, simd_threshold}; +/// +/// assert_eq!(backend_for(simd_threshold() - 1), Backend::Scalar); +/// # } +/// ``` +#[cfg(any(feature = "benchmarking", feature = "test-util"))] +#[must_use] +pub const fn simd_threshold() -> usize { + dispatch::SIMD_THRESHOLD +} + +/// A byte-scanning implementation. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(any(feature = "benchmarking", feature = "test-util"))] +/// # { +/// use http_headers_simd::Backend; +/// +/// assert_eq!(format!("{:?}", Backend::Scalar), "Scalar"); +/// # } +/// ``` +#[cfg(any(feature = "benchmarking", feature = "test-util"))] +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum Backend { + /// Generic safe Rust. + Scalar, + /// SSE2 on x86 or x86-64. + Sse2, + /// SSE4.2 on x86 or x86-64. + Sse42, + /// NEON on `AArch64`. + Neon, +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #[cfg(all( + any(feature = "benchmarking", feature = "test-util"), + any(target_arch = "x86", target_arch = "x86_64") + ))] + use std::arch; + #[cfg(not(feature = "std"))] + use std::format; + use std::time::Duration; + #[cfg(not(feature = "std"))] + use std::vec; + #[cfg(not(feature = "std"))] + use std::vec::Vec; + + use super::*; + use crate::{base64, list, range, scalar, uri}; + + fn lane_bytes(length: usize) -> impl Iterator { + // Miri keeps every byte in every lane across two blocks and a tail. + // Other lengths retain the boundaries of each byte class. + (u8::MIN..=u8::MAX).filter(move |byte| { + !cfg!(miri) + || length == 33 + || matches!( + *byte, + 0 | b'\t' | b'\n' | b'\r' | 0x1f..=b'0' | b'9'..=b'A' | b'Z'..=b'a' | b'z'..=0x80 | 0xff + ) + }) + } + + #[test] + fn token_byte_class_is_exhaustive() { + for byte in u8::MIN..=u8::MAX { + let expected = byte.is_ascii_alphanumeric() || b"!#$%&'*+-.^_`|~".contains(&byte); + assert_eq!(is_token(&[byte]), expected, "byte {byte:#04x}"); + } + assert!(!is_token(b"")); + } + + #[test] + fn token68_byte_class_is_exhaustive() { + for byte in u8::MIN..=u8::MAX { + let expected = byte.is_ascii_alphanumeric() || b"-._~+/".contains(&byte); + assert_eq!(is_token68(&[byte]), expected, "byte {byte:#04x}"); + } + assert!(!is_token68(b"")); + assert!(!is_token68(b"=")); + } + + #[test] + fn field_value_byte_class_is_exhaustive() { + for byte in u8::MIN..=u8::MAX { + let expected = byte == b'\t' || byte >= b' ' && byte != 0x7f; + assert_eq!(is_field_value(&[byte]), expected, "byte {byte:#04x}"); + } + assert!(is_field_value(b"")); + } + + #[test] + fn interesting_byte_class_is_exhaustive() { + for byte in u8::MIN..=u8::MAX { + let expected = b",;\"\\ \t".contains(&byte).then_some(0); + assert_eq!(find_interesting(&[byte]), expected, "byte {byte:#04x}"); + } + } + + #[test] + fn either_byte_search_matches_scalar_across_boundaries() { + for length in 0..=80 { + let mut bytes = vec![b'a'; length]; + assert_eq!(find_either(&bytes, b',', b'"'), None); + for position in 0..length { + bytes[position] = if position % 2 == 0 { b',' } else { b'"' }; + assert_eq!( + find_either(&bytes, b',', b'"'), + Some(position), + "length {length}, position {position}" + ); + bytes[position] = b'a'; + } + } + } + + #[test] + fn base64_byte_class_is_exhaustive() { + for byte in u8::MIN..=u8::MAX { + let expected = byte.is_ascii_alphanumeric() || byte == b'+' || byte == b'/'; + assert_eq!(base64::is_base64_alphabet_byte(byte), expected, "{byte:#04x}"); + let padded = [byte; 40]; + assert_eq!(all_base64_alphabet(&padded), expected, "{byte:#04x}"); + } + } + + #[test] + fn base64_boundary_lengths_match_scalar() { + for length in 0..=65 { + let mut bytes = vec![b'a'; length]; + if length != 0 { + bytes[length - 1] = b'='; + } + assert_eq!(all_base64_alphabet(&bytes), base64::all_base64_alphabet(&bytes), "length {length}"); + } + } + + #[test] + fn simple_uri_byte_class_is_exhaustive() { + for byte in u8::MIN..=u8::MAX { + let expected = byte.is_ascii_alphanumeric() || b"!#$&'()*+,-./:;=?@_~".contains(&byte); + let reference = [b'/', b'a', byte]; + assert_eq!(is_simple_uri_path(&reference), expected, "byte {byte:#04x}"); + } + assert!(is_simple_uri_path(b"/")); + assert!(!is_simple_uri_path(b"")); + assert!(!is_simple_uri_path(b"//host/path")); + assert!(!is_simple_uri_path(b"path")); + assert!(is_simple_uri_path(b"/docs/index.html?x=1#top")); + assert!(!is_simple_uri_path(b"/docs#one#two")); + assert!(!is_simple_uri_path(b"/caf%C3%A9")); + } + + #[test] + fn simple_uri_boundaries_and_lanes() { + for length in [16, 17, 31, 32, 33, 47, 48, 63, 64, 65] { + let mut bytes = vec![b'a'; length]; + bytes[0] = b'/'; + assert!(is_simple_uri_path(&bytes), "valid length {length}"); + + for lane in 1..length { + let mut invalid = bytes.clone(); + invalid[lane] = b'%'; + assert!(!is_simple_uri_path(&invalid), "escape lane {lane}, length {length}"); + + let mut fragment = bytes.clone(); + fragment[lane] = b'#'; + assert!(is_simple_uri_path(&fragment), "one hash at {lane}, length {length}"); + + for other in 1..length { + if other == lane { + continue; + } + let mut twice = fragment.clone(); + twice[other] = b'#'; + assert!(!is_simple_uri_path(&twice), "hashes at {lane} and {other}, length {length}"); + } + } + } + } + + #[test] + fn simd_lanes_match_scalar_for_every_byte() { + for length in [16, 17, 31, 32, 33, 48] { + for lane in 0..length { + for byte in lane_bytes(length) { + let mut bytes = vec![b'a'; length]; + bytes[lane] = byte; + let context = format!("byte {byte:#04x} at lane {lane}, length {length}"); + assert_eq!(is_token(&bytes), scalar::is_token(&bytes), "{context}"); + assert_eq!(all_base64_alphabet(&bytes), base64::all_base64_alphabet(&bytes), "{context}"); + assert_eq!(is_token68(&bytes), scalar::is_token68(&bytes), "{context}"); + assert_eq!(is_field_value(&bytes), scalar::is_field_value(&bytes), "{context}"); + assert_eq!(find_interesting(&bytes), scalar::find_interesting(&bytes), "{context}"); + } + } + } + } + + #[test] + fn simple_uri_lanes_match_scalar_for_every_byte() { + for length in [16, 17, 23, 31, 32, 33, 48, 49] { + for lane in 1..length { + for byte in lane_bytes(length) { + let mut bytes = vec![b'a'; length]; + bytes[0] = b'/'; + bytes[lane] = byte; + assert_eq!( + is_simple_uri_path(&bytes), + uri::is_simple_uri_path(&bytes), + "byte {byte:#04x} at lane {lane}, length {length}" + ); + } + } + } + } + + #[test] + fn ascii_text_conversion_matches_the_general_validator() { + for length in 0..=40 { + let bytes = vec![b'a'; length]; + assert_eq!(ascii_str(&bytes), Some(str::from_utf8(&bytes).expect("the probe is ASCII")),); + + for lane in 0..length { + for byte in [0x00_u8, 0x7f, 0x80, 0xc3, 0xff] { + let mut probe = bytes.clone(); + probe[lane] = byte; + assert_eq!( + ascii_str(&probe).is_some(), + byte.is_ascii(), + "byte {byte:#04x} at lane {lane}, length {length}" + ); + } + } + } + assert_eq!(ascii_str(b""), Some("")); + assert_eq!(ascii_str("caf\u{e9}".as_bytes()), None); + } + + #[test] + fn simple_uri_reference_accepts_the_absolute_shape() { + for reference in [ + "https://example.com", + "https://example.com/", + "https://example.com/a/b?c=d&e=f#g", + "http://example.com:8080/a", + "http://example.com:/a", + "ws+tls-1.0://h0st.example/a", + "https://example.com//double", + "/origin/relative?x=1#top", + ] { + assert!( + as_simple_uri_reference(reference.as_bytes()).is_some(), + "{reference} must be recognized" + ); + } + + for reference in [ + "", + "relative", + "//example.com/a", + "https://user@example.com/a", + "https://[::1]/a", + "https:///a", + "https://example.com:80x/a", + "https://exa%6dple.com/a", + "https://example.com/a%20b", + "https://example.com/a#b#c", + "1https://example.com/a", + "https:/example.com/a", + "https//example.com/a", + "mailto:user@example.com", + ] { + assert!( + as_simple_uri_reference(reference.as_bytes()).is_none(), + "{reference} must fall back" + ); + } + } + + #[test] + fn simple_uri_reference_lanes_match_scalar_for_every_byte() { + let prefix = b"https://example.com"; + for length in [24, 32, 33, 48, 49] { + for lane in 0..length { + for byte in lane_bytes(length) { + let mut bytes = vec![b'a'; length]; + bytes[..prefix.len()].copy_from_slice(prefix); + bytes[prefix.len()] = b'/'; + bytes[lane] = byte; + assert_eq!( + as_simple_uri_reference(&bytes).is_some(), + uri::is_simple_uri_reference(&bytes), + "byte {byte:#04x} at lane {lane}, length {length}" + ); + } + } + } + } + + #[test] + fn ascii_case_equality_is_exhaustive() { + for left in u8::MIN..=u8::MAX { + for right in u8::MIN..=u8::MAX { + assert_eq!( + eq_ignore_ascii_case(&[left], &[right]), + left.eq_ignore_ascii_case(&right), + "{left:#04x} versus {right:#04x}" + ); + } + } + assert!(!eq_ignore_ascii_case(b"a", b"aa")); + } + + #[test] + fn boundary_lengths_match_scalar() { + for length in 0..=65 { + let mut bytes = [b'a'; 65]; + if length != 0 { + bytes[length - 1] = b';'; + } + let bytes = &bytes[..length]; + assert_eq!(is_token(bytes), scalar::is_token(bytes)); + assert_eq!(is_field_value(bytes), scalar::is_field_value(bytes)); + assert_eq!(eq_ignore_ascii_case(bytes, bytes), scalar::eq_ignore_ascii_case(bytes, bytes)); + assert_eq!(find_interesting(bytes), scalar::find_interesting(bytes)); + assert_eq!(is_simple_uri_path(bytes), uri::is_simple_uri_path(bytes)); + } + } + + #[test] + fn valid_tokens_at_vector_boundaries() { + for length in [1, 15, 16, 17, 31, 32, 33, 47, 48, 63, 64, 65, 80] { + let bytes = vec![b'a'; length]; + assert!(is_token(&bytes), "valid token length {length}"); + } + } + + #[test] + fn every_lane_is_checked_at_vector_boundaries() { + for length in [16, 17, 31, 32, 33, 47, 48, 63, 64, 65] { + for lane in 0..length { + let mut bytes = vec![b'a'; length]; + bytes[lane] = b'('; + assert!(!is_token(&bytes), "invalid lane {lane}, length {length}"); + + bytes[lane] = b';'; + assert_eq!(find_interesting(&bytes), Some(lane), "interesting lane {lane}, length {length}"); + } + } + } + + #[test] + fn token68_boundaries_and_lanes() { + for length in [1, 15, 16, 17, 31, 32, 33, 47, 48, 63, 64, 65, 127, 128] { + let valid = vec![b'a'; length]; + assert!(is_token68(&valid), "all-data length {length}"); + + for lane in 0..length { + let mut invalid = valid.clone(); + invalid[lane] = b':'; + assert!(!is_token68(&invalid), "invalid lane {lane}, length {length}"); + + let mut padded = valid.clone(); + padded[lane..].fill(b'='); + assert_eq!(is_token68(&padded), lane != 0, "padding lane {lane}, length {length}"); + + if lane + 1 < length { + let next = lane + 1; + padded[next] = b'a'; + assert!(!is_token68(&padded), "data after padding at lane {next}, length {length}"); + } + } + } + } + + #[test] + fn equal_length_ascii_case_folding_crosses_vector_boundaries() { + let lower = b"abcdefghijklmnopqrstuvwxyz012345abcdefghijklmnopqrstuvwxyz012345"; + let upper = b"ABCDEFGHIJKLMNOPQRSTUVWXYZ012345ABCDEFGHIJKLMNOPQRSTUVWXYZ012345"; + assert_eq!(lower.len(), 64); + assert!(eq_ignore_ascii_case(lower, upper)); + + for lane in 0..lower.len() { + let mut different = *upper; + different[lane] = if lower[lane].is_ascii_alphabetic() { b'0' } else { b'X' }; + assert!(!eq_ignore_ascii_case(lower, &different), "different lane {lane}"); + } + + let mut non_ascii_left = [b'a'; 32]; + let mut non_ascii_right = [b'A'; 32]; + non_ascii_left[16] = 0x80; + non_ascii_right[16] = 0x80; + assert!(eq_ignore_ascii_case(&non_ascii_left, &non_ascii_right)); + non_ascii_right[16] = 0x81; + assert!(!eq_ignore_ascii_case(&non_ascii_left, &non_ascii_right)); + } + + #[test] + #[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] + fn differential_properties() { + bolero::check!() + .with_iterations(4_096) + .with_test_time(Duration::from_millis(400)) + .with_type::<(Vec, Vec)>() + .for_each(|(left, right)| { + assert_eq!(is_token(left), scalar::is_token(left)); + assert_eq!(is_token68(left), scalar::is_token68(left)); + assert_eq!(is_field_value(left), scalar::is_field_value(left)); + assert_eq!(eq_ignore_ascii_case(left, right), scalar::eq_ignore_ascii_case(left, right)); + assert_eq!(find_interesting(left), scalar::find_interesting(left)); + assert_eq!(is_simple_uri_path(left), uri::is_simple_uri_path(left)); + for empty in [EmptyMembers::Skip, EmptyMembers::Reject] { + let expected = list::oracle(left, empty); + assert_eq!(scan_token_list(left, empty), expected); + assert_eq!(list::scan_token_list(left, empty), expected); + } + for start in 0..=left.len() { + let scanned = scan_byte_range_set(left, start); + assert_eq!(scanned, range::scan_byte_range_set(left, start)); + assert!(!scanned || range::oracle(&left[start..])); + } + }); + } + + #[cfg(any(feature = "benchmarking", feature = "test-util"))] + #[test] + fn backend_helpers_report_threshold_dispatch() { + let expected_threshold = if cfg!(any(target_arch = "x86", target_arch = "x86_64")) { + 16 + } else { + 32 + }; + assert_eq!(simd_threshold(), expected_threshold); + assert_eq!(backend_for(simd_threshold() - 1), Backend::Scalar); + assert_eq!(backend(), backend_for(simd_threshold())); + + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + { + let selected = backend(); + assert!([Backend::Sse2, Backend::Sse42].contains(&selected)); + assert_eq!(selected == Backend::Sse42, arch::is_x86_feature_detected!("sse4.2")); + } + } + + /// Checks the accelerated range scanner against the scalar one. + /// + /// Every payload the kernels classify is compared against both the scalar + /// classification and the parser-shaped oracle, at every payload offset a + /// field line can hand over, and with the payload placed both inside a + /// borrowed window and inside a padded one. + #[test] + fn range_scan_matches_scalar_for_every_short_payload() { + let alphabet = b"0-, 1\t9"; + let mut line = Vec::new(); + for prefix in [&b"bytes="[..], b"b=", b"aaaaaaaaaaaaaaaabytes="] { + for length in 0..=4_usize { + for index in 0..alphabet.len().pow(u32::try_from(length).unwrap_or(0)) { + line.clear(); + line.extend_from_slice(prefix); + let mut rest = index; + for _step in 0..length { + line.push(alphabet[rest % alphabet.len()]); + rest /= alphabet.len(); + } + let start = prefix.len(); + let scanned = scan_byte_range_set(&line, start); + assert_eq!(scanned, range::scan_byte_range_set(&line, start), "{line:?} at {start}"); + if scanned { + assert!(range::oracle(&line[start..]), "{line:?} at {start}"); + } + } + } + } + } + + /// Slides every short byte pattern across a block boundary. + /// + /// The scanners carry token/whitespace state and whether token or empty + /// members have occurred across blocks. Patterns straddling a block + /// exercise those transitions at every nearby offset. + #[test] + fn token_list_scan_matches_scalar_across_block_boundaries() { + let alphabet = b"a,\t %"; + let mut pattern = Vec::new(); + slide_patterns(alphabet, 3, &mut pattern); + } + + /// Puts every byte in every position of every length that leaves a tail. + /// + /// A line whose length is not a multiple of the block width is finished by + /// reloading the last whole block and dropping the lanes already folded, so + /// an off-by-one in that shift shows up as one byte in one position of one + /// length disagreeing with the scalar state machine. + #[test] + fn token_list_tails_match_scalar_for_every_length_and_byte() { + for length in 16..=33_usize { + for lane in 0..length { + for byte in lane_bytes(length) { + let mut input = vec![b'a'; length]; + input[lane] = byte; + for empty in [EmptyMembers::Skip, EmptyMembers::Reject] { + let expected = list::oracle(&input, empty); + assert_eq!( + scan_token_list(&input, empty), + expected, + "length {length} lane {lane} byte {byte:#04x}" + ); + } + } + } + } + } + + /// Walks every tail a well-formed list can end with, across two blocks. + /// + /// The tail carries the state the earlier blocks left, so the interesting + /// cases are the ones where a member, a comma, or a whitespace run spans + /// the boundary between the last whole block and the reloaded one. + #[test] + fn token_list_tails_match_scalar_for_every_short_ending() { + let alphabet = b"a,\t x"; + for prefix in [&b"abc, def, ghij"[..], &b"a,b,c,d,e,f,g,h,i,j,k,l"[..], &b" a , b , c "[..]] { + let mut ending = Vec::new(); + walk_endings(alphabet, 4, prefix, &mut ending); + } + } + + fn walk_endings(alphabet: &[u8], length: usize, prefix: &[u8], ending: &mut Vec) { + if length == 0 { + let mut input = prefix.to_vec(); + input.extend_from_slice(ending); + for empty in [EmptyMembers::Skip, EmptyMembers::Reject] { + let expected = list::oracle(&input, empty); + assert_eq!(scan_token_list(&input, empty), expected, "{input:?}"); + } + return; + } + for &byte in alphabet { + ending.push(byte); + walk_endings(alphabet, length - 1, prefix, ending); + ending.pop(); + } + walk_endings(alphabet, 0, prefix, ending); + } + + fn slide_patterns(alphabet: &[u8], length: usize, pattern: &mut Vec) { + if length == 0 { + for &fill in b"a, \t" { + for lead in 0..=17 { + let mut input = vec![fill; lead]; + input.extend_from_slice(pattern); + input.resize(input.len() + 19, fill); + for empty in [EmptyMembers::Skip, EmptyMembers::Reject] { + let expected = list::oracle(&input, empty); + assert_eq!(scan_token_list(&input, empty), expected, "{input:?}"); + assert_eq!(list::scan_token_list(&input, empty), expected, "{input:?}"); + } + } + } + return; + } + for &byte in alphabet { + pattern.push(byte); + slide_patterns(alphabet, length - 1, pattern); + pattern.pop(); + } + slide_patterns(alphabet, 0, pattern); + } +} diff --git a/crates/http_headers_simd/src/arm.rs b/crates/http_headers_simd/src/arm.rs new file mode 100644 index 000000000..d4fcc47cc --- /dev/null +++ b/crates/http_headers_simd/src/arm.rs @@ -0,0 +1,541 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! `AArch64` NEON implementations of shared byte scanners. + +use core::arch::aarch64::*; + +use crate::list; +use crate::list::{EmptyMembers, ListScan, TokenListScan}; +use crate::range::{RangeMasks, WINDOW}; + +/// NEON byte vectors contain sixteen lanes. +const WIDTH: usize = 16; + +/// The same width counted in mask lanes. +const WIDTH_LANES: u32 = 16; + +/// # Safety +/// +/// The processor must support NEON. +#[target_feature(enable = "neon")] +pub(super) unsafe fn is_token(bytes: &[u8]) -> bool { + if bytes.is_empty() { + return false; + } + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { vld1q_u8(pointer) }; + if !all_set(token_mask(value)) { + return false; + } + offset += WIDTH; + } + crate::scalar::all_token_bytes(&bytes[offset..]) +} + +/// # Safety +/// +/// The processor must support NEON. +#[target_feature(enable = "neon")] +pub(super) unsafe fn is_token68(bytes: &[u8]) -> bool { + let mut offset = 0; + while offset + 2 * WIDTH <= bytes.len() { + let first_pointer = bytes.as_ptr().wrapping_add(offset); + // SAFETY: `offset + 2 * WIDTH <= bytes.len()` bounds this unaligned load. + let first = unsafe { vld1q_u8(first_pointer) }; + let second_pointer = bytes.as_ptr().wrapping_add(offset + WIDTH); + // SAFETY: `offset + 2 * WIDTH <= bytes.len()` bounds this unaligned load. + let second = unsafe { vld1q_u8(second_pointer) }; + if !all_set(token68_data_mask(first)) || !all_set(token68_data_mask(second)) { + return crate::scalar::token68_tail(&bytes[offset..], offset != 0); + } + offset += 2 * WIDTH; + } + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { vld1q_u8(pointer) }; + if !all_set(token68_data_mask(value)) { + return crate::scalar::token68_tail(&bytes[offset..], offset != 0); + } + offset += WIDTH; + } + crate::scalar::token68_tail(&bytes[offset..], offset != 0) +} + +/// # Safety +/// +/// The processor must support NEON. +#[target_feature(enable = "neon")] +pub(super) unsafe fn is_field_value(bytes: &[u8]) -> bool { + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { vld1q_u8(pointer) }; + let control = vandq_u8(vcltq_u8(value, vdupq_n_u8(0x20)), vmvnq_u8(vceqq_u8(value, vdupq_n_u8(9)))); + let invalid = vorrq_u8(control, vceqq_u8(value, vdupq_n_u8(0x7f))); + if any_set(invalid) { + return false; + } + offset += WIDTH; + } + crate::scalar::is_field_value(&bytes[offset..]) +} + +/// # Safety +/// +/// The processor must support NEON, and `right.len()` must equal `left.len()` +/// because vector loads from both slices are bounded using `left.len()`. +#[target_feature(enable = "neon")] +pub(super) unsafe fn eq_ignore_ascii_case(left: &[u8], right: &[u8]) -> bool { + let mut offset = 0; + while offset + WIDTH <= left.len() { + let left_pointer = left.as_ptr().wrapping_add(offset); + // SAFETY: the dispatcher supplies equal-length slices and the loop bounds this load. + let left_value = unsafe { vld1q_u8(left_pointer) }; + let right_pointer = right.as_ptr().wrapping_add(offset); + // SAFETY: the dispatcher supplies equal-length slices and the loop bounds this load. + let right_value = unsafe { vld1q_u8(right_pointer) }; + if !all_set(vceqq_u8(lower(left_value), lower(right_value))) { + return false; + } + offset += WIDTH; + } + crate::scalar::eq_ignore_ascii_case(&left[offset..], &right[offset..]) +} + +/// Checks that every byte is a base64 alphabet character. +/// +/// # Safety +/// +/// The processor must support NEON. +#[target_feature(enable = "neon")] +pub(super) unsafe fn all_base64_alphabet(bytes: &[u8]) -> bool { + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { vld1q_u8(pointer) }; + if !all_set(base64_accept_mask(value)) { + return false; + } + offset += WIDTH; + } + crate::base64::all_base64_alphabet(&bytes[offset..]) +} + +/// Scans a reference whose origin-relative prefix the dispatcher already checked. +/// +/// # Safety +/// +/// The processor must support NEON. +#[target_feature(enable = "neon")] +pub(super) unsafe fn is_simple_uri_tail(bytes: &[u8]) -> bool { + let mut offset = 0; + let mut hashes = 0_u32; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { vld1q_u8(pointer) }; + let accepted = simple_uri_accept_mask(value); + if !all_set(accepted) { + let fragments = vceqq_u8(value, vdupq_n_u8(b'#')); + if !all_set(vorrq_u8(accepted, fragments)) { + return false; + } + hashes += u32::from(hash_count(value)); + if hashes > 1 { + return false; + } + } + offset += WIDTH; + } + crate::uri::simple_uri_tail(&bytes[offset..], hashes) +} + +/// Scans a comma-separated token list one block at a time. +/// +/// # Safety +/// +/// The processor must support NEON, and `bytes.len()` must be at least `WIDTH` +/// because the trailing run is folded from a full vector of input. +#[target_feature(enable = "neon")] +pub(super) unsafe fn scan_token_list(bytes: &[u8], empty: EmptyMembers) -> TokenListScan { + let mut state = ListScan::new(); + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { vld1q_u8(pointer) }; + let (tokens, ows, commas) = list_masks(value); + if !state.push_block(tokens, ows, commas) { + return TokenListScan::Rejected; + } + offset += WIDTH; + } + if offset == bytes.len() { + return state.finish(&[], empty); + } + let final_offset = bytes.len() - WIDTH; + let pointer = bytes.as_ptr().wrapping_add(final_offset); + // SAFETY: the caller guarantees `bytes.len() >= WIDTH`, so this load ends at the slice end. + let value = unsafe { vld1q_u8(pointer) }; + let (tokens, ows, commas) = list_masks(value); + let shift = u32::try_from(offset - final_offset).unwrap_or(WIDTH_LANES); + if !state.push_partial_block(tokens >> shift, ows >> shift, commas >> shift, WIDTH_LANES - shift) { + return TokenListScan::Rejected; + } + state.finish(&[], empty) +} + +/// Scans a token list of one to two vectors as a single fold. +/// +/// # Safety +/// +/// The processor must support NEON, `bytes.len()` must be at least `WIDTH` and +/// at most twice it, and `TWO` must say whether the line runs past the first +/// vector, because both vectors are loaded whole. +#[target_feature(enable = "neon")] +pub(super) unsafe fn scan_short_token_list(bytes: &[u8], empty: EmptyMembers) -> TokenListScan { + let last = bytes.len() - WIDTH; + let head = bytes.as_ptr(); + let tail = bytes.as_ptr().wrapping_add(last); + // SAFETY: the caller guarantees a whole vector at offset zero. + let first = unsafe { vld1q_u8(head) }; + // SAFETY: the caller guarantees a whole vector at `last`. + let final_value = unsafe { vld1q_u8(tail) }; + let (first_tokens, first_ows, first_commas) = list_masks(first); + let mut state = ListScan::new(); + let folded = if TWO { + let (last_tokens, last_ows, last_commas) = list_masks(final_value); + let lanes = u32::try_from(last).unwrap_or(WIDTH_LANES); + let shift = WIDTH_LANES - lanes; + state.push_line( + list::join_lanes(first_tokens, last_tokens, shift), + list::join_lanes(first_ows, last_ows, shift), + list::join_lanes(first_commas, last_commas, shift), + WIDTH_LANES + lanes, + ) + } else { + state.push_block(first_tokens, first_ows, first_commas) + }; + if !folded { + return TokenListScan::Rejected; + } + state.finish(&[], empty) +} + +/// Splits one block into the token, whitespace, and comma lanes a list needs. +/// +/// # Safety +/// +/// The processor must support NEON. +#[target_feature(enable = "neon")] +fn list_masks(value: uint8x16_t) -> (u32, u32, u32) { + let tokens = movemask(token_mask(value)); + let ows = movemask(vorrq_u8(vceqq_u8(value, vdupq_n_u8(b' ')), vceqq_u8(value, vdupq_n_u8(b'\t')))); + let commas = movemask(vceqq_u8(value, vdupq_n_u8(b','))); + (tokens, ows, commas) +} + +/// # Safety +/// +/// The processor must support NEON. +#[target_feature(enable = "neon")] +pub(super) unsafe fn find_interesting(bytes: &[u8]) -> Option { + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { vld1q_u8(pointer) }; + let matches = b",;\"\\ \t" + .iter() + .fold(vdupq_n_u8(0), |mask, byte| vorrq_u8(mask, vceqq_u8(value, vdupq_n_u8(*byte)))); + if any_set(matches) { + let mut lanes = [0_u8; WIDTH]; + // SAFETY: `lanes` has exactly 16 writable bytes. + unsafe { vst1q_u8(lanes.as_mut_ptr(), matches) }; + return lanes.iter().position(|byte| *byte != 0).map(|index| offset + index); + } + offset += WIDTH; + } + crate::scalar::find_interesting(&bytes[offset..]).map(|index| offset + index) +} + +/// # Safety +/// +/// The processor must support NEON. +#[target_feature(enable = "neon")] +pub(super) unsafe fn find_either(bytes: &[u8], first: u8, second: u8) -> Option { + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { vld1q_u8(pointer) }; + let matches = vorrq_u8(vceqq_u8(value, vdupq_n_u8(first)), vceqq_u8(value, vdupq_n_u8(second))); + if any_set(matches) { + let mut lanes = [0_u8; WIDTH]; + // SAFETY: `lanes` has exactly 16 writable bytes. + unsafe { vst1q_u8(lanes.as_mut_ptr(), matches) }; + return lanes.iter().position(|byte| *byte != 0).map(|index| offset + index); + } + offset += WIDTH; + } + crate::scalar::find_either(&bytes[offset..], first, second).map(|index| offset + index) +} + +#[target_feature(enable = "neon")] +fn token_mask(value: uint8x16_t) -> uint8x16_t { + let mut mask = vorrq_u8(in_range(value, b'0', b'9'), in_range(value, b'A', b'Z')); + mask = vorrq_u8(mask, in_range(value, b'a', b'z')); + b"!#$%&'*+-.^_`|~" + .iter() + .fold(mask, |mask, byte| vorrq_u8(mask, vceqq_u8(value, vdupq_n_u8(*byte)))) +} + +#[target_feature(enable = "neon")] +fn token68_data_mask(value: uint8x16_t) -> uint8x16_t { + let digit = in_range(value, b'0', b'9'); + let folded = vorrq_u8(value, vdupq_n_u8(0x20)); + let alpha = in_range(folded, b'a', b'z'); + let plus_to_slash = in_range(value, b'+', b'/'); + let symbols = vbicq_u8(plus_to_slash, vceqq_u8(value, vdupq_n_u8(b','))); + let underscore = vceqq_u8(value, vdupq_n_u8(b'_')); + let tilde = vceqq_u8(value, vdupq_n_u8(b'~')); + vorrq_u8(vorrq_u8(digit, alpha), vorrq_u8(symbols, vorrq_u8(underscore, tilde))) +} + +/// Marks the lanes inside the base64 alphabet. +/// +/// `/` and the digits are contiguous, so the alphabet needs only three ranges and one compare. +#[target_feature(enable = "neon")] +fn base64_accept_mask(value: uint8x16_t) -> uint8x16_t { + let digit_or_slash = in_range(value, b'/', b'9'); + let upper = in_range(value, b'A', b'Z'); + let lower_case = in_range(value, b'a', b'z'); + let plus = vceqq_u8(value, vdupq_n_u8(b'+')); + vorrq_u8(vorrq_u8(digit_or_slash, upper), vorrq_u8(lower_case, plus)) +} + +/// Marks the lanes inside the separator-free URI subset. +/// +/// The `&`..`?` range covers sub-delims, digits, `:`, and `;` after excluding +/// `<` and `>`. Case folding merges the letter ranges; isolated bytes use +/// equality comparisons. `#` is excluded for separate fragment counting. +#[target_feature(enable = "neon")] +fn simple_uri_accept_mask(value: uint8x16_t) -> uint8x16_t { + let raised = vorrq_u8(value, vdupq_n_u8(2)); + let punctuation = vbicq_u8(in_range(value, b'&', b'?'), vceqq_u8(raised, vdupq_n_u8(0x3e))); + let bang = vceqq_u8(value, vdupq_n_u8(b'!')); + let dollar = vceqq_u8(value, vdupq_n_u8(b'$')); + let letter = in_range(vorrq_u8(value, vdupq_n_u8(0x20)), b'a', b'z'); + let at = vceqq_u8(value, vdupq_n_u8(b'@')); + let underscore = vceqq_u8(value, vdupq_n_u8(b'_')); + let tilde = vceqq_u8(value, vdupq_n_u8(b'~')); + vorrq_u8( + vorrq_u8(punctuation, vorrq_u8(bang, dollar)), + vorrq_u8(letter, vorrq_u8(at, vorrq_u8(underscore, tilde))), + ) +} + +/// Packs one bit per lane, lane zero in bit zero, from an all-ones lane mask. +/// +/// NEON has no move-mask instruction, so each lane keeps a single distinct bit +/// and the halves are summed across, which costs two horizontal adds. +#[target_feature(enable = "neon")] +fn movemask(mask: uint8x16_t) -> u32 { + const BITS: [u8; WIDTH] = [1, 2, 4, 8, 16, 32, 64, 128, 1, 2, 4, 8, 16, 32, 64, 128]; + + // SAFETY: `BITS` is exactly 16 bytes, so this unaligned load stays in bounds. + let bits = unsafe { vld1q_u8(BITS.as_ptr()) }; + let selected = vandq_u8(mask, bits); + let low = u32::from(vaddv_u8(vget_low_u8(selected))); + let high = u32::from(vaddv_u8(vget_high_u8(selected))); + low | (high << 8) +} + +/// Classifies one window for the byte-range grammar. +/// +/// Digits fall out of one biased unsigned comparison, and the three delimiters +/// and the leading-zero marker are plain equality comparisons. +/// +/// # Safety +/// +/// The processor must support NEON. +#[inline] +#[target_feature(enable = "neon")] +pub(super) unsafe fn range_masks(window: &[u8; WINDOW]) -> RangeMasks { + // SAFETY: the window is exactly 16 bytes, so this unaligned load stays in bounds. + let value = unsafe { vld1q_u8(window.as_ptr()) }; + let biased = vsubq_u8(value, vdupq_n_u8(b'0')); + let digits = vcleq_u8(biased, vdupq_n_u8(9)); + let zeros = vceqq_u8(value, vdupq_n_u8(b'0')); + let dashes = vceqq_u8(value, vdupq_n_u8(b'-')); + let commas = vceqq_u8(value, vdupq_n_u8(b',')); + let spaces = vorrq_u8(vceqq_u8(value, vdupq_n_u8(b' ')), vceqq_u8(value, vdupq_n_u8(b'\t'))); + RangeMasks { + digits: movemask(digits), + dashes: movemask(dashes), + commas: movemask(commas), + spaces: movemask(spaces), + zeros: movemask(zeros), + } +} + +#[target_feature(enable = "neon")] +fn hash_count(value: uint8x16_t) -> u8 { + vaddvq_u8(vshrq_n_u8::<7>(vceqq_u8(value, vdupq_n_u8(b'#')))) +} + +#[target_feature(enable = "neon")] +fn lower(value: uint8x16_t) -> uint8x16_t { + vorrq_u8(value, vandq_u8(in_range(value, b'A', b'Z'), vdupq_n_u8(0x20))) +} + +#[target_feature(enable = "neon")] +fn in_range(value: uint8x16_t, start: u8, end: u8) -> uint8x16_t { + vandq_u8(vcgeq_u8(value, vdupq_n_u8(start)), vcleq_u8(value, vdupq_n_u8(end))) +} + +#[target_feature(enable = "neon")] +fn any_set(value: uint8x16_t) -> bool { + vmaxvq_u8(value) != 0 +} + +#[target_feature(enable = "neon")] +fn all_set(value: uint8x16_t) -> bool { + vminvq_u8(value) == u8::MAX +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::arch; + use std::time::Duration; + #[cfg(not(feature = "std"))] + use std::vec::Vec; + + use super::*; + + /// Puts every byte value in every lane of a token-list block. + #[test] + fn token_list_lanes_match_scalar_for_every_byte() { + if !arch::is_aarch64_feature_detected!("neon") { + return; + } + for lane in 0..2 * WIDTH { + for byte in u8::MIN..=u8::MAX { + let mut block = [b'a'; 2 * WIDTH]; + block[lane] = byte; + for empty in [EmptyMembers::Skip, EmptyMembers::Reject] { + // SAFETY: runtime detection above establishes NEON support. + let list = unsafe { scan_token_list(&block, empty) }; + assert_eq!(list, crate::list::scan_token_list(&block, empty), "byte {byte:#04x} lane {lane}"); + } + } + } + } + + #[test] + #[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] + fn neon_matches_scalar() { + let available = arch::is_aarch64_feature_detected!("neon"); + if !available { + return; + } + bolero::check!() + .with_iterations(4_096) + .with_test_time(Duration::from_millis(400)) + .with_type::<(Vec, Vec)>() + .for_each(|(left, right)| { + // SAFETY: runtime detection above establishes NEON support. + let token = unsafe { is_token(left) }; + assert_eq!(token, crate::scalar::is_token(left)); + // SAFETY: runtime detection above establishes NEON support. + let token68 = unsafe { is_token68(left) }; + assert_eq!(token68, crate::scalar::is_token68(left)); + // SAFETY: runtime detection above establishes NEON support. + let field_value = unsafe { is_field_value(left) }; + assert_eq!(field_value, crate::scalar::is_field_value(left)); + if left.len() == right.len() { + // SAFETY: runtime detection establishes NEON support; lengths are equal. + let equal = unsafe { eq_ignore_ascii_case(left, right) }; + assert_eq!(equal, crate::scalar::eq_ignore_ascii_case(left, right)); + } + // SAFETY: runtime detection above establishes NEON support. + let interesting = unsafe { find_interesting(left) }; + assert_eq!(interesting, crate::scalar::find_interesting(left)); + // SAFETY: runtime detection above establishes NEON support. + let base64 = unsafe { all_base64_alphabet(left) }; + assert_eq!(base64, crate::base64::all_base64_alphabet(left)); + // SAFETY: runtime detection above establishes NEON support. + let simple = unsafe { is_simple_uri_tail(left) }; + assert_eq!(simple, crate::uri::simple_uri_tail(left, 0)); + if left.len() >= WIDTH { + for empty in [EmptyMembers::Skip, EmptyMembers::Reject] { + let expected = crate::list::scan_token_list(left, empty); + // SAFETY: runtime detection establishes NEON support; the length is checked. + let list = unsafe { scan_token_list(left, empty) }; + assert_eq!(list, expected); + if left.len() == WIDTH { + // SAFETY: detection above, and the line is exactly one vector. + let short = unsafe { scan_short_token_list::(left, empty) }; + assert_eq!(short, expected); + } else if left.len() <= 2 * WIDTH { + // SAFETY: detection above, and the line spans two vectors. + let short = unsafe { scan_short_token_list::(left, empty) }; + assert_eq!(short, expected); + } + } + } + }); + } + + #[test] + fn range_lanes_match_scalar_for_every_byte() { + if !arch::is_aarch64_feature_detected!("neon") { + return; + } + for lane in 0..WINDOW { + for byte in u8::MIN..=u8::MAX { + let mut window = [b'0'; WINDOW]; + window[lane] = byte; + // SAFETY: runtime detection above establishes NEON support. + let masks = unsafe { range_masks(&window) }; + assert_eq!(masks, RangeMasks::scalar(&window), "byte {byte:#04x} lane {lane}"); + } + } + } + + #[test] + fn base64_and_uri_lanes_match_neon() { + if !arch::is_aarch64_feature_detected!("neon") { + return; + } + for lane in 0..2 * WIDTH { + for byte in u8::MIN..=u8::MAX { + let mut base64_block = [b'A'; 2 * WIDTH]; + base64_block[lane] = byte; + assert_eq!( + // SAFETY: runtime detection above establishes NEON support. + unsafe { all_base64_alphabet(&base64_block) }, + crate::base64::all_base64_alphabet(&base64_block), + "base64 byte {byte:#04x} lane {lane}" + ); + + let mut uri_block = [b'a'; 2 * WIDTH]; + uri_block[lane] = byte; + assert_eq!( + // SAFETY: runtime detection above establishes NEON support. + unsafe { is_simple_uri_tail(&uri_block) }, + crate::uri::simple_uri_tail(&uri_block, 0), + "URI byte {byte:#04x} lane {lane}" + ); + } + } + } +} diff --git a/crates/http_headers_simd/src/base64.rs b/crates/http_headers_simd/src/base64.rs new file mode 100644 index 000000000..01d7763ea --- /dev/null +++ b/crates/http_headers_simd/src/base64.rs @@ -0,0 +1,35 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Scalar base64 alphabet classification used as the SIMD oracle. + +/// Returns whether one byte is a standard base64 alphabet character. +/// +/// Padding is deliberately excluded: callers know where `=` may appear and +/// check it positionally. +pub(super) fn is_base64_alphabet_byte(byte: u8) -> bool { + byte.is_ascii_alphanumeric() || byte == b'+' || byte == b'/' +} + +/// Runs the whole check without any architecture-specific acceleration. +pub(super) fn all_base64_alphabet(bytes: &[u8]) -> bool { + bytes.iter().copied().all(is_base64_alphabet_byte) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::*; + + #[test] + fn alphabet_accepts_data_but_not_padding_or_url_safe_symbols() { + assert!(all_base64_alphabet( + b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/" + )); + assert!(all_base64_alphabet(b"")); + for byte in [b'=', b'-', b'_', b' ', b'\n', 0x80] { + assert!(!is_base64_alphabet_byte(byte), "{byte:#04x}"); + assert!(!all_base64_alphabet(&[b'A', byte]), "{byte:#04x}"); + } + } +} diff --git a/crates/http_headers_simd/src/benchmarking.rs b/crates/http_headers_simd/src/benchmarking.rs new file mode 100644 index 000000000..789923f08 --- /dev/null +++ b/crates/http_headers_simd/src/benchmarking.rs @@ -0,0 +1,322 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Direct reference-backend access for differential benchmarks. +//! +//! Scalar functions provide benchmark references. Architecture-specific URI +//! functions return `None` when unavailable and bypass application dispatch. +//! +//! # Examples +//! +//! ``` +//! # #[cfg(feature = "benchmarking")] +//! # { +//! use http_headers_simd::benchmarking; +//! +//! assert_eq!( +//! benchmarking::is_token_scalar(b"content-type"), +//! http_headers_simd::is_token(b"content-type"), +//! ); +//! # } +//! ``` + +#[cfg(all(feature = "std", any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] +use std::arch; + +use crate::{EmptyMembers, TokenListScan, base64, list, range, scalar, uri}; + +/// Runs the production scalar token validator. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// assert!(http_headers_simd::benchmarking::is_token_scalar( +/// b"content-type" +/// )); +/// ``` +#[must_use] +pub fn is_token_scalar(bytes: &[u8]) -> bool { + scalar::is_token(bytes) +} + +/// Runs the production scalar `token68` validator. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// assert!(http_headers_simd::benchmarking::is_token68_scalar( +/// b"YWJjZA==" +/// )); +/// ``` +#[must_use] +pub fn is_token68_scalar(bytes: &[u8]) -> bool { + scalar::is_token68(bytes) +} + +/// Runs the production scalar field-value validator. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// assert!(http_headers_simd::benchmarking::is_field_value_scalar( +/// b"text/plain" +/// )); +/// ``` +#[must_use] +pub fn is_field_value_scalar(bytes: &[u8]) -> bool { + scalar::is_field_value(bytes) +} + +/// Runs the production scalar interesting-byte scanner. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// assert_eq!( +/// http_headers_simd::benchmarking::find_interesting_scalar(b"gzip, br"), +/// Some(4), +/// ); +/// ``` +#[must_use] +pub fn find_interesting_scalar(bytes: &[u8]) -> Option { + scalar::find_interesting(bytes) +} + +/// Runs the production scalar base64-alphabet scanner. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// assert!(http_headers_simd::benchmarking::all_base64_alphabet_scalar( +/// b"AZaz09+/" +/// )); +/// ``` +#[must_use] +pub fn all_base64_alphabet_scalar(bytes: &[u8]) -> bool { + base64::all_base64_alphabet(bytes) +} + +/// Runs the unaccelerated origin-relative reference check. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// assert!(http_headers_simd::benchmarking::is_simple_uri_path_scalar( +/// b"/docs?page=2", +/// )); +/// ``` +#[must_use] +pub fn is_simple_uri_path_scalar(bytes: &[u8]) -> bool { + uri::is_simple_uri_path(bytes) +} + +/// Runs the unaccelerated byte-range-set scan. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// assert!(http_headers_simd::benchmarking::scan_byte_range_set_scalar( +/// b"bytes=0-499", +/// 6, +/// )); +/// ``` +#[must_use] +pub fn scan_byte_range_set_scalar(bytes: &[u8], start: usize) -> bool { + range::scan_byte_range_set(bytes, start) +} + +/// Runs the unaccelerated comma-separated token list scan. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// # { +/// use http_headers_simd::{EmptyMembers, TokenListScan, benchmarking}; +/// +/// assert_eq!( +/// benchmarking::scan_token_list_scalar(b"gzip, br", EmptyMembers::Skip), +/// TokenListScan::Members, +/// ); +/// # } +/// ``` +#[must_use] +pub fn scan_token_list_scalar(bytes: &[u8], empty: EmptyMembers) -> TokenListScan { + list::scan_token_list(bytes, empty) +} + +/// Runs the SSE2 URI scanner directly when the current host supports it. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// # { +/// let result = http_headers_simd::benchmarking::is_simple_uri_path_sse2(b"/long/simple/path"); +/// assert!(result.is_none() || result == Some(true)); +/// # } +/// ``` +#[must_use] +pub fn is_simple_uri_path_sse2(bytes: &[u8]) -> Option { + #[cfg(all(feature = "std", any(target_arch = "x86", target_arch = "x86_64")))] + if bytes.len() >= 16 && uri::has_origin_relative_prefix(bytes) && arch::is_x86_feature_detected!("sse2") { + // SAFETY: runtime detection establishes SSE2 and the length check + // satisfies the scanner's trailing-block precondition. + return Some(unsafe { crate::x86::is_simple_uri_tail_sse2(bytes) }); + } + let _ = bytes; + None +} + +/// Runs the SSSE3 URI scanner directly when the current host supports it. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// # { +/// let result = http_headers_simd::benchmarking::is_simple_uri_path_ssse3(b"/long/simple/path"); +/// assert!(result.is_none() || result == Some(true)); +/// # } +/// ``` +#[must_use] +pub fn is_simple_uri_path_ssse3(bytes: &[u8]) -> Option { + #[cfg(all(feature = "std", any(target_arch = "x86", target_arch = "x86_64")))] + if bytes.len() >= 16 && uri::has_origin_relative_prefix(bytes) && arch::is_x86_feature_detected!("ssse3") { + // SAFETY: runtime detection establishes SSSE3 and the length check + // satisfies the scanner's trailing-block precondition. + return Some(unsafe { crate::x86::is_simple_uri_tail_ssse3(bytes) }); + } + let _ = bytes; + None +} + +/// Runs the SSE4.2 URI scanner directly when the current host supports it. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// # { +/// let result = http_headers_simd::benchmarking::is_simple_uri_path_sse42(b"/long/simple/path"); +/// assert!(result.is_none() || result == Some(true)); +/// # } +/// ``` +#[must_use] +pub fn is_simple_uri_path_sse42(bytes: &[u8]) -> Option { + #[cfg(all(feature = "std", any(target_arch = "x86", target_arch = "x86_64")))] + if bytes.len() >= 16 && uri::has_origin_relative_prefix(bytes) && arch::is_x86_feature_detected!("sse4.2") { + // SAFETY: runtime detection establishes SSE4.2 and the length + // check satisfies the scanner's trailing-block precondition. + return Some(unsafe { crate::x86::is_simple_uri_tail_sse42(bytes) }); + } + let _ = bytes; + None +} + +/// Runs the NEON URI scanner directly when the current host supports it. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(feature = "benchmarking")] +/// # { +/// let result = http_headers_simd::benchmarking::is_simple_uri_path_neon(b"/long/simple/path"); +/// assert!(result.is_none() || result == Some(true)); +/// # } +/// ``` +#[must_use] +pub fn is_simple_uri_path_neon(bytes: &[u8]) -> Option { + #[cfg(all(feature = "std", target_arch = "aarch64"))] + if bytes.len() >= 16 && uri::has_origin_relative_prefix(bytes) && arch::is_aarch64_feature_detected!("neon") { + // SAFETY: runtime detection establishes NEON and the length check + // satisfies the scanner's trailing-block precondition. + return Some(unsafe { crate::arm::is_simple_uri_tail(bytes) }); + } + let _ = bytes; + None +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #[cfg(all(not(feature = "std"), any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] + use std::arch; + #[cfg(not(feature = "std"))] + use std::vec; + + use super::*; + + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + fn assert_direct_uri_backends(path: &[u8], expected: bool) { + assert_eq!(is_simple_uri_path_sse2(path), Some(expected)); + assert_eq!( + is_simple_uri_path_ssse3(path), + arch::is_x86_feature_detected!("ssse3").then_some(expected) + ); + assert_eq!( + is_simple_uri_path_sse42(path), + arch::is_x86_feature_detected!("sse4.2").then_some(expected) + ); + assert_eq!(is_simple_uri_path_neon(path), None); + } + + #[cfg(target_arch = "aarch64")] + fn assert_direct_uri_backends(path: &[u8], expected: bool) { + assert_eq!(is_simple_uri_path_sse2(path), None); + assert_eq!(is_simple_uri_path_ssse3(path), None); + assert_eq!(is_simple_uri_path_sse42(path), None); + assert_eq!( + is_simple_uri_path_neon(path), + arch::is_aarch64_feature_detected!("neon").then_some(expected) + ); + } + + #[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] + fn assert_direct_uri_backends(path: &[u8], _expected: bool) { + assert_eq!(is_simple_uri_path_sse2(path), None); + assert_eq!(is_simple_uri_path_ssse3(path), None); + assert_eq!(is_simple_uri_path_sse42(path), None); + assert_eq!(is_simple_uri_path_neon(path), None); + } + + #[test] + fn scalar_entry_points_expose_the_reference_backends() { + assert!(is_token_scalar(b"token")); + assert!(!is_token_scalar(b"not a token")); + assert!(is_token68_scalar(b"abc+/==")); + assert!(!is_token68_scalar(b"=abc")); + assert!(is_field_value_scalar(&[b'\t', b' ', 0x80])); + assert!(!is_field_value_scalar(b"line\nbreak")); + assert_eq!(find_interesting_scalar(b"plain,value"), Some(5)); + assert!(all_base64_alphabet_scalar(b"AZaz09+/")); + assert!(!all_base64_alphabet_scalar(b"padding=")); + assert!(is_simple_uri_path_scalar(b"/path?query#fragment")); + assert!(!is_simple_uri_path_scalar(b"//authority/path")); + assert!(scan_byte_range_set_scalar(b"bytes=0-499", 6)); + assert!(!scan_byte_range_set_scalar(b"bytes=500-100", 6)); + assert_eq!(scan_token_list_scalar(b"gzip, deflate", EmptyMembers::Skip), TokenListScan::Members); + } + + #[test] + fn direct_uri_backends_validate_preconditions_and_match_scalar() { + let valid = vec![b'a'; 48]; + let mut path = valid; + path[0] = b'/'; + let expected = is_simple_uri_path_scalar(&path); + + assert_direct_uri_backends(&path, expected); + + assert_eq!(is_simple_uri_path_sse2(b"/short"), None); + assert_eq!(is_simple_uri_path_ssse3(b"relative-path-that-is-long"), None); + assert_eq!(is_simple_uri_path_sse42(b"//authority/path-that-is-long"), None); + assert_eq!(is_simple_uri_path_neon(b"/short"), None); + } +} diff --git a/crates/http_headers_simd/src/dispatch.rs b/crates/http_headers_simd/src/dispatch.rs new file mode 100644 index 000000000..bafd6a7b1 --- /dev/null +++ b/crates/http_headers_simd/src/dispatch.rs @@ -0,0 +1,1163 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Runtime and compile-time dispatch across scalar and SIMD scanners. + +use crate::list::{self, EmptyMembers, TokenListScan}; +use crate::range::{self, RangeMasks, WINDOW}; +use crate::{base64, scalar, uri}; + +/// The shortest input eligible for the shared byte-classification scanners. +/// +/// On an AMD EPYC 7763, x86-64 Criterion and Callgrind measurements at 16, +/// 17, and 31 bytes favored SIMD for token, `token68`, field-value, and +/// interesting-byte scans. Their 16-byte instruction counts fell by 40-63% +/// and wall-clock times by 31-65%. Architectures not measured here use a +/// two-vector crossover. +pub(super) const SIMD_THRESHOLD: usize = if cfg!(any(target_arch = "x86", target_arch = "x86_64")) { + 16 +} else { + 32 +}; + +/// The shortest input the ASCII case-insensitive equality scanner vectorizes. +/// +/// The same x86-64 measurements did not justify lowering this crossover: +/// SIMD saved 12 instructions at 16 and 17 bytes but cost 185 at 31 bytes. +/// Raising it to 48 saved work at 47 bytes but regressed 32- and 33-byte +/// cases, so 32 remains the best measured monotonic cutoff. `AArch64` uses 32 +/// because no measurements support another threshold. +const EQUALITY_SIMD_THRESHOLD: usize = 32; + +/// The shortest reference the URI subset scanner vectorizes. +/// +/// Redirect targets are usually short, and the scanner covers the trailing +/// bytes with one overlapping block rather than a scalar loop, so a single +/// vector's worth of input is already enough to pay for the dispatch. +const URI_SIMD_THRESHOLD: usize = 16; + +#[inline] +pub(super) fn is_token(bytes: &[u8]) -> bool { + if bytes.is_empty() { + return false; + } + if bytes.len() < SIMD_THRESHOLD { + return scalar::is_token(bytes); + } + finish_with(optimized_is_token(bytes), || scalar::is_token(bytes)) +} + +#[inline] +fn finish_with(optimized: Option, fallback: impl FnOnce() -> T) -> T { + optimized.unwrap_or_else(fallback) +} + +#[inline] +pub(super) fn is_token68(bytes: &[u8]) -> bool { + if bytes.len() < SIMD_THRESHOLD { + return scalar::is_token68(bytes); + } + finish_with(optimized_is_token68(bytes), || scalar::is_token68(bytes)) +} + +#[inline] +pub(super) fn is_field_value(bytes: &[u8]) -> bool { + if bytes.len() < SIMD_THRESHOLD { + return scalar::is_field_value(bytes); + } + finish_with(optimized_is_field_value(bytes), || scalar::is_field_value(bytes)) +} + +#[inline] +pub(super) fn eq_ignore_ascii_case(left: &[u8], right: &[u8]) -> bool { + if left.len() != right.len() { + return false; + } + if left.len() < EQUALITY_SIMD_THRESHOLD { + return scalar::eq_ignore_ascii_case(left, right); + } + finish_eq_ignore_ascii_case(left, right, optimized_eq_ignore_ascii_case(left, right)) +} + +#[inline] +fn finish_eq_ignore_ascii_case(left: &[u8], right: &[u8], optimized: Option) -> bool { + finish_with(optimized, || scalar::eq_ignore_ascii_case(left, right)) +} + +#[inline] +pub(super) fn find_interesting(bytes: &[u8]) -> Option { + if bytes.len() < SIMD_THRESHOLD { + return scalar::find_interesting(bytes); + } + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { find_interesting_x86(bytes, sse2_available()) } + } + #[cfg(target_arch = "aarch64")] + { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { find_interesting_arm(bytes, neon_available()) } + } + #[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] + { + scalar::find_interesting(bytes) + } +} + +#[inline] +pub(super) fn find_either(bytes: &[u8], first: u8, second: u8) -> Option { + if bytes.len() < SIMD_THRESHOLD { + return scalar::find_either(bytes, first, second); + } + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { find_either_x86(bytes, first, second, sse2_available()) } + } + #[cfg(target_arch = "aarch64")] + { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { find_either_arm(bytes, first, second, neon_available()) } + } + #[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] + scalar::find_either(bytes, first, second) +} + +/// Runs the `AArch64` byte-pair search with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON. +#[cfg(target_arch = "aarch64")] +unsafe fn find_either_arm(bytes: &[u8], first: u8, second: u8, neon: bool) -> Option { + if neon { + // SAFETY: the caller guarantees NEON support when `neon` is true. + return unsafe { crate::arm::find_either(bytes, first, second) }; + } + scalar::find_either(bytes, first, second) +} + +/// Runs the x86 byte-pair search with already-detected features. +/// +/// # Safety +/// +/// If `sse2` is true, the processor must support SSE2. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn find_either_x86(bytes: &[u8], first: u8, second: u8, sse2: bool) -> Option { + if sse2 { + // SAFETY: the caller guarantees SSE2 support when `sse2` is true. + return unsafe { crate::x86::find_either(bytes, first, second) }; + } + scalar::find_either(bytes, first, second) +} + +/// Runs `AArch64` interesting-byte dispatch with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON. +#[cfg(target_arch = "aarch64")] +unsafe fn find_interesting_arm(bytes: &[u8], neon: bool) -> Option { + if neon { + // SAFETY: the caller supplies the result of NEON feature detection. + return unsafe { crate::arm::find_interesting(bytes) }; + } + scalar::find_interesting(bytes) +} + +/// Runs the x86 interesting-byte dispatch with already-detected features. +/// +/// # Safety +/// +/// If `sse2` is true, the processor must support SSE2. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn find_interesting_x86(bytes: &[u8], sse2: bool) -> Option { + if sse2 { + // SAFETY: the caller guarantees SSE2 support when `sse2` is true. + return unsafe { crate::x86::find_interesting(bytes) }; + } + scalar::find_interesting(bytes) +} + +#[inline] +pub(super) fn all_base64_alphabet(bytes: &[u8]) -> bool { + if bytes.len() < URI_SIMD_THRESHOLD { + return base64::all_base64_alphabet(bytes); + } + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + { + // SAFETY: the feature flags come from runtime or compile-time detection, + // and the length check satisfies both backends' trailing-block requirement. + unsafe { all_base64_alphabet_x86(bytes, sse42_available(), sse2_available()) } + } + #[cfg(target_arch = "aarch64")] + { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { all_base64_alphabet_arm(bytes, neon_available()) } + } + #[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] + { + base64::all_base64_alphabet(bytes) + } +} + +/// Runs `AArch64` base64 dispatch with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON. +#[cfg(target_arch = "aarch64")] +unsafe fn all_base64_alphabet_arm(bytes: &[u8], neon: bool) -> bool { + if neon { + // SAFETY: the caller supplies the result of NEON feature detection. + return unsafe { crate::arm::all_base64_alphabet(bytes) }; + } + base64::all_base64_alphabet(bytes) +} + +/// Runs x86 base64 dispatch with already-detected features. +/// +/// # Safety +/// +/// Every true feature flag must name an instruction set the processor supports, +/// and `bytes` must hold at least one whole vector. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn all_base64_alphabet_x86(bytes: &[u8], packed_ranges: bool, baseline: bool) -> bool { + if packed_ranges { + // SAFETY: the caller guarantees SSE4.2 support and a whole input vector. + return unsafe { crate::x86::all_base64_alphabet_sse42(bytes) }; + } + if baseline { + // SAFETY: the caller guarantees SSE2 support and a whole input vector. + return unsafe { crate::x86::all_base64_alphabet_sse2(bytes) }; + } + base64::all_base64_alphabet(bytes) +} + +/// The shortest field line the token-list scanner vectorizes. +/// +/// The scanner folds whole blocks and leaves the trailing bytes to the scalar +/// state machine, so a line shorter than one vector would pay for dispatch and +/// then run the scalar scan anyway. +const LIST_SIMD_THRESHOLD: usize = WIDTH; + +/// The number of bytes one accelerated block covers on every architecture. +const WIDTH: usize = 16; + +/// The longest line a single fold covers, which is two accelerated blocks. +/// +/// The grammar costs more to fold than the class masks cost to build, so a +/// line this short is classified as two overlapping vectors and folded once, +/// while longer lines run the block loop. +const LIST_SHORT_LIMIT: usize = 2 * WIDTH; + +#[inline] +pub(super) fn scan_token_list(bytes: &[u8], empty: EmptyMembers) -> TokenListScan { + if bytes.len() < LIST_SIMD_THRESHOLD { + return list::scan_token_list(bytes, empty); + } + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + { + // SAFETY: the feature flags come from runtime or compile-time detection, + // and the input length satisfies every selected scanner's precondition. + unsafe { scan_token_list_x86(bytes, empty, sse42_available(), sse2_available()) } + } + #[cfg(target_arch = "aarch64")] + { + // SAFETY: the feature flag comes from runtime or compile-time detection, + // and the input contains at least one whole vector. + unsafe { scan_token_list_arm(bytes, empty, neon_available()) } + } + #[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] + { + list::scan_token_list(bytes, empty) + } +} + +/// Runs `AArch64` token-list dispatch with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON and `bytes` must hold at +/// least one whole vector. +#[cfg(target_arch = "aarch64")] +unsafe fn scan_token_list_arm(bytes: &[u8], empty: EmptyMembers, neon: bool) -> TokenListScan { + if neon { + if bytes.len() <= LIST_SHORT_LIMIT { + if bytes.len() == WIDTH { + // SAFETY: the caller establishes NEON support and the line is + // exactly one whole vector. + return unsafe { crate::arm::scan_short_token_list::(bytes, empty) }; + } + // SAFETY: the caller establishes NEON support and the length checks + // bound the line to two whole vectors. + return unsafe { crate::arm::scan_short_token_list::(bytes, empty) }; + } + // SAFETY: the caller establishes NEON support and the length check in + // `scan_token_list` satisfies the trailing-block requirement. + return unsafe { crate::arm::scan_token_list(bytes, empty) }; + } + list::scan_token_list(bytes, empty) +} + +/// Runs x86 token-list dispatch with already-detected features. +/// +/// # Safety +/// +/// Every true feature flag must name an instruction set the processor supports, +/// and `bytes` must hold at least one whole vector. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn scan_token_list_x86(bytes: &[u8], empty: EmptyMembers, packed_ranges: bool, baseline: bool) -> TokenListScan { + if packed_ranges { + if bytes.len() <= LIST_SHORT_LIMIT { + if bytes.len() == WIDTH { + // SAFETY: the caller guarantees SSE4.2 and exactly one vector. + return unsafe { crate::x86::scan_short_token_list_sse42(bytes, empty, false) }; + } + // SAFETY: the caller guarantees SSE4.2 and at most two vectors. + return unsafe { crate::x86::scan_short_token_list_sse42(bytes, empty, true) }; + } + // SAFETY: the caller guarantees SSE4.2 and at least one vector. + return unsafe { crate::x86::scan_token_list_sse42(bytes, empty) }; + } + if baseline { + if bytes.len() <= LIST_SHORT_LIMIT { + if bytes.len() == WIDTH { + // SAFETY: the caller guarantees SSE2 and exactly one vector. + return unsafe { crate::x86::scan_short_token_list_sse2(bytes, empty, false) }; + } + // SAFETY: the caller guarantees SSE2 and at most two vectors. + return unsafe { crate::x86::scan_short_token_list_sse2(bytes, empty, true) }; + } + // SAFETY: the caller guarantees SSE2 and at least one vector. + return unsafe { crate::x86::scan_token_list_sse2(bytes, empty) }; + } + list::scan_token_list(bytes, empty) +} + +/// Scans a `bytes=` field line for a range set that needs no parsing. +#[inline] +pub(super) fn scan_byte_range_set(bytes: &[u8], start: usize) -> bool { + let Some(len) = bytes.len().checked_sub(start) else { + return false; + }; + // A payload of nothing and a payload past one window are both out of the + // subset, and one wrapping subtraction folds the pair into one comparison. + if len.wrapping_sub(1) >= WINDOW { + return false; + } + let Some(window) = range::window_of(bytes) else { + return short_byte_range_set(bytes, len); + }; + range::accepts(window, &classify_range(window), len) +} + +/// Scans a field line shorter than one window by padding it into one. +#[cold] +#[inline(never)] +fn short_byte_range_set(bytes: &[u8], len: usize) -> bool { + let mut padded = [range::FILLER; WINDOW]; + if !range::fill_window(bytes, &mut padded) { + return false; + } + range::accepts(&padded, &classify_range(&padded), len) +} + +/// Classifies one range window with the widest available kernel. +#[inline] +fn classify_range(window: &[u8; WINDOW]) -> RangeMasks { + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { classify_range_x86(window, sse2_available()) } + } + #[cfg(target_arch = "aarch64")] + { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { classify_range_arm(window, neon_available()) } + } + #[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] + { + RangeMasks::scalar(window) + } +} + +/// Runs `AArch64` range classification with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON. +#[cfg(target_arch = "aarch64")] +unsafe fn classify_range_arm(window: &[u8; WINDOW], neon: bool) -> RangeMasks { + if neon { + // SAFETY: the caller supplies the result of NEON feature detection. + return unsafe { crate::arm::range_masks(window) }; + } + RangeMasks::scalar(window) +} + +/// Runs x86 range classification with already-detected features. +/// +/// # Safety +/// +/// If `sse2` is true, the processor must support SSE2. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn classify_range_x86(window: &[u8; WINDOW], sse2: bool) -> RangeMasks { + if sse2 { + // SAFETY: the caller guarantees SSE2 support when `sse2` is true. + return unsafe { crate::x86::range_masks(window) }; + } + RangeMasks::scalar(window) +} + +#[inline] +pub(super) fn is_simple_uri_path(bytes: &[u8]) -> bool { + if !uri::has_origin_relative_prefix(bytes) { + return false; + } + simple_uri_tail(bytes) +} + +#[inline] +pub(super) fn is_simple_uri_reference(bytes: &[u8]) -> bool { + if uri::has_origin_relative_prefix(bytes) { + return simple_uri_tail(bytes); + } + match uri::simple_authority_end(bytes) { + Some(end) => simple_uri_tail(&bytes[end..]), + None => false, + } +} + +/// Classifies every byte of a path, query, and fragment run. +fn simple_uri_tail(bytes: &[u8]) -> bool { + if bytes.len() < URI_SIMD_THRESHOLD { + return uri::simple_uri_tail(bytes, 0); + } + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + { + // SAFETY: the feature flags come from runtime or compile-time detection, + // and the length check satisfies every backend's trailing-block requirement. + unsafe { is_simple_uri_path_x86(bytes, sse42_available(), ssse3_available(), sse2_available()) } + } + #[cfg(target_arch = "aarch64")] + { + // SAFETY: the feature flag comes from runtime or compile-time detection, + // and the input contains at least one whole vector. + unsafe { is_simple_uri_path_arm(bytes, neon_available()) } + } + #[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] + { + uri::simple_uri_tail(bytes, 0) + } +} + +/// Runs `AArch64` URI dispatch with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON and `bytes` must hold at +/// least one whole vector. +#[cfg(target_arch = "aarch64")] +unsafe fn is_simple_uri_path_arm(bytes: &[u8], neon: bool) -> bool { + if neon { + // SAFETY: the caller supplies the result of NEON feature detection. + return unsafe { crate::arm::is_simple_uri_tail(bytes) }; + } + uri::simple_uri_tail(bytes, 0) +} + +/// Runs x86 URI dispatch with already-detected features. +/// +/// # Safety +/// +/// Every true feature flag must name an instruction set the processor supports, +/// and `bytes` must hold at least one whole vector. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn is_simple_uri_path_x86(bytes: &[u8], packed_ranges: bool, shuffle_table: bool, baseline: bool) -> bool { + if packed_ranges { + // SAFETY: the caller guarantees SSE4.2 and at least one vector. + return unsafe { crate::x86::is_simple_uri_tail_sse42(bytes) }; + } + if shuffle_table { + // SAFETY: the caller guarantees SSSE3 and at least one vector. + return unsafe { crate::x86::is_simple_uri_tail_ssse3(bytes) }; + } + if baseline { + // SAFETY: the caller guarantees SSE2 and at least one vector. + return unsafe { crate::x86::is_simple_uri_tail_sse2(bytes) }; + } + uri::simple_uri_tail(bytes, 0) +} + +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +fn optimized_is_token(bytes: &[u8]) -> Option { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { optimized_is_token_x86(bytes, sse2_available()) } +} + +/// Runs x86 token dispatch with already-detected features. +/// +/// # Safety +/// +/// If `sse2` is true, the processor must support SSE2. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn optimized_is_token_x86(bytes: &[u8], sse2: bool) -> Option { + if sse2 { + // SAFETY: the caller guarantees SSE2 support when `sse2` is true. + return Some(unsafe { crate::x86::is_token(bytes) }); + } + None +} + +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +fn optimized_is_token68(bytes: &[u8]) -> Option { + // SAFETY: the feature flags come from runtime or compile-time detection. + unsafe { optimized_is_token68_x86(bytes, sse42_available(), sse2_available()) } +} + +/// Runs x86 `token68` dispatch with already-detected features. +/// +/// # Safety +/// +/// Every true feature flag must name an instruction set the processor supports. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn optimized_is_token68_x86(bytes: &[u8], packed_ranges: bool, baseline: bool) -> Option { + if packed_ranges { + // SAFETY: the caller guarantees SSE4.2 support when `sse42` is true. + return Some(unsafe { crate::x86::is_token68_sse42(bytes) }); + } + if baseline { + // SAFETY: the caller guarantees SSE2 support when `sse2` is true. + return Some(unsafe { crate::x86::is_token68_sse2(bytes) }); + } + None +} + +#[cfg(target_arch = "aarch64")] +fn optimized_is_token68(bytes: &[u8]) -> Option { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { optimized_is_token68_arm(bytes, neon_available()) } +} + +/// Runs `AArch64` `token68` dispatch with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON. +#[cfg(target_arch = "aarch64")] +unsafe fn optimized_is_token68_arm(bytes: &[u8], neon: bool) -> Option { + if neon { + // SAFETY: the caller supplies the result of NEON feature detection. + return Some(unsafe { crate::arm::is_token68(bytes) }); + } + None +} + +#[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] +fn optimized_is_token68(_: &[u8]) -> Option { + None +} + +#[cfg(target_arch = "aarch64")] +fn optimized_is_token(bytes: &[u8]) -> Option { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { optimized_is_token_arm(bytes, neon_available()) } +} + +/// Runs `AArch64` token dispatch with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON. +#[cfg(target_arch = "aarch64")] +unsafe fn optimized_is_token_arm(bytes: &[u8], neon: bool) -> Option { + if neon { + // SAFETY: the caller supplies the result of NEON feature detection. + return Some(unsafe { crate::arm::is_token(bytes) }); + } + None +} + +#[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] +fn optimized_is_token(_: &[u8]) -> Option { + None +} + +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +fn optimized_is_field_value(bytes: &[u8]) -> Option { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { optimized_is_field_value_x86(bytes, sse2_available()) } +} + +/// Runs x86 field-value dispatch with already-detected features. +/// +/// # Safety +/// +/// If `sse2` is true, the processor must support SSE2. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn optimized_is_field_value_x86(bytes: &[u8], sse2: bool) -> Option { + if sse2 { + // SAFETY: the caller guarantees SSE2 support when `sse2` is true. + return Some(unsafe { crate::x86::is_field_value(bytes) }); + } + None +} + +#[cfg(target_arch = "aarch64")] +fn optimized_is_field_value(bytes: &[u8]) -> Option { + // SAFETY: the feature flag comes from runtime or compile-time detection. + unsafe { optimized_is_field_value_arm(bytes, neon_available()) } +} + +/// Runs `AArch64` field-value dispatch with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON. +#[cfg(target_arch = "aarch64")] +unsafe fn optimized_is_field_value_arm(bytes: &[u8], neon: bool) -> Option { + if neon { + // SAFETY: the caller supplies the result of NEON feature detection. + return Some(unsafe { crate::arm::is_field_value(bytes) }); + } + None +} + +#[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] +fn optimized_is_field_value(_: &[u8]) -> Option { + None +} + +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +fn optimized_eq_ignore_ascii_case(left: &[u8], right: &[u8]) -> Option { + // SAFETY: the feature flag comes from runtime or compile-time detection, + // and the public dispatcher establishes equal lengths. + unsafe { optimized_eq_ignore_ascii_case_x86(left, right, sse2_available()) } +} + +/// Runs x86 ASCII-case dispatch with already-detected features. +/// +/// # Safety +/// +/// If `sse2` is true, the processor must support SSE2. `left` and `right` must +/// have equal lengths. +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +unsafe fn optimized_eq_ignore_ascii_case_x86(left: &[u8], right: &[u8], sse2: bool) -> Option { + if sse2 { + // SAFETY: the caller guarantees SSE2 support when `sse2` is true and + // equal-length slices. + return Some(unsafe { crate::x86::eq_ignore_ascii_case(left, right) }); + } + None +} +#[cfg(target_arch = "aarch64")] +fn optimized_eq_ignore_ascii_case(left: &[u8], right: &[u8]) -> Option { + // SAFETY: the feature flag comes from runtime or compile-time detection, + // and the public dispatcher establishes equal lengths. + unsafe { optimized_eq_ignore_ascii_case_arm(left, right, neon_available()) } +} + +/// Runs `AArch64` ASCII-case dispatch with already-detected features. +/// +/// # Safety +/// +/// If `neon` is true, the processor must support NEON. `left` and `right` must +/// have equal lengths. +#[cfg(target_arch = "aarch64")] +unsafe fn optimized_eq_ignore_ascii_case_arm(left: &[u8], right: &[u8], neon: bool) -> Option { + if neon { + // SAFETY: the caller supplies the result of NEON feature detection. + return Some(unsafe { crate::arm::eq_ignore_ascii_case(left, right) }); + } + None +} + +#[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] +fn optimized_eq_ignore_ascii_case(_: &[u8], _: &[u8]) -> Option { + None +} + +#[cfg(target_arch = "x86_64")] +const fn sse2_available() -> bool { + true +} + +#[cfg(target_arch = "x86")] +fn sse2_available() -> bool { + #[cfg(feature = "std")] + { + use std::arch; + + arch::is_x86_feature_detected!("sse2") + } + #[cfg(not(feature = "std"))] + { + cfg!(target_feature = "sse2") + } +} + +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +fn ssse3_available() -> bool { + #[cfg(target_feature = "ssse3")] + { + true + } + #[cfg(all(not(target_feature = "ssse3"), feature = "std"))] + { + use std::arch; + + arch::is_x86_feature_detected!("ssse3") + } + #[cfg(all(not(target_feature = "ssse3"), not(feature = "std")))] + { + cached_ssse3_available() + } +} + +/// Caches SSSE3 detection where `std` feature detection is unavailable. +#[cfg(all( + any(target_arch = "x86", target_arch = "x86_64"), + not(feature = "std"), + not(target_feature = "ssse3") +))] +fn cached_ssse3_available() -> bool { + use core::sync::atomic::{AtomicU8, Ordering}; + + static AVAILABLE: AtomicU8 = AtomicU8::new(0); + let cached = AVAILABLE.load(Ordering::Relaxed); + if cached != 0 { + return cached == 2; + } + #[cfg(target_arch = "x86")] + let features = core::arch::x86::__cpuid(1); + #[cfg(target_arch = "x86_64")] + let features = core::arch::x86_64::__cpuid(1); + const SSSE3_BIT: u32 = 1 << 9; + let available = features.ecx & SSSE3_BIT != 0; + AVAILABLE.store(if available { 2 } else { 1 }, Ordering::Relaxed); + available +} + +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +fn sse42_available() -> bool { + #[cfg(target_feature = "sse4.2")] + { + true + } + #[cfg(all(not(target_feature = "sse4.2"), feature = "std"))] + { + use std::arch; + + arch::is_x86_feature_detected!("sse4.2") + } + #[cfg(all(not(target_feature = "sse4.2"), not(feature = "std")))] + { + cached_sse42_available() + } +} + +#[cfg(all( + any(target_arch = "x86", target_arch = "x86_64"), + not(feature = "std"), + not(target_feature = "sse4.2") +))] +fn cached_sse42_available() -> bool { + use core::sync::atomic::{AtomicU8, Ordering}; + + static AVAILABLE: AtomicU8 = AtomicU8::new(0); + let cached = AVAILABLE.load(Ordering::Relaxed); + if cached != 0 { + return cached == 2; + } + #[cfg(target_arch = "x86")] + let features = core::arch::x86::__cpuid(1); + #[cfg(target_arch = "x86_64")] + let features = core::arch::x86_64::__cpuid(1); + const SSE42_BIT: u32 = 1 << 20; + let available = features.ecx & SSE42_BIT != 0; + AVAILABLE.store(if available { 2 } else { 1 }, Ordering::Relaxed); + available +} + +#[cfg(target_arch = "aarch64")] +fn neon_available() -> bool { + #[cfg(feature = "std")] + { + use std::arch; + + arch::is_aarch64_feature_detected!("neon") + } + #[cfg(not(feature = "std"))] + { + cfg!(target_feature = "neon") + } +} + +#[cfg(any(feature = "benchmarking", feature = "test-util"))] +pub(super) fn backend() -> crate::Backend { + backend_for(SIMD_THRESHOLD) +} + +#[cfg(any(feature = "benchmarking", feature = "test-util"))] +pub(super) fn backend_for(len: usize) -> crate::Backend { + if len < SIMD_THRESHOLD { + return crate::Backend::Scalar; + } + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + { + backend_for_x86(sse42_available(), sse2_available()) + } + #[cfg(target_arch = "aarch64")] + { + backend_for_arm(neon_available()) + } + #[cfg(not(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64")))] + { + crate::Backend::Scalar + } +} + +#[cfg(all(target_arch = "aarch64", any(feature = "benchmarking", feature = "test-util")))] +const fn backend_for_arm(neon: bool) -> crate::Backend { + if neon { crate::Backend::Neon } else { crate::Backend::Scalar } +} + +#[cfg(all( + any(target_arch = "x86", target_arch = "x86_64"), + any(feature = "benchmarking", feature = "test-util") +))] +fn backend_for_x86(packed_ranges: bool, baseline: bool) -> crate::Backend { + if packed_ranges { + return crate::Backend::Sse42; + } + if baseline { + return crate::Backend::Sse2; + } + crate::Backend::Scalar +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + use std::arch; + #[cfg(not(feature = "std"))] + use std::vec; + + use crate::list::{EmptyMembers, TokenListScan}; + + #[test] + fn dispatch_covers_scalar_and_vector_length_boundaries() { + assert!(!super::is_token(b"")); + assert!(super::is_token(b"short-token")); + assert!(!super::is_token(b"short token")); + + let valid = vec![b'a'; super::SIMD_THRESHOLD + 17]; + assert!(super::is_token(&valid)); + assert!(super::is_token68(&valid)); + assert!(super::is_field_value(&valid)); + assert!(super::eq_ignore_ascii_case(&valid, &valid)); + assert_eq!(super::find_interesting(&valid), None); + assert!(super::all_base64_alphabet(&valid)); + + let mut interesting = valid; + interesting[super::SIMD_THRESHOLD + 3] = b';'; + assert_eq!(super::find_interesting(&interesting), Some(super::SIMD_THRESHOLD + 3)); + } + + #[test] + fn tuned_scanners_match_scalar_around_both_crossovers() { + for length in [15, 16, 17, 31, 32, 33, 47, 48] { + let valid = vec![b'a'; length]; + let folded = vec![b'A'; length]; + + assert_eq!(super::is_token(&valid), crate::scalar::is_token(&valid)); + assert_eq!(super::is_token68(&valid), crate::scalar::is_token68(&valid)); + assert_eq!(super::is_field_value(&valid), crate::scalar::is_field_value(&valid)); + assert_eq!( + super::eq_ignore_ascii_case(&valid, &folded), + crate::scalar::eq_ignore_ascii_case(&valid, &folded) + ); + assert_eq!(super::find_interesting(&valid), crate::scalar::find_interesting(&valid)); + + let mut rejected = valid.clone(); + rejected[length - 1] = b' '; + assert_eq!(super::is_token(&rejected), crate::scalar::is_token(&rejected)); + assert_eq!(super::is_token68(&rejected), crate::scalar::is_token68(&rejected)); + assert_eq!(super::find_interesting(&rejected), crate::scalar::find_interesting(&rejected)); + + rejected[length - 1] = b'\n'; + assert_eq!(super::is_field_value(&rejected), crate::scalar::is_field_value(&rejected)); + + let mut unequal = folded; + unequal[length - 1] = b'b'; + assert_eq!( + super::eq_ignore_ascii_case(&valid, &unequal), + crate::scalar::eq_ignore_ascii_case(&valid, &unequal) + ); + } + } + + #[test] + fn dispatch_rejects_invalid_range_offsets_and_scans_list_shapes() { + assert!(!super::scan_byte_range_set(b"bytes=0-1", 64)); + assert!(!super::scan_byte_range_set(b"bytes=", 6)); + assert!(!super::scan_byte_range_set(b"1-2", 0)); + assert!(super::scan_byte_range_set(b"bytes=1-2", 6)); + assert!(super::scan_byte_range_set(b"bytes=1234-5678", 6)); + + for length in [16, 17, 32, 33, 48] { + let input = vec![b'a'; length]; + assert_eq!(super::scan_token_list(&input, EmptyMembers::Skip), TokenListScan::Members); + } + let mut rejected = vec![b'a'; 48]; + rejected[31] = b' '; + assert_eq!(super::scan_token_list(&rejected, EmptyMembers::Skip), TokenListScan::Rejected); + } + + #[test] + fn uri_dispatch_checks_prefix_short_tail_and_vector_tail() { + assert!(!super::is_simple_uri_path(b"relative")); + assert!(super::is_simple_uri_path(b"/short#tail")); + + let mut path = vec![b'a'; 49]; + path[0] = b'/'; + assert!(super::is_simple_uri_path(&path)); + path[40] = b'%'; + assert!(!super::is_simple_uri_path(&path)); + } + + #[cfg(any(feature = "benchmarking", feature = "test-util"))] + #[test] + fn backend_selection_respects_the_scalar_threshold() { + assert_eq!(super::backend_for(super::SIMD_THRESHOLD - 1), crate::Backend::Scalar); + assert_eq!(super::backend(), super::backend_for(super::SIMD_THRESHOLD)); + } + + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + #[test] + fn x86_feature_detection_matches_compile_time_or_runtime_support() { + #[cfg(target_arch = "x86_64")] + assert!(super::sse2_available()); + #[cfg(target_arch = "x86")] + assert_eq!( + super::sse2_available(), + cfg!(target_feature = "sse2") || cfg!(feature = "std") && arch::is_x86_feature_detected!("sse2") + ); + assert_eq!( + super::ssse3_available(), + cfg!(target_feature = "ssse3") || arch::is_x86_feature_detected!("ssse3") + ); + assert_eq!( + super::sse42_available(), + cfg!(target_feature = "sse4.2") || arch::is_x86_feature_detected!("sse4.2") + ); + } + + #[cfg(any(target_arch = "x86", target_arch = "x86_64"))] + #[test] + fn lower_x86_dispatch_tiers_and_scalar_fallbacks_are_directly_tested() { + #[cfg(target_arch = "x86")] + if !arch::is_x86_feature_detected!("sse2") { + return; + } + + let long = [b'a'; 48]; + assert!(super::finish_with(None, || true)); + assert!(!super::finish_with(Some(false), || true)); + assert!(super::finish_eq_ignore_ascii_case(&long, &long, None)); + + let mut interesting = long; + interesting[37] = b','; + assert_eq!( + // SAFETY: x86-64 guarantees SSE2; the x86 guard above detects it. + unsafe { super::find_interesting_x86(&interesting, true) }, + Some(37) + ); + assert_eq!( + // SAFETY: no feature instruction is used when the flag is false. + unsafe { super::find_interesting_x86(&interesting, false) }, + Some(37) + ); + assert_eq!( + // SAFETY: x86-64 guarantees SSE2; the x86 guard above detects it. + unsafe { super::find_either_x86(&interesting, b',', b'"', true) }, + Some(37) + ); + assert_eq!( + // SAFETY: no feature instruction is used when the flag is false. + unsafe { super::find_either_x86(&interesting, b',', b'"', false) }, + Some(37) + ); + + assert!( + // SAFETY: x86-64 guarantees SSE2; the x86 guard above detects it, + // and the input holds at least one whole vector. + unsafe { super::all_base64_alphabet_x86(&long, false, true) } + ); + assert!( + // SAFETY: false feature flags select the scalar implementation. + unsafe { super::all_base64_alphabet_x86(&long, false, false) } + ); + + for input in [&[b'a'; 16][..], &[b'a'; 17][..], &[b'a'; 33][..]] { + assert_eq!( + // SAFETY: x86-64 guarantees SSE2; the x86 guard above detects it, + // and every input holds at least one whole vector. + unsafe { super::scan_token_list_x86(input, EmptyMembers::Skip, false, true) }, + TokenListScan::Members + ); + assert_eq!( + // SAFETY: false feature flags select the scalar implementation. + unsafe { super::scan_token_list_x86(input, EmptyMembers::Skip, false, false) }, + TokenListScan::Members + ); + } + + let window = *b"0123456789-, \txy"; + assert_eq!( + // SAFETY: x86-64 guarantees SSE2; the x86 guard above detects it. + unsafe { super::classify_range_x86(&window, true) }, + crate::range::RangeMasks::scalar(&window) + ); + assert_eq!( + // SAFETY: a false feature flag selects scalar classification. + unsafe { super::classify_range_x86(&window, false) }, + crate::range::RangeMasks::scalar(&window) + ); + + assert!( + // SAFETY: x86-64 guarantees SSE2; the x86 guard above detects it, + // and the input holds at least one whole vector. + unsafe { super::is_simple_uri_path_x86(&long, false, false, true) } + ); + assert!( + // SAFETY: false feature flags select the scalar implementation. + unsafe { super::is_simple_uri_path_x86(&long, false, false, false) } + ); + let ssse3 = arch::is_x86_feature_detected!("ssse3"); + assert_eq!( + ssse3.then(|| { + // SAFETY: runtime detection establishes SSSE3 support and the input + // holds at least one whole vector. + unsafe { super::is_simple_uri_path_x86(&long, false, true, true) } + }), + ssse3.then_some(true) + ); + + assert_eq!( + // SAFETY: a false feature flag selects no optimized token backend. + unsafe { super::optimized_is_token_x86(&long, false) }, + None + ); + assert_eq!( + // SAFETY: x86-64 guarantees SSE2; the x86 guard above detects it. + unsafe { super::optimized_is_token68_x86(&long, false, true) }, + Some(true) + ); + assert_eq!( + // SAFETY: false feature flags select no optimized `token68` backend. + unsafe { super::optimized_is_token68_x86(&long, false, false) }, + None + ); + assert_eq!( + // SAFETY: a false feature flag selects no optimized field-value backend. + unsafe { super::optimized_is_field_value_x86(&long, false) }, + None + ); + assert_eq!( + // SAFETY: a false feature flag selects no optimized case-folding backend. + unsafe { super::optimized_eq_ignore_ascii_case_x86(&long, &long, false) }, + None + ); + + #[cfg(any(feature = "benchmarking", feature = "test-util"))] + { + assert_eq!(super::backend_for_x86(false, true), crate::Backend::Sse2); + assert_eq!(super::backend_for_x86(false, false), crate::Backend::Scalar); + } + } + + #[cfg(all( + not(feature = "std"), + any(target_arch = "x86", target_arch = "x86_64"), + not(target_feature = "ssse3") + ))] + #[test] + fn no_std_ssse3_detection_is_stable_after_caching() { + let expected = arch::is_x86_feature_detected!("ssse3"); + assert_eq!(super::ssse3_available(), expected); + for _ in 0..8 { + assert_eq!(super::ssse3_available(), expected); + } + } + + #[cfg(all( + not(feature = "std"), + any(target_arch = "x86", target_arch = "x86_64"), + not(target_feature = "sse4.2") + ))] + #[test] + fn no_std_sse42_detection_is_stable_after_caching() { + let expected = arch::is_x86_feature_detected!("sse4.2"); + assert_eq!(super::sse42_available(), expected); + for _ in 0..8 { + assert_eq!(super::sse42_available(), expected); + } + } + + #[cfg(all(not(feature = "std"), target_arch = "x86"))] + #[test] + fn no_std_sse2_detection_matches_target_features() { + assert_eq!(super::sse2_available(), cfg!(target_feature = "sse2")); + } + + #[cfg(all(not(feature = "std"), target_arch = "aarch64"))] + #[test] + fn no_std_neon_detection_matches_target_features() { + assert_eq!(super::neon_available(), cfg!(target_feature = "neon")); + } + + #[cfg(target_arch = "aarch64")] + #[test] + fn arm_scalar_fallbacks_are_directly_tested() { + let long = [b'a'; 48]; + let mut interesting = long; + interesting[37] = b','; + // SAFETY: the false feature flag selects only the scalar fallback. + assert_eq!(unsafe { super::find_interesting_arm(&interesting, false) }, Some(37)); + // SAFETY: the false feature flag selects only the scalar fallback. + assert_eq!(unsafe { super::find_either_arm(&interesting, b',', b'"', false) }, Some(37)); + // SAFETY: the false feature flag selects only the scalar fallback. + assert!(unsafe { super::all_base64_alphabet_arm(&long, false) }); + + for input in [&[b'a'; 16][..], &[b'a'; 17][..], &[b'a'; 33][..]] { + // SAFETY: the false feature flag selects only the scalar fallback, + // and every input contains at least one whole vector. + let scan = unsafe { super::scan_token_list_arm(input, EmptyMembers::Skip, false) }; + assert_eq!(scan, TokenListScan::Members); + } + + let window = *b"0123456789-, \txy"; + // SAFETY: the false feature flag selects only the scalar fallback. + let masks = unsafe { super::classify_range_arm(&window, false) }; + assert_eq!(masks, crate::range::RangeMasks::scalar(&window)); + // SAFETY: the false feature flag selects only the scalar fallback, + // and the input contains at least one whole vector. + assert!(unsafe { super::is_simple_uri_path_arm(&long, false) }); + // SAFETY: the false feature flag selects only the scalar fallback. + assert_eq!(unsafe { super::optimized_is_token_arm(&long, false) }, None); + // SAFETY: the false feature flag selects only the scalar fallback. + assert_eq!(unsafe { super::optimized_is_token68_arm(&long, false) }, None); + // SAFETY: the false feature flag selects only the scalar fallback. + assert_eq!(unsafe { super::optimized_is_field_value_arm(&long, false) }, None); + // SAFETY: the false feature flag selects only the scalar fallback, and + // the slices have equal lengths. + assert_eq!(unsafe { super::optimized_eq_ignore_ascii_case_arm(&long, &long, false) }, None); + + #[cfg(any(feature = "benchmarking", feature = "test-util"))] + assert_eq!(super::backend_for_arm(false), crate::Backend::Scalar); + } +} diff --git a/crates/http_headers_simd/src/lib.rs b/crates/http_headers_simd/src/lib.rs new file mode 100644 index 000000000..084fb85b7 --- /dev/null +++ b/crates/http_headers_simd/src/lib.rs @@ -0,0 +1,61 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +#![cfg_attr(not(feature = "std"), no_std)] +#![cfg_attr(coverage_nightly, feature(coverage_attribute))] +#![deny(unsafe_op_in_unsafe_fn)] +#![doc(hidden)] +#![doc(html_logo_url = "https://media.githubusercontent.com/media/microsoft/oxidizer/refs/heads/main/crates/http_headers_simd/logo.png")] +#![doc( + html_favicon_url = "https://media.githubusercontent.com/media/microsoft/oxidizer/refs/heads/main/crates/http_headers_simd/favicon.ico" +)] + +//! SIMD implementation details for the +//! [`http_headers`](https://docs.rs/http_headers) crate. +//! +//! **Do not depend on this crate directly.** Use `http_headers` instead. +//! +//! The default `std` feature enables runtime CPU-feature detection. +//! With default features disabled, x86 and x86-64 use compile-time features +//! plus cached runtime `CPUID` detection for SSSE3 and SSE4.2 when those features +//! are not enabled at compile time. SSE2 is guaranteed on x86-64 and requires +//! compile-time support on x86. `AArch64` NEON availability follows compile-time +//! target features. Unavailable accelerated paths fall back to scalar code. +//! +//! The `benchmarking` and `test-util` features expose unstable repository +//! instrumentation only. + +#[cfg(test)] +extern crate std; + +mod api; +mod base64; +mod dispatch; +mod list; +mod range; +mod scalar; +mod uri; + +#[cfg(any(target_arch = "x86", target_arch = "x86_64"))] +mod x86; + +#[cfg(target_arch = "aarch64")] +mod arm; + +#[cfg(all(feature = "benchmarking", feature = "std"))] +pub mod tracking; + +/// Direct reference-backend access for differential benchmarks. +#[cfg(feature = "benchmarking")] +pub mod benchmarking; + +#[cfg(any(feature = "benchmarking", feature = "test-util"))] +#[doc(inline)] +pub use api::{Backend, backend, backend_for, simd_threshold}; +#[doc(inline)] +pub use api::{ + all_base64_alphabet, as_simple_uri_reference, ascii_str, eq_ignore_ascii_case, find_either, find_interesting, is_field_value, + is_simple_uri_path, is_token, is_token68, scan_byte_range_set, scan_token_list, +}; +#[doc(inline)] +pub use list::{EmptyMembers, TokenListScan}; diff --git a/crates/http_headers_simd/src/list.rs b/crates/http_headers_simd/src/list.rs new file mode 100644 index 000000000..92f4fba90 --- /dev/null +++ b/crates/http_headers_simd/src/list.rs @@ -0,0 +1,455 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Comma-separated lists of bare HTTP tokens. +//! +//! The grammar accepted here is the `#rule` expansion of [RFC 9110 section +//! 5.6.1.2] restricted to members that are bare tokens: a field line is a +//! sequence of comma-delimited members, each surrounded by optional +//! whitespace, and each member is either empty or one token. Quoted strings, +//! parameters, and any other byte leave the subset, so a rejected line means +//! "not a simple token list" and never "definitely malformed": callers that +//! need a diagnostic re-scan the line with their own parser. +//! +//! [RFC 9110 section 5.6.1.2]: https://www.rfc-editor.org/rfc/rfc9110#section-5.6.1.2 + +/// How a list scan treats zero-length members. +/// +/// # Examples +/// +/// ``` +/// use http_headers_simd::{EmptyMembers, TokenListScan, scan_token_list}; +/// +/// assert_eq!( +/// scan_token_list(b"gzip,,br", EmptyMembers::Skip), +/// TokenListScan::Members, +/// ); +/// ``` +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum EmptyMembers { + /// Empty members are ignored, which is what `#rule` expansion prescribes. + /// + /// `a,,b` and `,a,` both hold the members `a` and `b`. + Skip, + /// An empty member rejects the whole field line. + /// + /// A line that holds nothing but optional whitespace is still reported as + /// [`TokenListScan::Empty`] rather than rejected, because it carries no + /// member at all and callers report that as a missing value. + Reject, +} + +/// What one field line turned out to be. +/// +/// # Examples +/// +/// ``` +/// use http_headers_simd::{EmptyMembers, TokenListScan, scan_token_list}; +/// +/// assert_eq!( +/// scan_token_list(b"gzip;q=1", EmptyMembers::Skip), +/// TokenListScan::Rejected, +/// ); +/// ``` +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum TokenListScan { + /// A valid list holding at least one token member. + Members, + /// A valid list holding no member at all. + Empty, + /// Not a simple token list under the requested empty-member policy. + Rejected, +} + +/// The byte classes a token list is built from. +/// Dense indices keep the classification table cache-small and directly indexable. +const CLASS_OTHER: usize = 0; +const CLASS_TOKEN: usize = 1; +const CLASS_OWS: usize = 2; +const CLASS_COMMA: usize = 3; + +/// The current member is still empty: the line just started or a comma just +/// closed the member before it. +const STATE_MEMBER_DUE: u8 = 0; +/// A token run is in progress. +const STATE_IN_TOKEN: u8 = 1; +/// Whitespace opened right after a token, so only a comma may follow. +const STATE_AFTER_TOKEN: u8 = 2; +/// The line left the subset. The state absorbs, so the scan needs no branch to +/// leave the loop and the outcome is still exact. +const STATE_REJECTED: u8 = 3; + +/// A non-empty member has been seen. +const FLAG_MEMBER: u8 = 4; +/// An empty member has been seen. +const FLAG_EMPTY: u8 = 8; + +/// Drives the scan one byte at a time. +/// +/// The entry for a state and a byte holds the next state in its low two bits +/// and the flags that byte raises above them, so a byte costs one indexed load, +/// one `or` into the sticky flags, and one mask to form the next row. +static LIST_TRANSITION: [u8; 4 * 256] = list_transition_table(); + +const fn list_transition_table() -> [u8; 4 * 256] { + let mut table = [STATE_REJECTED; 4 * 256]; + let mut state = 0_usize; + while state < 4 { + let mut byte = 0_usize; + while byte < 256 { + #[expect(clippy::cast_possible_truncation, reason = "the loop bound keeps the index inside a byte")] + let class = byte_class(byte as u8); + table[(state << 8) | byte] = match (state, class) { + (_, CLASS_OTHER) | (3, _) | (2, CLASS_TOKEN) => STATE_REJECTED, + (0, CLASS_TOKEN) => STATE_IN_TOKEN | FLAG_MEMBER, + (_, CLASS_TOKEN) => STATE_IN_TOKEN, + (0, CLASS_OWS) => STATE_MEMBER_DUE, + (_, CLASS_OWS) => STATE_AFTER_TOKEN, + (0, _) => STATE_MEMBER_DUE | FLAG_EMPTY, + _ => STATE_MEMBER_DUE, + }; + byte += 1; + } + state += 1; + } + table +} + +/// Maps one byte to its role inside a comma-separated token list. +const fn byte_class(byte: u8) -> usize { + if is_token_byte(byte) { + CLASS_TOKEN + } else if byte == b',' { + CLASS_COMMA + } else if byte == b' ' || byte == b'\t' { + CLASS_OWS + } else { + CLASS_OTHER + } +} + +/// Returns whether one byte is an RFC 9110 `tchar`. +const fn is_token_byte(byte: u8) -> bool { + matches!( + byte, + b'!' | b'#' + | b'$' + | b'%' + | b'&' + | b'\'' + | b'*' + | b'+' + | b'-' + | b'.' + | b'^' + | b'_' + | b'`' + | b'|' + | b'~' + | b'0'..=b'9' + | b'A'..=b'Z' + | b'a'..=b'z' + ) +} + +/// Runs the whole scan without any architecture-specific acceleration. +pub(super) fn scan_token_list(bytes: &[u8], empty: EmptyMembers) -> TokenListScan { + ListScan::new().finish(bytes, empty) +} + +/// Joins the class masks of two overlapping vectors into one lane mask. +/// +/// `low` covers the first vector of the line and `high` the vector that ends +/// it, so the two overlap by `shift` lanes and dropping that many lanes from +/// `high` leaves exactly the lanes the line still owes. +pub(super) fn join_lanes(low: u32, high: u32, shift: u32) -> u64 { + u64::from(low) | (u64::from(high >> shift) << 16) +} + +/// The state a token-list scan carries between blocks and bytes. +/// +/// Accelerated scanners fold whole blocks with [`ListScan::push_block`] and +/// hand the trailing bytes to [`ListScan::finish`], so the grammar itself +/// lives here once rather than in every architecture's kernel. +#[derive(Clone, Copy, Debug)] +pub(super) struct ListScan { + /// Where the scan stands between members, as one of the `STATE_` values. + state: u8, + /// The `FLAG_` bits the bytes fed one at a time have raised so far. + flags: u8, + /// Every token lane a block has reported, kept as raw bits so a block + /// costs one `or` rather than a comparison and a shift. + member_lanes: u64, + /// Every lane at which a block closed an empty member, kept as raw bits + /// for the same reason. + empty_lanes: u64, +} + +impl ListScan { + /// Starts a scan at the beginning of a field line. + pub(super) const fn new() -> Self { + Self { + state: STATE_MEMBER_DUE, + flags: 0, + member_lanes: 0, + empty_lanes: 0, + } + } + + /// Folds one 16-byte block described by its per-lane class masks. + /// + /// `tokens`, `ows`, and `commas` hold one bit per lane, lane zero in bit + /// zero, and every other bit must be clear. Returns `false` as soon as the + /// block cannot belong to a token list, which happens either because a lane + /// is outside the three classes or because whitespace alone separates two + /// members. + pub(super) fn push_block(&mut self, tokens: u32, ows: u32, commas: u32) -> bool { + self.push_lanes(u64::from(tokens), u64::from(ows), u64::from(commas), 16) + } + + /// Folds the `lanes` low lanes of a block that covers only part of a line. + /// + /// A scanner reaching a trailing run shorter than one vector reloads the + /// last whole vector of the line and shifts its masks down, so the lanes + /// it already folded fall away and lane zero is the first byte still + /// owed. `lanes` must be between one and sixteen and every mask bit at or + /// above it must be clear. + pub(super) fn push_partial_block(&mut self, tokens: u32, ows: u32, commas: u32, lanes: u32) -> bool { + self.push_lanes(u64::from(tokens), u64::from(ows), u64::from(commas), lanes) + } + + /// Folds a whole line of at most two vectors as one run of `lanes` lanes. + /// + /// The grammar's cost is dominated by the fold rather than by the class + /// masks, so a line that spans two vectors is cheaper to classify twice + /// and fold once than to fold twice. `lanes` must be between one and + /// thirty-two and every mask bit at or above it must be clear. + pub(super) fn push_line(&mut self, tokens: u64, ows: u64, commas: u64, lanes: u32) -> bool { + self.push_lanes(tokens, ows, commas, lanes) + } + + /// Folds `lanes` lanes, which is the whole grammar for a block of any width. + #[expect(clippy::inline_always, reason = "the whole-block caller folds the lane count into its constants")] + #[inline(always)] + fn push_lanes(&mut self, tokens: u64, ows: u64, commas: u64, lanes: u32) -> bool { + // The two carry-sensitive facts about the byte before the block are + // whether the current member is still empty and whether whitespace may + // no longer be followed by a token, and both are one bit, so the whole + // block folds without a branch until the single rejection test. + let due = u64::from(self.state == STATE_MEMBER_DUE); + // A block is only ever folded into a state the previous block or byte + // left behind, and both leave one of the first three states, so the + // high bit of the state alone says whether whitespace closed a token. + let after_token = u64::from(self.state >> 1); + let leading_ows = ows & 1; + + // Whitespace runs opening right after a token get a start bit, and + // adding those bits to the run carries through it, so the carry lands + // on the byte that ends the run. A token there is a member that no + // comma separated from the one before it. Landing bits are the only + // bits the sum can place outside the whitespace mask, so intersecting + // with the token mask needs no further masking. + let after_tokens = tokens << 1; + let token_runs = (after_tokens & ows) | (leading_ows & (due ^ 1)); + let carried = ows + token_runs; + + // The same carry finds empty members: a run that opens right after a + // comma and lands on another comma spans a member holding nothing, and + // adjacent commas are the same thing with no whitespace between. + let after_commas = commas << 1; + let empty_runs = (after_commas & ows) | (leading_ows & due); + let closed_empty = ((ows + empty_runs) & commas) | (after_commas & commas) | (commas & due); + + // A lane outside the three classes leaves a hole in the union, a token + // in lane zero cannot follow whitespace that closed a token, and a + // landing bit on a token is a member with no comma before it. + let all = (1_u64 << lanes) - 1; + let rejected = ((tokens | ows | commas) ^ all) | (after_token & tokens) | (carried & tokens); + if rejected != 0 { + self.state = STATE_REJECTED; + return false; + } + + // A block ends in a token, in a whitespace run that a token opened, or + // with the next member due; the first two are mutually exclusive + // because one lane cannot be both. + self.state = u8::from(tokens & (1_u64 << (lanes - 1)) != 0) | (u8::from(carried & (1_u64 << lanes) != 0) << 1); + self.member_lanes |= tokens; + self.empty_lanes |= closed_empty; + true + } + + /// Folds the trailing bytes and reports the outcome for the whole line. + pub(super) fn finish(self, tail: &[u8], empty: EmptyMembers) -> TokenListScan { + let mut row = usize::from(self.state) << 8; + let mut flags = self.flags; + for &byte in tail { + let entry = LIST_TRANSITION[row | usize::from(byte)]; + flags |= entry; + row = usize::from(entry & 3) << 8; + } + #[expect(clippy::cast_possible_truncation, reason = "the row is one of the four states shifted into place")] + let state = (row >> 8) as u8; + if state == STATE_REJECTED { + return TokenListScan::Rejected; + } + + let members = flags & FLAG_MEMBER != 0 || self.member_lanes != 0; + // A line that ends with a member still due ends on an empty member, + // unless nothing at all preceded it and the line simply has no members. + let empties = flags & FLAG_EMPTY != 0 || self.empty_lanes != 0 || state == STATE_MEMBER_DUE && members; + match empty { + EmptyMembers::Skip => { + if members { + TokenListScan::Members + } else { + TokenListScan::Empty + } + } + EmptyMembers::Reject => { + if !members && !empties { + TokenListScan::Empty + } else if empties { + TokenListScan::Rejected + } else { + TokenListScan::Members + } + } + } + } +} + +/// Splits on commas and checks every trimmed member, which is the grammar +/// spelled the obvious way rather than the way the scanners walk it. +/// +/// Every differential test in this crate compares an implementation against +/// this, so it is deliberately the slow and literal reading of the rule. +#[cfg(test)] +pub(super) fn oracle(bytes: &[u8], empty: EmptyMembers) -> TokenListScan { + let mut members = 0_usize; + let mut empties = 0_usize; + for member in bytes.split(|byte| *byte == b',') { + let member = trim_ows(member); + if member.is_empty() { + empties += 1; + } else if member.iter().copied().all(is_token_byte) { + members += 1; + } else { + return TokenListScan::Rejected; + } + } + match empty { + EmptyMembers::Skip => { + if members == 0 { + TokenListScan::Empty + } else { + TokenListScan::Members + } + } + EmptyMembers::Reject => { + if members == 0 && empties == 1 { + TokenListScan::Empty + } else if empties == 0 { + TokenListScan::Members + } else { + TokenListScan::Rejected + } + } + } +} + +#[cfg(test)] +fn trim_ows(bytes: &[u8]) -> &[u8] { + let start = bytes.iter().position(|byte| *byte != b' ' && *byte != b'\t').unwrap_or(bytes.len()); + let end = bytes + .iter() + .rposition(|byte| *byte != b' ' && *byte != b'\t') + .map_or(start, |index| index + 1); + &bytes[start..end] +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::hint; + #[cfg(not(feature = "std"))] + use std::vec::Vec; + + use super::*; + + #[test] + fn token_class_matches_the_shared_scalar_validator() { + for byte in u8::MIN..=u8::MAX { + assert_eq!( + byte_class(byte) == CLASS_TOKEN, + crate::scalar::all_token_bytes(&[byte]), + "byte {byte:#04x}" + ); + } + } + + #[test] + fn runtime_transition_table_matches_the_static_table() { + let generated = hint::black_box(list_transition_table()); + assert_eq!(generated, LIST_TRANSITION); + + for state in 0..4 { + for byte in u8::MIN..=u8::MAX { + let entry = generated[(state << 8) | usize::from(byte)]; + assert_eq!(entry & !(3 | FLAG_MEMBER | FLAG_EMPTY), 0); + if state == usize::from(STATE_REJECTED) { + assert_eq!(entry, STATE_REJECTED); + } + } + } + } + + #[test] + fn overlapping_lane_masks_join_without_repeating_lanes() { + assert_eq!(join_lanes(0x0000_ffff, 0xffff_0000, 16), 0x0000_0000_ffff_ffff); + assert_eq!(join_lanes(0x0000_00ff, 0x0000_ff00, 8), 0x0000_0000_00ff_00ff); + } + + #[test] + fn scan_matches_the_oracle_for_every_short_alphabet_string() { + let alphabet = b"a,\t %\0"; + let mut buffer = Vec::new(); + for length in 0..=4 { + walk(alphabet, length, &mut buffer); + } + } + + fn walk(alphabet: &[u8], length: usize, buffer: &mut Vec) { + if length == 0 { + for empty in [EmptyMembers::Skip, EmptyMembers::Reject] { + assert_eq!(scan_token_list(buffer, empty), oracle(buffer, empty), "{buffer:?} {empty:?}"); + } + return; + } + for &byte in alphabet { + buffer.push(byte); + walk(alphabet, length - 1, buffer); + buffer.pop(); + } + } + + #[test] + fn whitespace_separated_members_need_a_comma() { + assert_eq!(scan_token_list(b"gzip deflate", EmptyMembers::Skip), TokenListScan::Rejected); + assert_eq!(scan_token_list(b"gzip , deflate", EmptyMembers::Skip), TokenListScan::Members); + } + + #[test] + fn empty_members_follow_the_policy() { + assert_eq!(scan_token_list(b"gzip,,deflate", EmptyMembers::Skip), TokenListScan::Members); + assert_eq!(scan_token_list(b"gzip,,deflate", EmptyMembers::Reject), TokenListScan::Rejected); + assert_eq!(scan_token_list(b" \t ", EmptyMembers::Reject), TokenListScan::Empty); + assert_eq!(scan_token_list(b"", EmptyMembers::Skip), TokenListScan::Empty); + } + + #[test] + fn quoted_and_control_bytes_leave_the_subset() { + assert_eq!(scan_token_list(b"\"gzip\"", EmptyMembers::Skip), TokenListScan::Rejected); + assert_eq!(scan_token_list(b"gzip;q=1", EmptyMembers::Skip), TokenListScan::Rejected); + } +} diff --git a/crates/http_headers_simd/src/range.rs b/crates/http_headers_simd/src/range.rs new file mode 100644 index 000000000..f8a10b9c3 --- /dev/null +++ b/crates/http_headers_simd/src/range.rs @@ -0,0 +1,463 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Guaranteed-valid scanner for the common `bytes=` byte-range-set. +//! +//! A field line such as `bytes=0-499, 1000-` is a comma separated list of +//! `first-last`, `first-`, and `-suffix` specifications made only of decimal +//! digits, a hyphen, commas, and optional whitespace. The scanner classifies +//! one sixteen byte window with a handful of masks and then decides the whole +//! grammar — including `first <= last` — with branch-free bit arithmetic, so a +//! well formed line never walks the general parser. +//! +//! Every answer is conservative: `false` means "not provably a valid range +//! set", never "malformed", so callers fall back to the parser that produces +//! the diagnostic. + +/// Bytes classified in one pass. +pub(super) const WINDOW: usize = 16; + +/// Byte that fills the lanes before a line shorter than one window. +/// +/// It belongs to no class the scan reads, so it can neither open a member nor +/// satisfy a rule. +pub(super) const FILLER: u8 = b'x'; + +/// Per-lane classification of one window. +/// +/// Each field holds one bit per lane, with lane zero in the least significant +/// bit. The vector kernels only fill these in; every grammar decision is made +/// once, in [`accepts`], so the scalar, SSE2, and NEON paths cannot disagree +/// about anything but the masks themselves. +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(super) struct RangeMasks { + /// Lanes holding `0`–`9`. + pub(super) digits: u32, + /// Lanes holding `-`. + pub(super) dashes: u32, + /// Lanes holding `,`. + pub(super) commas: u32, + /// Lanes holding a space or a horizontal tab. + pub(super) spaces: u32, + /// Lanes holding `0`. + pub(super) zeros: u32, +} + +impl RangeMasks { + /// Classifies a window without vector instructions. + #[inline] + pub(super) fn scalar(window: &[u8; WINDOW]) -> Self { + let mut masks = Self { + digits: 0, + dashes: 0, + commas: 0, + spaces: 0, + zeros: 0, + }; + for (lane, &byte) in window.iter().enumerate() { + let bit = 1_u32 << lane; + match byte { + b'0'..=b'9' => { + masks.digits |= bit; + if byte == b'0' { + masks.zeros |= bit; + } + } + b'-' => masks.dashes |= bit, + b',' => masks.commas |= bit, + b' ' | b'\t' => masks.spaces |= bit, + _other => {} + } + } + masks + } +} + +/// Returns whether the last `len` lanes of `window` are a valid range set. +/// +/// The window is classified by `masks`, and only its final `len` lanes belong +/// to the payload, so a caller that holds a longer field line can pass the +/// window that ends with it and leave the earlier lanes alone. +/// +/// Every rule contributes to one rejection mask that is tested once, so a line +/// that is going to be accepted takes a single branch on its way through. +#[expect(clippy::inline_always, reason = "the padded path pays for the call and the masks it cannot fold")] +#[inline(always)] +pub(super) fn accepts(window: &[u8; WINDOW], masks: &RangeMasks, len: usize) -> bool { + if len == 0 || len > WINDOW { + return false; + } + let payload = (0xffff_u64 << (WINDOW - len)) & 0xffff; + let digits = u64::from(masks.digits) & payload; + let dashes = u64::from(masks.dashes) & payload; + let commas = u64::from(masks.commas) & payload; + let spaces = u64::from(masks.spaces) & payload; + let after_digit = digits << 1; + let before_digit = digits >> 1; + + // Members are the runs of digits and hyphens, and every one of them walks + // over its leading digits onto the hyphen that splits it. A run that walks + // onto a comma, onto whitespace, or off the end of the line holds no + // hyphen, and a second hyphen is neither the landing of its member nor the + // start of one. + let member = digits | dashes; + let starts = member & !(member << 1); + let landing = digits.wrapping_add(starts & digits) & !digits; + + let rejected = + // A byte outside the four classes leaves the subset. + (payload & !(member | commas | spaces)) + // Whitespace is only recognized where it follows a comma. + | (spaces & !(commas << 1)) + // Every member holds exactly one hyphen. + | (landing & !dashes) + | (dashes & !landing & !starts) + // A hyphen needs a digit on at least one side. + | (dashes & !after_digit & !before_digit) + // A run that starts with a zero and continues would compare as a + // shorter number than it is, so leave those to the general parser. + | (u64::from(masks.zeros) & payload & before_digit & !after_digit) + // A line of nothing but separators holds no member at all. + | (starts.wrapping_sub(1) >> (u64::BITS - 1)); + if rejected != 0 { + return false; + } + + // Only a specification with digits on both sides of its hyphen can be + // inverted; the rest are already ordered. + let mut compare = dashes & after_digit & before_digit; + while compare != 0 { + if !ordered(window, !digits, compare.trailing_zeros()) { + return false; + } + compare &= compare - 1; + } + true +} + +/// Returns whether the specification whose hyphen sits at `dash` is ordered. +/// +/// `gaps` marks every lane that is not a digit, so shifting it until the lane +/// beside the hyphen reaches the end of the word turns each run length into a +/// count of leading or trailing zeros. Neither run carries a leading zero by +/// the time this runs, so the shorter run is the smaller number and equal runs +/// compare byte by byte. +#[inline] +fn ordered(window: &[u8; WINDOW], gaps: u64, dash: u32) -> bool { + let first_len = (gaps << (u64::BITS - dash)).leading_zeros(); + let last_len = (gaps >> (dash + 1)).trailing_zeros(); + if first_len != last_len { + return first_len < last_len; + } + same_length_ordered(window, dash, first_len) +} + +/// Compares two runs of the same length byte by byte. +/// +/// Equal lengths are the only case that survives the length comparison, and a +/// well formed set usually spells its bounds with different widths, so this +/// stays out of the straight-line path. +#[cold] +#[inline(never)] +fn same_length_ordered(window: &[u8; WINDOW], dash: u32, len: u32) -> bool { + let first = (dash - len) as usize; + let last = dash as usize + 1; + for step in 0..len as usize { + let left = window[first + step]; + let right = window[last + step]; + if left != right { + return left < right; + } + } + true +} + +/// Scans a field line for a range set without vector instructions. +/// +/// The dispatcher reaches the scalar classifier directly, so this entry only +/// serves the differential tests and the benchmark harness. +#[cfg(any(test, feature = "benchmarking"))] +pub(super) fn scan_byte_range_set(bytes: &[u8], start: usize) -> bool { + with_window(bytes, start, |window, len| accepts(window, &RangeMasks::scalar(window), len)) +} + +/// Calls `scan` with the window that ends the field line, if one exists. +/// +/// A line at least [`WINDOW`] bytes long lends its own tail, and a shorter one +/// is copied into a padded window whose leading lanes fall outside the +/// payload; either way the payload occupies the final `len` lanes. +#[cfg(any(test, feature = "benchmarking"))] +#[inline] +pub(super) fn with_window(bytes: &[u8], start: usize, scan: fn(&[u8; WINDOW], usize) -> bool) -> bool { + let Some(len) = bytes.len().checked_sub(start) else { + return false; + }; + if len == 0 || len > WINDOW { + return false; + } + if let Some(window) = window_of(bytes) { + return scan(window, len); + } + let mut padded = [FILLER; WINDOW]; + if !fill_window(bytes, &mut padded) { + return false; + } + scan(&padded, len) +} + +/// Right aligns a line shorter than one window inside a padded window. +/// +/// Two eight byte copies place the line without the branch tree a variable +/// length copy needs; they overlap for lines shorter than sixteen bytes, which +/// writes the same bytes twice and costs nothing. Lines of fewer than eight +/// bytes have no such pair and are left to the caller's own parser, and the +/// filler is a byte no rule reads, so the lanes before the line decide +/// nothing. +#[inline] +pub(super) fn fill_window(bytes: &[u8], window: &mut [u8; WINDOW]) -> bool { + if !(8..=WINDOW).contains(&bytes.len()) { + return false; + } + let offset = WINDOW - bytes.len(); + let tail = bytes.len() - 8; + window[offset..offset + 8].copy_from_slice(&bytes[..8]); + window[WINDOW - 8..].copy_from_slice(&bytes[tail..]); + true +} + +/// Borrows the window a line ends with, if the line is long enough. +/// +/// Asking the slice for its final chunk keeps the bound check to the one +/// comparison the length needs, where re-slicing and converting would also +/// make the compiler prove the pointer arithmetic afresh. +pub(super) fn window_of(bytes: &[u8]) -> Option<&[u8; WINDOW]> { + bytes.last_chunk::() +} + +/// Mirrors the general parser closely enough to check the scanner against it. +#[cfg(test)] +pub(super) fn oracle(payload: &[u8]) -> bool { + fn number(bytes: &[u8]) -> Option { + if bytes.is_empty() { + return None; + } + bytes.iter().try_fold(0_u64, |value, byte| { + value + .checked_mul(10)? + .checked_add(u64::from(byte.checked_sub(b'0').filter(|d| *d <= 9)?)) + }) + } + + fn trim(bytes: &[u8]) -> &[u8] { + let mut slice = bytes; + while let [b' ' | b'\t', rest @ ..] = slice { + slice = rest; + } + while let [rest @ .., b' ' | b'\t'] = slice { + slice = rest; + } + slice + } + + let mut count = 0_usize; + for item in payload.split(|byte| *byte == b',') { + let item = trim(item); + if item.is_empty() { + continue; + } + let Some(dash) = item.iter().position(|byte| *byte == b'-') else { + return false; + }; + if dash == 0 { + if number(&item[1..]).is_none() { + return false; + } + } else { + if item[dash + 1..].contains(&b'-') { + return false; + } + let Some(first) = number(&item[..dash]) else { + return false; + }; + let last = &item[dash + 1..]; + if !last.is_empty() { + let Some(last) = number(last) else { + return false; + }; + if last < first { + return false; + } + } + } + count += 1; + } + count != 0 +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::ops::RangeInclusive; + #[cfg(not(feature = "std"))] + use std::vec::Vec; + + use super::*; + + /// Runs the scanner the way a field line reaches it. + fn scan(payload: &[u8]) -> bool { + let mut line = Vec::from(&b"bytes="[..]); + line.extend_from_slice(payload); + scan_byte_range_set(&line, 6) + } + + #[test] + fn accepts_the_specification_examples() { + for payload in [ + &b"0-499"[..], + b"0-499, 1000-", + b"1000-", + b"-500", + b"0-0", + b"9-9", + b"0-499,500-999", + b"0-1,\t2-3", + b"0-1,,2-3", + b"0-1,", + b",0-1", + b"12-345", + b"499-499", + ] { + assert!(scan(payload), "rejected {payload:?}"); + assert!(oracle(payload), "oracle rejected {payload:?}"); + } + } + + #[test] + fn leaves_everything_unusual_to_the_general_parser() { + for payload in [ + &b""[..], + b"-", + b"0", + b"500-100", + b"0--1", + b"0-1-2", + b"0- 1", + b"0 - 1", + b"00-1", + b"1-09", + b"bytes=0-1", + b"0-1;q=1", + ] { + assert!(!scan(payload), "accepted {payload:?}"); + } + } + + #[test] + fn helpers_reject_unrepresentable_windows_and_offsets() { + fn expected_window(window: &[u8; WINDOW], len: usize) -> bool { + len == WINDOW && window == b"0123456789abcdef" + } + + let mut padded = [FILLER; WINDOW]; + assert!(!accepts(&padded, &RangeMasks::scalar(&padded), 0)); + assert!(!accepts(&padded, &RangeMasks::scalar(&padded), WINDOW + 1)); + assert!(!with_window(b"short", 6, expected_window)); + assert!(!with_window(b"payload", 7, expected_window)); + assert!(!with_window(b"0123456789abcdefg", 0, expected_window)); + assert!(!fill_window(b"1234567", &mut padded)); + assert!(!fill_window(b"0123456789abcdefg", &mut padded)); + + assert!(fill_window(b"1234-678", &mut padded)); + assert_eq!(&padded[WINDOW - 8..], b"1234-678"); + assert!(with_window(b"prefix0123456789abcdef", 6, expected_window)); + } + + #[test] + fn oracle_rejects_each_malformed_number_shape() { + for payload in [ + &b", ,\t,"[..], + b"1", + b"-", + b"-x", + b"x-1", + b"1-2-3", + b"1-x", + b"2-1", + b"18446744073709551616-", + b"1-18446744073709551616", + b"184467440737095516150-", + b"1-184467440737095516150", + ] { + assert!(!oracle(payload), "{payload:?}"); + } + assert!(oracle(b" \t1-2\t ")); + } + + #[test] + fn comma_before_payload_does_not_legalize_leading_space() { + let mut window = [FILLER; WINDOW]; + window[WINDOW - 5] = b','; + window[WINDOW - 4..].copy_from_slice(b" 0-1"); + assert!(!accepts(&window, &RangeMasks::scalar(&window), 4)); + } + + #[test] + fn every_acceptance_is_one_the_general_parser_shares() { + const ALPHABET: &[u8] = b"0-,1 9\t2"; + const NARROW: &[u8] = b"0-,1 "; + + fn sweep(alphabet: &[u8], lengths: RangeInclusive) { + let mut payload = Vec::new(); + for length in lengths { + let total = alphabet.len().pow(u32::try_from(length).unwrap_or(0)); + for index in 0..total { + payload.clear(); + let mut rest = index; + for _step in 0..length { + payload.push(alphabet[rest % alphabet.len()]); + rest /= alphabet.len(); + } + if scan(&payload) { + assert!(oracle(&payload), "accepted invalid {payload:?}"); + } + } + } + } + + // The wide alphabet covers every class, and the narrow one reaches the + // lengths that hold two specifications and a separator between them. + let wide_max = if cfg!(miri) { 3 } else { 5 }; + let narrow_max = if cfg!(miri) { 7 } else { 8 }; + sweep(ALPHABET, 0..=wide_max); + sweep(NARROW, 6..=narrow_max); + } + + #[test] + fn ordering_is_decided_for_every_short_pair() { + fn digits(value: u32, into: &mut Vec) { + if value >= 10 { + digits(value / 10, into); + } + into.push(b'0' + u8::try_from(value % 10).unwrap_or(0)); + } + + let mut payload = Vec::new(); + for first in 0..=120_u32 { + for last in 0..=120_u32 { + payload.clear(); + digits(first, &mut payload); + payload.push(b'-'); + digits(last, &mut payload); + assert_eq!(scan(&payload), first <= last, "{payload:?}"); + assert_eq!(oracle(&payload), first <= last, "{payload:?}"); + } + } + } + + #[test] + fn the_window_must_cover_the_whole_payload() { + // Seventeen payload bytes cannot be classified in one window. + assert!(!scan(b"10000000-20000000")); + assert!(scan(b"1000000-2000000")); + } +} diff --git a/crates/http_headers_simd/src/scalar.rs b/crates/http_headers_simd/src/scalar.rs new file mode 100644 index 000000000..631424886 --- /dev/null +++ b/crates/http_headers_simd/src/scalar.rs @@ -0,0 +1,87 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Portable scalar validators and reference implementations. + +// RFC 9110 `tchar`, indexed by the upper two bits of each byte. +const TOKEN_BITMAP: [u64; 4] = [0x03ff_6cfa_0000_0000, 0x57ff_ffff_c7ff_fffe, 0, 0]; +// RFC 9110 `token68` data characters, indexed by the upper two bits. +const TOKEN68_BITMAP: [u64; 4] = [0x03ff_e800_0000_0000, 0x47ff_fffe_87ff_fffe, 0, 0]; + +pub(super) fn is_token(bytes: &[u8]) -> bool { + !bytes.is_empty() && all_token_bytes(bytes) +} + +pub(super) fn all_token_bytes(bytes: &[u8]) -> bool { + bytes.iter().copied().all(is_token_byte) +} + +pub(super) fn is_token68(bytes: &[u8]) -> bool { + token68_tail(bytes, false) +} + +pub(super) fn token68_tail(bytes: &[u8], mut has_data: bool) -> bool { + for (index, &byte) in bytes.iter().enumerate() { + if !is_token68_data_byte(byte) { + return has_data && byte == b'=' && bytes[index + 1..].iter().all(|byte| *byte == b'='); + } + has_data = true; + } + has_data +} + +fn is_token68_data_byte(byte: u8) -> bool { + let word = TOKEN68_BITMAP[usize::from(byte >> 6)]; + word & (1_u64 << (byte & 63)) != 0 +} + +fn is_token_byte(byte: u8) -> bool { + let word = TOKEN_BITMAP[usize::from(byte >> 6)]; + word & (1_u64 << (byte & 63)) != 0 +} + +pub(super) fn is_field_value(bytes: &[u8]) -> bool { + bytes.iter().copied().all(|byte| byte == b'\t' || byte >= b' ' && byte != 0x7f) +} + +pub(super) fn eq_ignore_ascii_case(left: &[u8], right: &[u8]) -> bool { + left.eq_ignore_ascii_case(right) +} + +pub(super) fn find_interesting(bytes: &[u8]) -> Option { + bytes + .iter() + .position(|byte| matches!(byte, b',' | b';' | b'"' | b'\\' | b' ' | b'\t')) +} + +pub(super) fn find_either(bytes: &[u8], first: u8, second: u8) -> Option { + bytes.iter().position(|byte| *byte == first || *byte == second) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::*; + + #[test] + fn token68_requires_data_before_contiguous_padding() { + for valid in [&b"a"[..], b"a=", b"a==", b"AZ09-._~+/=="] { + assert!(is_token68(valid), "{valid:?}"); + } + for invalid in [&b""[..], b"=", b"==", b"a=a", b"a= ="] { + assert!(!is_token68(invalid), "{invalid:?}"); + } + } + + #[test] + fn scalar_helpers_handle_offsets_and_non_ascii_bytes() { + assert!(all_token_bytes(b"AZaz09!#$%&'*+-.^_`|~")); + assert!(!all_token_bytes(b"token()")); + assert!(is_field_value(&[b'\t', b' ', b'~', 0x80, 0xff])); + assert!(!is_field_value(&[0x7f])); + assert!(eq_ignore_ascii_case(b"HeAdEr", b"hEaDeR")); + assert!(!eq_ignore_ascii_case(&[0x80], &[0x81])); + assert_eq!(find_interesting(b"token;parameter"), Some(5)); + assert_eq!(find_interesting(b"token"), None); + } +} diff --git a/crates/http_headers_simd/src/tracking.rs b/crates/http_headers_simd/src/tracking.rs new file mode 100644 index 000000000..764bc82ab --- /dev/null +++ b/crates/http_headers_simd/src/tracking.rs @@ -0,0 +1,455 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Process-wide allocation accounting for the workspace performance scripts. +//! +//! This module exists so allocation claims in `docs/PERF.md` can be measured +//! without any project-authored `unsafe` outside this crate. It is compiled +//! only under the `benchmarking` feature and is never part of a normal build. + +use std::alloc::{GlobalAlloc, Layout, System}; +use std::hint; +use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}; + +#[repr(align(64))] +struct Counters { + locked: AtomicBool, + allocations: AtomicU64, + allocated_bytes: AtomicU64, + live_bytes: AtomicUsize, + peak_live_bytes: AtomicUsize, +} + +static COUNTERS: Counters = Counters { + locked: AtomicBool::new(false), + allocations: AtomicU64::new(0), + allocated_bytes: AtomicU64::new(0), + live_bytes: AtomicUsize::new(0), + peak_live_bytes: AtomicUsize::new(0), +}; + +struct CounterGuard; + +impl Drop for CounterGuard { + fn drop(&mut self) { + COUNTERS.locked.store(false, Ordering::Release); + } +} + +fn lock_counters() -> CounterGuard { + lock_counters_with_attempt_hook(|_| {}) +} + +fn lock_counters_with_attempt_hook(mut after_attempt: impl FnMut(bool)) -> CounterGuard { + loop { + let acquired = COUNTERS + .locked + .compare_exchange_weak(false, true, Ordering::Acquire, Ordering::Relaxed) + .is_ok(); + after_attempt(acquired); + if acquired { + return CounterGuard; + } + hint::spin_loop(); + } +} + +/// A snapshot of the counters maintained by [`TrackingAllocator`]. +/// +/// Fields are public so benchmark binaries can consume snapshots as passive +/// data without accessor overhead. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(all(feature = "benchmarking", feature = "std"))] +/// # { +/// let stats = http_headers_simd::tracking::AllocationStats::default(); +/// assert_eq!(stats.allocations, 0); +/// assert_eq!(stats.live_bytes, 0); +/// # } +/// ``` +#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)] +pub struct AllocationStats { + /// The number of successful allocations, including growing reallocations. + pub allocations: u64, + /// The total number of bytes handed out, including reallocation growth. + pub allocated_bytes: u64, + /// The number of bytes currently allocated and not yet released. + pub live_bytes: usize, + /// The largest observed value of [`AllocationStats::live_bytes`]. + pub peak_live_bytes: usize, +} + +/// A [`GlobalAlloc`] that forwards to the system allocator and counts requests. +/// +/// Install it in a measurement binary with the `#[global_allocator]` attribute. +/// Counting is process-wide. Counter updates and snapshots are serialized, so +/// concurrent allocations are accounted coherently. +/// +/// # Examples +/// +/// ```no_run +/// # #[cfg(all(feature = "benchmarking", feature = "std"))] +/// # { +/// use http_headers_simd::tracking::TrackingAllocator; +/// +/// #[global_allocator] +/// static ALLOCATOR: TrackingAllocator = TrackingAllocator::new(); +/// # } +/// ``` +#[derive(Clone, Copy, Debug, Default)] +pub struct TrackingAllocator; + +impl TrackingAllocator { + /// Creates the allocator. + /// + /// # Examples + /// + /// ``` + /// # #[cfg(all(feature = "benchmarking", feature = "std"))] + /// # { + /// let allocator = http_headers_simd::tracking::TrackingAllocator::new(); + /// let _ = allocator; + /// # } + /// ``` + #[must_use] + pub const fn new() -> Self { + Self + } +} + +// SAFETY: every method forwards its arguments unchanged to `System`, which is +// a correct `GlobalAlloc`, and returns exactly what `System` returned. The +// counter updates neither allocate nor touch the returned memory, so all of +// `GlobalAlloc`'s pointer, layout, and aliasing obligations are discharged by +// `System` itself. +unsafe impl GlobalAlloc for TrackingAllocator { + unsafe fn alloc(&self, layout: Layout) -> *mut u8 { + // SAFETY: `layout` is forwarded unchanged from the caller, which the + // `GlobalAlloc` contract already requires to be valid for `System`. + let pointer = unsafe { System.alloc(layout) }; + record_allocation(pointer, layout.size()) + } + + unsafe fn alloc_zeroed(&self, layout: Layout) -> *mut u8 { + // SAFETY: `layout` is forwarded unchanged from the caller, which the + // `GlobalAlloc` contract already requires to be valid for `System`. + let pointer = unsafe { System.alloc_zeroed(layout) }; + record_allocation(pointer, layout.size()) + } + + unsafe fn dealloc(&self, ptr: *mut u8, layout: Layout) { + { + let _guard = lock_counters(); + let _ = COUNTERS.live_bytes.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |live| { + Some(live.saturating_sub(layout.size())) + }); + } + // SAFETY: `ptr` and `layout` are forwarded unchanged from the + // caller, which the `GlobalAlloc` contract already requires to name a + // block this allocator returned for that layout. + unsafe { System.dealloc(ptr, layout) }; + } + + unsafe fn realloc(&self, ptr: *mut u8, layout: Layout, new_size: usize) -> *mut u8 { + // SAFETY: all three arguments are forwarded unchanged from the caller, + // which the `GlobalAlloc` contract already requires to be valid. + let replacement = unsafe { System.realloc(ptr, layout, new_size) }; + record_reallocation(replacement, layout.size(), new_size) + } +} + +#[inline] +fn record_allocation(pointer: *mut u8, size: usize) -> *mut u8 { + if !pointer.is_null() { + record(size); + } + pointer +} + +#[inline] +fn record_reallocation(replacement: *mut u8, old_size: usize, new_size: usize) -> *mut u8 { + if replacement.is_null() { + return replacement; + } + if let Some(growth) = new_size.checked_sub(old_size) { + record(growth); + } else { + let shrink = old_size.saturating_sub(new_size); + let _guard = lock_counters(); + let _ = COUNTERS + .live_bytes + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |live| Some(live.saturating_sub(shrink))); + } + replacement +} + +fn record(bytes: usize) { + let _guard = lock_counters(); + let _ = COUNTERS.allocations.fetch_add(1, Ordering::Relaxed); + let _ = COUNTERS.allocated_bytes.fetch_add(bytes as u64, Ordering::Relaxed); + let live = COUNTERS.live_bytes.fetch_add(bytes, Ordering::Relaxed) + bytes; + let _ = COUNTERS.peak_live_bytes.fetch_max(live, Ordering::Relaxed); +} + +/// Reads the current counters. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(all(feature = "benchmarking", feature = "std"))] +/// # { +/// let stats = http_headers_simd::tracking::allocation_stats(); +/// assert!(stats.peak_live_bytes >= stats.live_bytes); +/// # } +/// ``` +#[must_use] +pub fn allocation_stats() -> AllocationStats { + let _guard = lock_counters(); + AllocationStats { + allocations: COUNTERS.allocations.load(Ordering::Relaxed), + allocated_bytes: COUNTERS.allocated_bytes.load(Ordering::Relaxed), + live_bytes: COUNTERS.live_bytes.load(Ordering::Relaxed), + peak_live_bytes: COUNTERS.peak_live_bytes.load(Ordering::Relaxed), + } +} + +/// Resets the peak watermark to the currently live byte count. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(all(feature = "benchmarking", feature = "std"))] +/// # { +/// use http_headers_simd::tracking::{allocation_stats, reset_peak_live_bytes}; +/// +/// reset_peak_live_bytes(); +/// let stats = allocation_stats(); +/// assert_eq!(stats.peak_live_bytes, stats.live_bytes); +/// # } +/// ``` +pub fn reset_peak_live_bytes() { + let _guard = lock_counters(); + COUNTERS + .peak_live_bytes + .store(COUNTERS.live_bytes.load(Ordering::Relaxed), Ordering::Relaxed); +} + +/// Returns the counter deltas accumulated between two snapshots. +/// +/// # Examples +/// +/// ``` +/// # #[cfg(all(feature = "benchmarking", feature = "std"))] +/// # { +/// use http_headers_simd::tracking::{AllocationStats, allocation_delta}; +/// +/// let before = AllocationStats { +/// allocations: 2, +/// allocated_bytes: 32, +/// live_bytes: 8, +/// peak_live_bytes: 16, +/// }; +/// let after = AllocationStats { +/// allocations: 5, +/// allocated_bytes: 80, +/// live_bytes: 24, +/// peak_live_bytes: 40, +/// }; +/// let delta = allocation_delta(before, after); +/// assert_eq!(delta.allocations, 3); +/// assert_eq!(delta.allocated_bytes, 48); +/// assert_eq!(delta.live_bytes, 16); +/// assert_eq!(delta.peak_live_bytes, 32); +/// # } +/// ``` +#[must_use] +pub fn allocation_delta(before: AllocationStats, after: AllocationStats) -> AllocationStats { + AllocationStats { + allocations: after.allocations.saturating_sub(before.allocations), + allocated_bytes: after.allocated_bytes.saturating_sub(before.allocated_bytes), + live_bytes: after.live_bytes.saturating_sub(before.live_bytes), + peak_live_bytes: after.peak_live_bytes.saturating_sub(before.live_bytes), + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::alloc::{GlobalAlloc, Layout}; + use std::sync::mpsc::{self, TryRecvError}; + use std::sync::{Mutex, PoisonError}; + use std::{ptr, slice, thread}; + + use super::{ + AllocationStats, TrackingAllocator, allocation_delta, allocation_stats, lock_counters, lock_counters_with_attempt_hook, record, + record_allocation, record_reallocation, reset_peak_live_bytes, + }; + + static TEST_LOCK: Mutex<()> = Mutex::new(()); + + #[test] + fn concurrent_updates_produce_a_coherent_snapshot() { + const THREADS: usize = 4; + const UPDATES: usize = if cfg!(miri) { 32 } else { 1_000 }; + let _serial = TEST_LOCK.lock().unwrap_or_else(PoisonError::into_inner); + + let before = allocation_stats(); + thread::scope(|scope| { + for _ in 0..THREADS { + scope.spawn(|| { + for _ in 0..UPDATES { + record(1); + } + }); + } + }); + let delta = allocation_delta(before, allocation_stats()); + + assert_eq!(delta.allocations, (THREADS * UPDATES) as u64); + assert_eq!(delta.allocated_bytes, (THREADS * UPDATES) as u64); + assert_eq!(delta.live_bytes, THREADS * UPDATES); + } + + #[test] + fn global_alloc_methods_preserve_memory_and_counters() { + let _serial = TEST_LOCK.lock().unwrap_or_else(PoisonError::into_inner); + let allocator = TrackingAllocator::new(); + let layout = Layout::from_size_align(32, 8).expect("valid test layout"); + reset_peak_live_bytes(); + let before = allocation_stats(); + + // SAFETY: `layout` is valid and the returned pointer is checked before use. + let pointer = unsafe { allocator.alloc_zeroed(layout) }; + assert!(!pointer.is_null()); + assert!( + // SAFETY: the allocation above contains 32 initialized bytes. + unsafe { slice::from_raw_parts(pointer, 32) }.iter().all(|byte| *byte == 0) + ); + + // SAFETY: `pointer` and `layout` identify the live allocation above. + let pointer = unsafe { allocator.realloc(pointer, layout, 64) }; + assert!(!pointer.is_null()); + let grown = Layout::from_size_align(64, 8).expect("valid grown layout"); + // SAFETY: `pointer` and `grown` identify the live reallocated block. + unsafe { allocator.dealloc(pointer, grown) }; + + let delta = allocation_delta(before, allocation_stats()); + assert_eq!(delta.allocations, 2); + assert_eq!(delta.allocated_bytes, 64); + assert_eq!(delta.live_bytes, 0); + assert_eq!(delta.peak_live_bytes, 64); + } + + #[test] + fn allocating_and_shrinking_preserve_memory_and_counters() { + let _serial = TEST_LOCK.lock().unwrap_or_else(PoisonError::into_inner); + let allocator = TrackingAllocator; + let layout = Layout::from_size_align(64, 8).expect("valid test layout"); + reset_peak_live_bytes(); + let before = allocation_stats(); + + // SAFETY: `layout` is valid and the returned pointer is checked before use. + let pointer = unsafe { allocator.alloc(layout) }; + assert!(!pointer.is_null()); + // SAFETY: the allocation above contains 64 writable bytes. + unsafe { pointer.write_bytes(0x5a, 64) }; + + // SAFETY: `pointer` and `layout` identify the live allocation above. + let pointer = unsafe { allocator.realloc(pointer, layout, 16) }; + assert!(!pointer.is_null()); + assert!( + // SAFETY: the replacement allocation contains the preserved 16-byte prefix. + unsafe { slice::from_raw_parts(pointer, 16) }.iter().all(|byte| *byte == 0x5a) + ); + let shrunk = Layout::from_size_align(16, 8).expect("valid shrunk layout"); + // SAFETY: `pointer` and `shrunk` identify the live reallocated block. + unsafe { allocator.dealloc(pointer, shrunk) }; + + let delta = allocation_delta(before, allocation_stats()); + assert_eq!(delta.allocations, 1); + assert_eq!(delta.allocated_bytes, 64); + assert_eq!(delta.live_bytes, 0); + assert_eq!(delta.peak_live_bytes, 64); + } + + #[test] + fn allocation_delta_saturates_independent_counters() { + let before = AllocationStats { + allocations: 10, + allocated_bytes: 20, + live_bytes: 30, + peak_live_bytes: 40, + }; + let after = AllocationStats { + allocations: 5, + allocated_bytes: 25, + live_bytes: 10, + peak_live_bytes: 35, + }; + assert_eq!( + allocation_delta(before, after), + AllocationStats { + allocations: 0, + allocated_bytes: 5, + live_bytes: 0, + peak_live_bytes: 5, + } + ); + } + + #[test] + fn failed_allocator_operations_leave_counters_unchanged() { + let _serial = TEST_LOCK.lock().unwrap_or_else(PoisonError::into_inner); + let layout = Layout::from_size_align(32, 8).expect("valid test layout"); + let before = allocation_stats(); + + let allocated = record_allocation(ptr::null_mut(), layout.size()); + assert!(allocated.is_null()); + let zeroed = record_allocation(ptr::null_mut(), layout.size()); + assert!(zeroed.is_null()); + let reallocated = record_reallocation(ptr::null_mut(), layout.size(), 64); + assert!(reallocated.is_null()); + assert_eq!(allocation_stats(), before); + } + + #[test] + fn counter_lock_waits_until_the_current_guard_is_released() { + let _serial = TEST_LOCK.lock().unwrap_or_else(PoisonError::into_inner); + let guard = lock_counters(); + let (attempted_send, attempted_receive) = mpsc::sync_channel(0); + let (resume_send, resume_receive) = mpsc::sync_channel(0); + let (acquired_send, acquired_receive) = mpsc::sync_channel(0); + + thread::scope(|scope| { + let waiter = scope.spawn(move || { + let waiter_guard = lock_counters_with_attempt_hook(|acquired| { + attempted_send.send(acquired).unwrap(); + resume_receive.recv().unwrap(); + }); + acquired_send.send(()).unwrap(); + drop(waiter_guard); + }); + + let acquired_while_locked = attempted_receive.recv().unwrap(); + let completed_while_locked = acquired_receive.try_recv(); + drop(guard); + + // Weak compare-exchange can fail spuriously after the guard is released. + let mut acquired = acquired_while_locked; + while !acquired { + resume_send.send(()).unwrap(); + acquired = attempted_receive.recv().unwrap(); + } + let completed_before_resume = acquired_receive.try_recv(); + resume_send.send(()).unwrap(); + acquired_receive.recv().unwrap(); + waiter.join().unwrap(); + + assert!(!acquired_while_locked); + assert_eq!(completed_while_locked, Err(TryRecvError::Empty)); + assert_eq!(completed_before_resume, Err(TryRecvError::Empty)); + }); + } +} diff --git a/crates/http_headers_simd/src/uri.rs b/crates/http_headers_simd/src/uri.rs new file mode 100644 index 000000000..01f024999 --- /dev/null +++ b/crates/http_headers_simd/src/uri.rs @@ -0,0 +1,144 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Scalar URI-reference subset classification used as the SIMD oracle. + +/// Bytes needing no percent-decoding in a path, query, or fragment. +/// +/// This is `unreserved`, `sub-delims`, and `:`, `@`, `/`, `?`, `#`; `%` is +/// omitted so escapes fall back to full parsing. +/// +/// [RFC 3986 section 2]: https://www.rfc-editor.org/rfc/rfc3986#section-2 +const SIMPLE_URI_BITMAP: [u64; 4] = [0xafff_ffda_0000_0000, 0x47ff_fffe_87ff_ffff, 0, 0]; + +/// Returns whether one byte belongs to the separator-free URI subset. +pub(super) fn is_simple_uri_byte(byte: u8) -> bool { + let word = SIMPLE_URI_BITMAP[usize::from(byte >> 6)]; + word & (1_u64 << (byte & 63)) != 0 +} + +/// Returns whether `bytes` may open an origin-relative reference. +/// +/// Rejects `//authority` because its host follows a different grammar. +pub(super) fn has_origin_relative_prefix(bytes: &[u8]) -> bool { + bytes.first() == Some(&b'/') && bytes.get(1) != Some(&b'/') +} + +/// Checks the remaining bytes of a reference whose prefix already matched. +/// +/// `hashes` counts accepted `#` bytes; at most one may appear overall. +pub(super) fn simple_uri_tail(bytes: &[u8], hashes: u32) -> bool { + let mut seen = hashes; + for &byte in bytes { + if !is_simple_uri_byte(byte) { + return false; + } + if byte == b'#' { + if seen != 0 { + return false; + } + seen = 1; + } + } + true +} + +/// Bytes that may appear in a `reg-name` host. +/// +/// This is `unreserved` plus `sub-delims`; `%`, `:`, `@`, and `[` are omitted +/// so escapes, ports, userinfo, and IP literals are recognized positionally. +/// +/// [RFC 3986 section 2]: https://www.rfc-editor.org/rfc/rfc3986#section-2 +const HOST_BITMAP: [u64; 4] = [0x2bff_7fd2_0000_0000, 0x47ff_fffe_87ff_fffe, 0, 0]; + +/// Returns whether one byte belongs to the escape-free `reg-name` subset. +fn is_host_byte(byte: u8) -> bool { + let word = HOST_BITMAP[usize::from(byte >> 6)]; + word & (1_u64 << (byte & 63)) != 0 +} + +/// Returns the offset of the path in a `scheme "://" host [":" port]` prefix. +/// +/// `None` leaves userinfo, IP literals, percent escapes, and empty hosts to the +/// full parser. Bytes after the returned offset use origin-relative rules. +pub(super) fn simple_authority_end(bytes: &[u8]) -> Option { + if !bytes.first()?.is_ascii_alphabetic() { + return None; + } + let mut index = 1; + while let Some(&byte) = bytes.get(index) { + if byte == b':' { + break; + } + if !byte.is_ascii_alphanumeric() && !matches!(byte, b'+' | b'-' | b'.') { + return None; + } + index += 1; + } + if bytes.get(index..index.checked_add(3)?) != Some(b"://".as_slice()) { + return None; + } + index += 3; + + let host_start = index; + while bytes.get(index).copied().is_some_and(is_host_byte) { + index += 1; + } + if index == host_start { + return None; + } + if bytes.get(index) == Some(&b':') { + index += 1; + while bytes.get(index).copied().is_some_and(|byte| byte.is_ascii_digit()) { + index += 1; + } + } + match bytes.get(index) { + None | Some(b'/' | b'?' | b'#') => Some(index), + Some(_) => None, + } +} + +/// Runs the whole check without any architecture-specific acceleration. +#[cfg(any(feature = "benchmarking", test))] +pub(super) fn is_simple_uri_path(bytes: &[u8]) -> bool { + has_origin_relative_prefix(bytes) && simple_uri_tail(bytes, 0) +} + +/// Runs the whole reference check without architecture-specific acceleration. +#[cfg(test)] +pub(super) fn is_simple_uri_reference(bytes: &[u8]) -> bool { + if has_origin_relative_prefix(bytes) { + return simple_uri_tail(bytes, 0); + } + match simple_authority_end(bytes) { + Some(end) => simple_uri_tail(&bytes[end..], 0), + None => false, + } +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use super::*; + + #[test] + fn origin_relative_prefix_excludes_authorities_and_relative_paths() { + assert!(has_origin_relative_prefix(b"/")); + assert!(has_origin_relative_prefix(b"/path")); + assert!(!has_origin_relative_prefix(b"")); + assert!(!has_origin_relative_prefix(b"path")); + assert!(!has_origin_relative_prefix(b"//example.com/path")); + } + + #[test] + fn tail_rejects_escapes_controls_and_multiple_fragments() { + assert!(simple_uri_tail(b"/a/b?x=1#top", 0)); + assert!(simple_uri_tail(b"continuation", 1)); + assert!(!simple_uri_tail(b"#second", 1)); + assert!(!simple_uri_tail(b"/percent%20escape", 0)); + assert!(!simple_uri_tail(b"/control\n", 0)); + assert!(is_simple_uri_path(b"/docs/index.html?x=1#top")); + assert!(!is_simple_uri_path(b"/docs#one#two")); + } +} diff --git a/crates/http_headers_simd/src/x86.rs b/crates/http_headers_simd/src/x86.rs new file mode 100644 index 000000000..d1b1fcc15 --- /dev/null +++ b/crates/http_headers_simd/src/x86.rs @@ -0,0 +1,1247 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! x86 SSE2, SSSE3, and SSE4.2 implementations of shared byte scanners. + +#[cfg(target_arch = "x86")] +use core::arch::x86::*; +#[cfg(target_arch = "x86_64")] +use core::arch::x86_64::*; + +use crate::list; +use crate::list::{EmptyMembers, ListScan, TokenListScan}; +use crate::range::{RangeMasks, WINDOW}; + +/// All x86 backends in this module process one 128-bit vector at a time. +const WIDTH: usize = 16; + +/// The same width counted in mask lanes. +const WIDTH_LANES: u32 = 16; + +/// # Safety +/// +/// The processor must support SSE2. +#[target_feature(enable = "sse2")] +pub(super) unsafe fn is_token(bytes: &[u8]) -> bool { + if bytes.is_empty() { + return false; + } + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + if token_mask(value) != 0xffff { + return false; + } + offset += WIDTH; + } + if offset == 0 || bytes.len() - offset < 8 { + return crate::scalar::all_token_bytes(&bytes[offset..]); + } + let pointer = bytes.as_ptr().wrapping_add(bytes.len() - WIDTH).cast(); + // SAFETY: at least one full vector was consumed and this load ends at the slice end. + token_mask(unsafe { _mm_loadu_si128(pointer) }) == 0xffff +} + +/// # Safety +/// +/// The processor must support SSE2. +#[target_feature(enable = "sse2")] +pub(super) unsafe fn is_token68_sse2(bytes: &[u8]) -> bool { + let mut offset = 0; + while offset + 2 * WIDTH <= bytes.len() { + let first_pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + 2 * WIDTH <= bytes.len()` bounds this unaligned load. + let first = unsafe { _mm_loadu_si128(first_pointer) }; + let second_pointer = bytes.as_ptr().wrapping_add(offset + WIDTH).cast(); + // SAFETY: `offset + 2 * WIDTH <= bytes.len()` bounds this unaligned load. + let second = unsafe { _mm_loadu_si128(second_pointer) }; + if token68_data_mask(first) != 0xffff || token68_data_mask(second) != 0xffff { + return crate::scalar::token68_tail(&bytes[offset..], offset != 0); + } + offset += 2 * WIDTH; + } + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + if token68_data_mask(value) != 0xffff { + return crate::scalar::token68_tail(&bytes[offset..], offset != 0); + } + offset += WIDTH; + } + crate::scalar::token68_tail(&bytes[offset..], offset != 0) +} + +/// # Safety +/// +/// The processor must support SSE4.2. +#[target_feature(enable = "sse4.2")] +pub(super) unsafe fn is_token68_sse42(bytes: &[u8]) -> bool { + if bytes.len() < WIDTH { + return crate::scalar::is_token68(bytes); + } + let mut offset = 0; + while offset + 2 * WIDTH <= bytes.len() { + let first_pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + 2 * WIDTH <= bytes.len()` bounds this unaligned load. + let first = unsafe { _mm_loadu_si128(first_pointer) }; + let second_pointer = bytes.as_ptr().wrapping_add(offset + WIDTH).cast(); + // SAFETY: `offset + 2 * WIDTH <= bytes.len()` bounds this unaligned load. + let second = unsafe { _mm_loadu_si128(second_pointer) }; + if token68_valid_mask_sse42(first) != 0xffff || token68_valid_mask_sse42(second) != 0xffff { + return crate::scalar::token68_tail(&bytes[offset..], offset != 0); + } + offset += 2 * WIDTH; + } + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + if token68_valid_mask_sse42(value) != 0xffff { + return crate::scalar::token68_tail(&bytes[offset..], offset != 0); + } + offset += WIDTH; + } + if offset == bytes.len() { + return offset != 0; + } + let final_offset = bytes.len() - WIDTH; + let pointer = bytes.as_ptr().wrapping_add(final_offset).cast(); + // SAFETY: `bytes.len() >= WIDTH` and this load ends exactly at the slice end. + let value = unsafe { _mm_loadu_si128(pointer) }; + let valid = token68_valid_mask_sse42(value); + if valid == 0xffff { + return true; + } + let equals = _mm_movemask_epi8(_mm_cmpeq_epi8(value, _mm_set1_epi8(b'='.cast_signed()))); + if equals == 0 { + return false; + } + let first = equals.trailing_zeros() as usize; + let expected = (0xffff_i32 << first) & 0xffff; + valid | equals == 0xffff && equals == expected && final_offset + first != 0 +} + +/// # Safety +/// +/// The processor must support SSE2. +#[target_feature(enable = "sse2")] +pub(super) unsafe fn is_field_value(bytes: &[u8]) -> bool { + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + let ascii = _mm_cmpgt_epi8(value, _mm_set1_epi8(-1)); + let control = _mm_and_si128(ascii, _mm_cmpgt_epi8(_mm_set1_epi8(0x20), value)); + let invalid_control = _mm_andnot_si128(_mm_cmpeq_epi8(value, _mm_set1_epi8(9)), control); + let del = _mm_cmpeq_epi8(value, _mm_set1_epi8(0x7f)); + if _mm_movemask_epi8(_mm_or_si128(invalid_control, del)) != 0 { + return false; + } + offset += WIDTH; + } + if offset == 0 || bytes.len() - offset == 1 { + return crate::scalar::is_field_value(&bytes[offset..]); + } + let pointer = bytes.as_ptr().wrapping_add(bytes.len() - WIDTH).cast(); + // SAFETY: at least one full vector was consumed and this load ends at the slice end. + let value = unsafe { _mm_loadu_si128(pointer) }; + let ascii = _mm_cmpgt_epi8(value, _mm_set1_epi8(-1)); + let control = _mm_and_si128(ascii, _mm_cmpgt_epi8(_mm_set1_epi8(0x20), value)); + let invalid_control = _mm_andnot_si128(_mm_cmpeq_epi8(value, _mm_set1_epi8(9)), control); + let del = _mm_cmpeq_epi8(value, _mm_set1_epi8(0x7f)); + _mm_movemask_epi8(_mm_or_si128(invalid_control, del)) == 0 +} + +/// # Safety +/// +/// The processor must support SSE2, and `right.len()` must equal `left.len()` +/// because vector loads from both slices are bounded using `left.len()`. +#[target_feature(enable = "sse2")] +pub(super) unsafe fn eq_ignore_ascii_case(left: &[u8], right: &[u8]) -> bool { + let mut offset = 0; + while offset + WIDTH <= left.len() { + let left_pointer = left.as_ptr().wrapping_add(offset).cast(); + // SAFETY: the dispatcher supplies equal-length slices and the loop bounds this load. + let left_value = unsafe { _mm_loadu_si128(left_pointer) }; + let right_pointer = right.as_ptr().wrapping_add(offset).cast(); + // SAFETY: the dispatcher supplies equal-length slices and the loop bounds this load. + let right_value = unsafe { _mm_loadu_si128(right_pointer) }; + if _mm_movemask_epi8(_mm_cmpeq_epi8(lower(left_value), lower(right_value))) != 0xffff { + return false; + } + offset += WIDTH; + } + crate::scalar::eq_ignore_ascii_case(&left[offset..], &right[offset..]) +} + +/// # Safety +/// +/// The processor must support SSE2. +#[target_feature(enable = "sse2")] +pub(super) unsafe fn find_interesting(bytes: &[u8]) -> Option { + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + let matches = b",;\"\\ \t".iter().fold(_mm_setzero_si128(), |mask, byte| { + _mm_or_si128(mask, _mm_cmpeq_epi8(value, _mm_set1_epi8(byte.cast_signed()))) + }); + let found = _mm_movemask_epi8(matches); + if found != 0 { + return Some(offset + found.trailing_zeros() as usize); + } + offset += WIDTH; + } + crate::scalar::find_interesting(&bytes[offset..]).map(|index| offset + index) +} + +/// # Safety +/// +/// The processor must support SSE2. +#[target_feature(enable = "sse2")] +pub(super) unsafe fn find_either(bytes: &[u8], first: u8, second: u8) -> Option { + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + let matches = _mm_or_si128( + _mm_cmpeq_epi8(value, _mm_set1_epi8(first.cast_signed())), + _mm_cmpeq_epi8(value, _mm_set1_epi8(second.cast_signed())), + ); + let found = _mm_movemask_epi8(matches); + if found != 0 { + return Some(offset + found.trailing_zeros() as usize); + } + offset += WIDTH; + } + crate::scalar::find_either(&bytes[offset..], first, second).map(|index| offset + index) +} + +/// Scans a reference with one nibble-table lookup per lane. +/// +/// # Safety +/// +/// The processor must support SSSE3, and `bytes.len()` must be at least +/// `WIDTH` because every block is loaded relative to a full vector of input. +#[target_feature(enable = "ssse3")] +pub(super) unsafe fn is_simple_uri_tail_ssse3(bytes: &[u8]) -> bool { + if bytes.len() <= 2 * WIDTH { + let head_pointer = bytes.as_ptr().cast(); + // SAFETY: the caller guarantees at least `WIDTH` bytes. + let head = unsafe { _mm_loadu_si128(head_pointer) }; + let final_offset = bytes.len() - WIDTH; + let tail_pointer = bytes.as_ptr().wrapping_add(final_offset).cast(); + // SAFETY: this load ends exactly at the end of the slice. + let tail = unsafe { _mm_loadu_si128(tail_pointer) }; + let head_rejects = simple_uri_reject_mask_ssse3(head); + let tail_rejects = simple_uri_reject_mask_ssse3(tail); + // Lanes shared by both blocks are rejected in `head` too, so the overlap only has to + // be masked off once a rejected lane forces the fragment count to be exact. + if head_rejects | tail_rejects == 0 { + return true; + } + let repeated = (1_u32 << (2 * WIDTH - bytes.len())) - 1; + let (Some(head_hashes), Some(tail_hashes)) = (fragment_count(head, head_rejects), fragment_count(tail, tail_rejects & !repeated)) + else { + return false; + }; + return head_hashes + tail_hashes <= 1; + } + let mut offset = 0; + let mut hashes = 0_u32; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + let rejects = simple_uri_reject_mask_ssse3(value); + if rejects != 0 { + let Some(count) = fragment_count(value, rejects) else { + return false; + }; + hashes += count; + if hashes > 1 { + return false; + } + } + offset += WIDTH; + } + if offset == bytes.len() { + return true; + } + let final_offset = bytes.len() - WIDTH; + let pointer = bytes.as_ptr().wrapping_add(final_offset).cast(); + // SAFETY: the caller guarantees `bytes.len() >= WIDTH`, so this load ends at the slice end. + let value = unsafe { _mm_loadu_si128(pointer) }; + let repeated = (1_u32 << (offset - final_offset)) - 1; + let rejects = simple_uri_reject_mask_ssse3(value) & !repeated; + if rejects == 0 { + return true; + } + let Some(count) = fragment_count(value, rejects) else { + return false; + }; + hashes + count <= 1 +} + +/// Checks that every byte is a base64 alphabet character. +/// +/// # Safety +/// +/// The processor must support SSE4.2, and `bytes.len()` must be at least `WIDTH` +/// because every block is loaded relative to a full vector of input. +#[target_feature(enable = "sse4.2")] +pub(super) unsafe fn all_base64_alphabet_sse42(bytes: &[u8]) -> bool { + if bytes.len() <= 2 * WIDTH { + let head_pointer = bytes.as_ptr().cast(); + // SAFETY: the caller guarantees at least `WIDTH` bytes. + let head = unsafe { _mm_loadu_si128(head_pointer) }; + let tail_pointer = bytes.as_ptr().wrapping_add(bytes.len() - WIDTH).cast(); + // SAFETY: this load ends exactly at the end of the slice. + let tail = unsafe { _mm_loadu_si128(tail_pointer) }; + return base64_reject_mask_sse42(head) | base64_reject_mask_sse42(tail) == 0; + } + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + if base64_reject_mask_sse42(value) != 0 { + return false; + } + offset += WIDTH; + } + if offset == bytes.len() { + return true; + } + let pointer = bytes.as_ptr().wrapping_add(bytes.len() - WIDTH).cast(); + // SAFETY: the caller guarantees `bytes.len() >= WIDTH`, so this load ends at the slice end. + let value = unsafe { _mm_loadu_si128(pointer) }; + base64_reject_mask_sse42(value) == 0 +} + +/// Checks that every byte is a base64 alphabet character using range compares. +/// +/// # Safety +/// +/// The processor must support SSE2, and `bytes.len()` must be at least `WIDTH` +/// because every block is loaded relative to a full vector of input. +#[target_feature(enable = "sse2")] +pub(super) unsafe fn all_base64_alphabet_sse2(bytes: &[u8]) -> bool { + if bytes.len() <= 2 * WIDTH { + let head_pointer = bytes.as_ptr().cast(); + // SAFETY: the caller guarantees at least `WIDTH` bytes. + let head = unsafe { _mm_loadu_si128(head_pointer) }; + let tail_pointer = bytes.as_ptr().wrapping_add(bytes.len() - WIDTH).cast(); + // SAFETY: this load ends exactly at the end of the slice. + let tail = unsafe { _mm_loadu_si128(tail_pointer) }; + return base64_reject_mask(head) | base64_reject_mask(tail) == 0; + } + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + if base64_reject_mask(value) != 0 { + return false; + } + offset += WIDTH; + } + if offset == bytes.len() { + return true; + } + let pointer = bytes.as_ptr().wrapping_add(bytes.len() - WIDTH).cast(); + // SAFETY: the caller guarantees `bytes.len() >= WIDTH`, so this load ends at the slice end. + let value = unsafe { _mm_loadu_si128(pointer) }; + base64_reject_mask(value) == 0 +} + +/// Scans a reference with one packed range comparison per block. +/// +/// # Safety +/// +/// The processor must support SSE4.2, and `bytes.len()` must be at least `WIDTH` +/// because every block is loaded relative to a full vector of input. +#[target_feature(enable = "sse4.2")] +pub(super) unsafe fn is_simple_uri_tail_sse42(bytes: &[u8]) -> bool { + if bytes.len() <= 2 * WIDTH { + let head_pointer = bytes.as_ptr().cast(); + // SAFETY: the caller guarantees at least `WIDTH` bytes. + let head = unsafe { _mm_loadu_si128(head_pointer) }; + let final_offset = bytes.len() - WIDTH; + let tail_pointer = bytes.as_ptr().wrapping_add(final_offset).cast(); + // SAFETY: this load ends exactly at the end of the slice. + let tail = unsafe { _mm_loadu_si128(tail_pointer) }; + let head_rejects = simple_uri_reject_mask_sse42(head); + let tail_rejects = simple_uri_reject_mask_sse42(tail); + // Lanes shared by both blocks are rejected in `head` too, so the overlap only has to + // be masked off once a rejected lane forces the fragment count to be exact. + if head_rejects | tail_rejects == 0 { + return true; + } + let repeated = (1_u32 << (2 * WIDTH - bytes.len())) - 1; + let (Some(head_hashes), Some(tail_hashes)) = (fragment_count(head, head_rejects), fragment_count(tail, tail_rejects & !repeated)) + else { + return false; + }; + return head_hashes + tail_hashes <= 1; + } + let mut offset = 0; + let mut hashes = 0_u32; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + let rejects = simple_uri_reject_mask_sse42(value); + if rejects != 0 { + let Some(count) = fragment_count(value, rejects) else { + return false; + }; + hashes += count; + if hashes > 1 { + return false; + } + } + offset += WIDTH; + } + if offset == bytes.len() { + return true; + } + let final_offset = bytes.len() - WIDTH; + let pointer = bytes.as_ptr().wrapping_add(final_offset).cast(); + // SAFETY: the caller guarantees `bytes.len() >= WIDTH`, so this load ends at the slice end. + let value = unsafe { _mm_loadu_si128(pointer) }; + let repeated = (1_u32 << (offset - final_offset)) - 1; + let rejects = simple_uri_reject_mask_sse42(value) & !repeated; + if rejects == 0 { + return true; + } + let Some(count) = fragment_count(value, rejects) else { + return false; + }; + hashes + count <= 1 +} + +/// Scans a reference with range compares instead of a nibble table. +/// +/// # Safety +/// +/// The processor must support SSE2, and `bytes.len()` must be at least `WIDTH` +/// because every block is loaded relative to a full vector of input. +#[target_feature(enable = "sse2")] +pub(super) unsafe fn is_simple_uri_tail_sse2(bytes: &[u8]) -> bool { + if bytes.len() <= 2 * WIDTH { + let head_pointer = bytes.as_ptr().cast(); + // SAFETY: the caller guarantees at least `WIDTH` bytes. + let head = unsafe { _mm_loadu_si128(head_pointer) }; + let final_offset = bytes.len() - WIDTH; + let tail_pointer = bytes.as_ptr().wrapping_add(final_offset).cast(); + // SAFETY: this load ends exactly at the end of the slice. + let tail = unsafe { _mm_loadu_si128(tail_pointer) }; + let head_rejects = simple_uri_reject_mask(head); + let tail_rejects = simple_uri_reject_mask(tail); + // Lanes shared by both blocks are rejected in `head` too, so the overlap only has to + // be masked off once a rejected lane forces the fragment count to be exact. + if head_rejects | tail_rejects == 0 { + return true; + } + let repeated = (1_u32 << (2 * WIDTH - bytes.len())) - 1; + let (Some(head_hashes), Some(tail_hashes)) = (fragment_count(head, head_rejects), fragment_count(tail, tail_rejects & !repeated)) + else { + return false; + }; + return head_hashes + tail_hashes <= 1; + } + let mut offset = 0; + let mut hashes = 0_u32; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + let rejects = simple_uri_reject_mask(value); + if rejects != 0 { + let Some(count) = fragment_count(value, rejects) else { + return false; + }; + hashes += count; + if hashes > 1 { + return false; + } + } + offset += WIDTH; + } + if offset == bytes.len() { + return true; + } + let final_offset = bytes.len() - WIDTH; + let pointer = bytes.as_ptr().wrapping_add(final_offset).cast(); + // SAFETY: the caller guarantees `bytes.len() >= WIDTH`, so this load ends at the slice end. + let value = unsafe { _mm_loadu_si128(pointer) }; + let repeated = (1_u32 << (offset - final_offset)) - 1; + let rejects = simple_uri_reject_mask(value) & !repeated; + if rejects == 0 { + return true; + } + let Some(count) = fragment_count(value, rejects) else { + return false; + }; + hashes + count <= 1 +} + +/// Scans a comma-separated token list with one packed range comparison per block. +/// +/// # Safety +/// +/// The processor must support SSE4.2, and `bytes.len()` must be at least +/// `WIDTH` because the trailing run is folded from a full vector of input. +#[target_feature(enable = "sse4.2")] +pub(super) unsafe fn scan_token_list_sse42(bytes: &[u8], empty: EmptyMembers) -> TokenListScan { + let mut state = ListScan::new(); + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + if !state.push_block(token_mask_sse42(value), ows_mask(value), comma_mask(value)) { + return TokenListScan::Rejected; + } + offset += WIDTH; + } + if offset == bytes.len() { + return state.finish(&[], empty); + } + let final_offset = bytes.len() - WIDTH; + let pointer = bytes.as_ptr().wrapping_add(final_offset).cast(); + // SAFETY: the caller guarantees `bytes.len() >= WIDTH`, so this load ends at the slice end. + let value = unsafe { _mm_loadu_si128(pointer) }; + let shift = shift_of(offset, final_offset); + if !state.push_partial_block( + token_mask_sse42(value) >> shift, + ows_mask(value) >> shift, + comma_mask(value) >> shift, + WIDTH_LANES - shift, + ) { + return TokenListScan::Rejected; + } + state.finish(&[], empty) +} + +/// Scans a token list of one to two vectors as a single fold. +/// +/// # Safety +/// +/// The processor must support SSE4.2, `bytes.len()` must be at least `WIDTH` +/// and at most twice it, and `two` must say whether the line runs past the +/// first vector, because both vectors are loaded whole. +#[target_feature(enable = "sse4.2")] +#[inline] +pub(super) unsafe fn scan_short_token_list_sse42(bytes: &[u8], empty: EmptyMembers, two: bool) -> TokenListScan { + let last = bytes.len() - WIDTH; + // SAFETY: the caller guarantees at least `WIDTH` bytes from either end. + let (first, final_value) = unsafe { short_windows(bytes, last) }; + let mut state = ListScan::new(); + let folded = if two { + let shift = WIDTH_LANES - lanes_of(last); + state.push_line( + list::join_lanes(token_mask_sse42(first), token_mask_sse42(final_value), shift), + list::join_lanes(ows_mask(first), ows_mask(final_value), shift), + list::join_lanes(comma_mask(first), comma_mask(final_value), shift), + WIDTH_LANES + lanes_of(last), + ) + } else { + state.push_block(token_mask_sse42(first), ows_mask(first), comma_mask(first)) + }; + if !folded { + return TokenListScan::Rejected; + } + state.finish(&[], empty) +} + +/// Scans a token list of one to two vectors as a single fold. +/// +/// # Safety +/// +/// The processor must support SSE2, `bytes.len()` must be at least `WIDTH` +/// and at most twice it, and `two` must say whether the line runs past the +/// first vector, because both vectors are loaded whole. +#[target_feature(enable = "sse2")] +#[inline] +pub(super) unsafe fn scan_short_token_list_sse2(bytes: &[u8], empty: EmptyMembers, two: bool) -> TokenListScan { + let last = bytes.len() - WIDTH; + // SAFETY: the caller guarantees at least `WIDTH` bytes from either end. + let (first, final_value) = unsafe { short_windows(bytes, last) }; + let mut state = ListScan::new(); + let folded = if two { + let shift = WIDTH_LANES - lanes_of(last); + state.push_line( + list::join_lanes(token_mask(first).cast_unsigned(), token_mask(final_value).cast_unsigned(), shift), + list::join_lanes(ows_mask(first), ows_mask(final_value), shift), + list::join_lanes(comma_mask(first), comma_mask(final_value), shift), + WIDTH_LANES + lanes_of(last), + ) + } else { + state.push_block(token_mask(first).cast_unsigned(), ows_mask(first), comma_mask(first)) + }; + if !folded { + return TokenListScan::Rejected; + } + state.finish(&[], empty) +} + +/// Loads the vector a line opens with and the vector it ends with. +/// +/// # Safety +/// +/// The processor must support SSE2, and `bytes` must hold at least `WIDTH` +/// bytes from offset zero and at least `WIDTH` bytes from `last`. +#[target_feature(enable = "sse2")] +unsafe fn short_windows(bytes: &[u8], last: usize) -> (__m128i, __m128i) { + let head = bytes.as_ptr().cast(); + let tail = bytes.as_ptr().wrapping_add(last).cast(); + // SAFETY: the caller guarantees a whole vector at offset zero. + let first = unsafe { _mm_loadu_si128(head) }; + // SAFETY: the caller guarantees a whole vector at `last`. + let final_value = unsafe { _mm_loadu_si128(tail) }; + (first, final_value) +} + +/// The lane count an offset within one vector stands for. +fn lanes_of(offset: usize) -> u32 { + u32::try_from(offset).unwrap_or(WIDTH_LANES) +} + +/// The number of lanes an overlapping final block repeats. +/// +/// Both offsets index the same line and the reload starts no later than the +/// bytes still owed, so the difference is smaller than one vector and the +/// masks keep every lane the scan has not folded yet. +fn shift_of(offset: usize, final_offset: usize) -> u32 { + u32::try_from(offset - final_offset).unwrap_or(WIDTH_LANES) +} + +/// Scans a comma-separated token list with range compares instead of a nibble table. +/// +/// # Safety +/// +/// The processor must support SSE2, and `bytes.len()` must be at least `WIDTH` +/// because the trailing run is folded from a full vector of input. +#[target_feature(enable = "sse2")] +pub(super) unsafe fn scan_token_list_sse2(bytes: &[u8], empty: EmptyMembers) -> TokenListScan { + let mut state = ListScan::new(); + let mut offset = 0; + while offset + WIDTH <= bytes.len() { + let pointer = bytes.as_ptr().wrapping_add(offset).cast(); + // SAFETY: `offset + WIDTH <= bytes.len()` permits this unaligned 16-byte load. + let value = unsafe { _mm_loadu_si128(pointer) }; + let tokens = token_mask(value).cast_unsigned(); + if !state.push_block(tokens, ows_mask(value), comma_mask(value)) { + return TokenListScan::Rejected; + } + offset += WIDTH; + } + if offset == bytes.len() { + return state.finish(&[], empty); + } + let final_offset = bytes.len() - WIDTH; + let pointer = bytes.as_ptr().wrapping_add(final_offset).cast(); + // SAFETY: the caller guarantees `bytes.len() >= WIDTH`, so this load ends at the slice end. + let value = unsafe { _mm_loadu_si128(pointer) }; + let shift = shift_of(offset, final_offset); + if !state.push_partial_block( + token_mask(value).cast_unsigned() >> shift, + ows_mask(value) >> shift, + comma_mask(value) >> shift, + WIDTH_LANES - shift, + ) { + return TokenListScan::Rejected; + } + state.finish(&[], empty) +} + +/// Classifies one window for the byte-range grammar. +/// +/// Digits fall out of one biased unsigned clamp, and the three delimiters and +/// the leading-zero marker are plain equality comparisons, so the whole +/// classification is five compares and five movemasks. +/// +/// # Safety +/// +/// The processor must support SSE2. +#[inline] +#[target_feature(enable = "sse2")] +pub(super) unsafe fn range_masks(window: &[u8; WINDOW]) -> RangeMasks { + let pointer = window.as_ptr().cast(); + // SAFETY: the window is exactly 16 bytes, so this unaligned load stays in bounds. + let value = unsafe { _mm_loadu_si128(pointer) }; + let biased = _mm_sub_epi8(value, _mm_set1_epi8(b'0'.cast_signed())); + let clamped = _mm_min_epu8(biased, _mm_set1_epi8(9)); + let digits = _mm_cmpeq_epi8(clamped, biased); + let zeros = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'0'.cast_signed())); + let dashes = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'-'.cast_signed())); + let commas = _mm_cmpeq_epi8(value, _mm_set1_epi8(b','.cast_signed())); + let spaces = _mm_or_si128( + _mm_cmpeq_epi8(value, _mm_set1_epi8(b' '.cast_signed())), + _mm_cmpeq_epi8(value, _mm_set1_epi8(b'\t'.cast_signed())), + ); + RangeMasks { + digits: _mm_movemask_epi8(digits).cast_unsigned(), + dashes: _mm_movemask_epi8(dashes).cast_unsigned(), + commas: _mm_movemask_epi8(commas).cast_unsigned(), + spaces: _mm_movemask_epi8(spaces).cast_unsigned(), + zeros: _mm_movemask_epi8(zeros).cast_unsigned(), + } +} + +/// Marks the `tchar` lanes with one packed range comparison and one compare. +/// +/// The eight ranges hold every `tchar` except `|` and `~`, which are the only +/// two bytes that survive forcing bit one on and comparing against `~`. A `NUL` +/// lane truncates the implicit string, so lanes after it report no token; that +/// only ever removes token lanes, and a `NUL` is outside every list class, so +/// the block is rejected before the classification is consulted. +#[target_feature(enable = "sse4.2")] +fn token_mask_sse42(value: __m128i) -> u32 { + const RANGES: &[u8; 16] = b"!!#'*+-.09AZ^`az"; + + let ranges_pointer = RANGES.as_ptr().cast(); + // SAFETY: `RANGES` is exactly 16 bytes, so this unaligned load stays in bounds. + let ranges = unsafe { _mm_loadu_si128(ranges_pointer) }; + let in_ranges = _mm_cmpistrm::<{ _SIDD_UBYTE_OPS | _SIDD_CMP_RANGES | _SIDD_UNIT_MASK | _SIDD_POSITIVE_POLARITY }>(ranges, value); + let raised = _mm_or_si128(value, _mm_set1_epi8(2)); + let bar_or_tilde = _mm_cmpeq_epi8(raised, _mm_set1_epi8(b'~'.cast_signed())); + _mm_movemask_epi8(_mm_or_si128(in_ranges, bar_or_tilde)).cast_unsigned() +} + +/// Marks the optional-whitespace lanes. +#[target_feature(enable = "sse2")] +fn ows_mask(value: __m128i) -> u32 { + let space = _mm_cmpeq_epi8(value, _mm_set1_epi8(b' '.cast_signed())); + let tab = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'\t'.cast_signed())); + _mm_movemask_epi8(_mm_or_si128(space, tab)).cast_unsigned() +} + +/// Marks the comma lanes. +#[target_feature(enable = "sse2")] +fn comma_mask(value: __m128i) -> u32 { + _mm_movemask_epi8(_mm_cmpeq_epi8(value, _mm_set1_epi8(b','.cast_signed()))).cast_unsigned() +} + +#[target_feature(enable = "sse2")] +fn token_mask(value: __m128i) -> i32 { + let mut mask = _mm_or_si128(in_range(value, b'0', b'9'), in_range(value, b'A', b'Z')); + mask = _mm_or_si128(mask, in_range(value, b'a', b'z')); + let mask = b"!#$%&'*+-.^_`|~".iter().fold(mask, |mask, byte| { + _mm_or_si128(mask, _mm_cmpeq_epi8(value, _mm_set1_epi8(byte.cast_signed()))) + }); + _mm_movemask_epi8(mask) +} + +#[target_feature(enable = "sse2")] +fn token68_data_mask(value: __m128i) -> i32 { + let digit = in_range(value, b'0', b'9'); + let folded = _mm_or_si128(value, _mm_set1_epi8(0x20)); + let alpha = in_range(folded, b'a', b'z'); + let plus_to_slash = in_range(value, b'+', b'/'); + let symbols = _mm_andnot_si128(_mm_cmpeq_epi8(value, _mm_set1_epi8(b','.cast_signed())), plus_to_slash); + let underscore = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'_'.cast_signed())); + let tilde = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'~'.cast_signed())); + _mm_movemask_epi8(_mm_or_si128( + _mm_or_si128(digit, alpha), + _mm_or_si128(symbols, _mm_or_si128(underscore, tilde)), + )) +} + +/// Marks the token68 data lanes with a single packed range comparison. +/// +/// The `=` padding is deliberately excluded so callers can locate it themselves. +#[target_feature(enable = "sse4.2")] +fn token68_valid_mask_sse42(value: __m128i) -> i32 { + const RANGES: [u8; 16] = [ + b'0', b'9', b'A', b'Z', b'a', b'z', b'+', b'+', b'-', b'/', b'_', b'_', b'~', b'~', 0, 0, + ]; + + let ranges_pointer = RANGES.as_ptr().cast(); + // SAFETY: `RANGES` is exactly 16 bytes, so this unaligned load stays in bounds. + let ranges = unsafe { _mm_loadu_si128(ranges_pointer) }; + let matched = _mm_cmpistrm::<{ _SIDD_UBYTE_OPS | _SIDD_CMP_RANGES | _SIDD_UNIT_MASK | _SIDD_POSITIVE_POLARITY }>(ranges, value); + _mm_movemask_epi8(matched) +} + +/// Counts the `#` lanes among the rejected ones, or reports a lane that cannot start a fragment. +/// +/// Returns `None` when any rejected lane holds something other than `#`, which means the +/// reference is outside the guaranteed-valid subset and needs the full parser. +#[target_feature(enable = "sse2")] +fn fragment_count(value: __m128i, rejects: u32) -> Option { + let hashes = hash_mask(value) & rejects; + (rejects & !hashes == 0).then(|| hashes.count_ones()) +} + +/// Marks the lanes outside the base64 alphabet with one packed range comparison. +/// +/// `/` and the digits are contiguous, so the alphabet needs only four ranges. A `NUL` lane ends +/// the implicit string and forces every later lane to be reported as rejected, which is the +/// right answer because `NUL` is outside the alphabet and already rejects the whole input. +#[target_feature(enable = "sse4.2")] +fn base64_reject_mask_sse42(value: __m128i) -> u32 { + const RANGES: [u8; 16] = [b'+', b'+', b'/', b'9', b'A', b'Z', b'a', b'z', 0, 0, 0, 0, 0, 0, 0, 0]; + + let ranges_pointer = RANGES.as_ptr().cast(); + // SAFETY: `RANGES` is exactly 16 bytes, so this unaligned load stays in bounds. + let ranges = unsafe { _mm_loadu_si128(ranges_pointer) }; + let rejected = _mm_cmpistrm::<{ _SIDD_UBYTE_OPS | _SIDD_CMP_RANGES | _SIDD_BIT_MASK | _SIDD_NEGATIVE_POLARITY }>(ranges, value); + _mm_cvtsi128_si32(rejected).cast_unsigned() +} + +/// Marks the lanes outside the base64 alphabet with range compares. +#[target_feature(enable = "sse2")] +fn base64_reject_mask(value: __m128i) -> u32 { + let digit_or_slash = in_range(value, b'/', b'9'); + let letter = in_range(_mm_or_si128(value, _mm_set1_epi8(0x20)), b'a', b'z'); + let plus = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'+'.cast_signed())); + let accepted = _mm_or_si128(_mm_or_si128(digit_or_slash, letter), plus); + !_mm_movemask_epi8(accepted).cast_unsigned() & 0xffff +} + +/// Marks the lanes outside the separator-free URI subset with one packed range comparison. +/// +/// The subset is exactly eight ranges, which is the widest set `pcmpistrm` can hold, so a single +/// instruction classifies a whole block. `#` is excluded because callers count fragments +/// separately. A `NUL` lane ends the implicit string, and every lane past it is reported as +/// rejected, which only ever sends the caller to the full parser. +#[target_feature(enable = "sse4.2")] +fn simple_uri_reject_mask_sse42(value: __m128i) -> u32 { + const RANGES: &[u8; 16] = b"!!$$&;==?Z__az~~"; + + let ranges_pointer = RANGES.as_ptr().cast(); + // SAFETY: `RANGES` is exactly 16 bytes, so this unaligned load stays in bounds. + let ranges = unsafe { _mm_loadu_si128(ranges_pointer) }; + let rejected = _mm_cmpistrm::<{ _SIDD_UBYTE_OPS | _SIDD_CMP_RANGES | _SIDD_BIT_MASK | _SIDD_NEGATIVE_POLARITY }>(ranges, value); + _mm_cvtsi128_si32(rejected).cast_unsigned() +} + +/// Marks the lanes outside the separator-free URI subset with one nibble-table lookup. +#[target_feature(enable = "ssse3")] +fn simple_uri_reject_mask_ssse3(value: __m128i) -> u32 { + let low_nibble = _mm_set_epi64x(0x7cd4_5c54_5cfc_fcfc, 0xfcfc_f8fc_f8f8_fcb8_u64.cast_signed()); + let high_nibble = _mm_set_epi64x(0, 0x8040_2010_0804_0201_u64.cast_signed()); + let low = _mm_and_si128(value, _mm_set1_epi8(0x0f)); + let high = _mm_and_si128(_mm_srli_epi16::<4>(value), _mm_set1_epi8(0x0f)); + let accepted = _mm_and_si128(_mm_shuffle_epi8(low_nibble, low), _mm_shuffle_epi8(high_nibble, high)); + _mm_movemask_epi8(_mm_cmpeq_epi8(accepted, _mm_setzero_si128())).cast_unsigned() & 0xffff +} + +/// Marks the lanes outside the separator-free URI subset. +/// +/// The `&`..`?` range covers sub-delims, digits, `:`, and `;` after excluding +/// `<` and `>`. Case folding merges the letter ranges; isolated bytes use +/// equality comparisons. `#` is excluded for separate fragment counting. +#[target_feature(enable = "sse2")] +fn simple_uri_reject_mask(value: __m128i) -> u32 { + let raised = _mm_or_si128(value, _mm_set1_epi8(2)); + let punctuation = _mm_andnot_si128(_mm_cmpeq_epi8(raised, _mm_set1_epi8(0x3e)), in_range(value, b'&', b'?')); + let bang = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'!'.cast_signed())); + let dollar = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'$'.cast_signed())); + let letter = in_range(_mm_or_si128(value, _mm_set1_epi8(0x20)), b'a', b'z'); + let at = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'@'.cast_signed())); + let underscore = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'_'.cast_signed())); + let tilde = _mm_cmpeq_epi8(value, _mm_set1_epi8(b'~'.cast_signed())); + let accepted = _mm_or_si128( + _mm_or_si128(punctuation, _mm_or_si128(bang, dollar)), + _mm_or_si128(letter, _mm_or_si128(at, _mm_or_si128(underscore, tilde))), + ); + !_mm_movemask_epi8(accepted).cast_unsigned() & 0xffff +} + +#[target_feature(enable = "sse2")] +fn hash_mask(value: __m128i) -> u32 { + _mm_movemask_epi8(_mm_cmpeq_epi8(value, _mm_set1_epi8(b'#'.cast_signed()))).cast_unsigned() +} + +#[target_feature(enable = "sse2")] +fn lower(value: __m128i) -> __m128i { + _mm_or_si128(value, _mm_and_si128(in_range(value, b'A', b'Z'), _mm_set1_epi8(0x20))) +} + +#[target_feature(enable = "sse2")] +fn in_range(value: __m128i, start: u8, end: u8) -> __m128i { + let above_start = _mm_cmpgt_epi8(value, _mm_set1_epi8(start.wrapping_sub(1).cast_signed())); + let below_end = _mm_cmpgt_epi8(_mm_set1_epi8(end.wrapping_add(1).cast_signed()), value); + _mm_and_si128(above_start, below_end) +} + +#[cfg(test)] +#[cfg_attr(coverage_nightly, coverage(off))] +mod tests { + use std::arch; + use std::time::Duration; + #[cfg(not(feature = "std"))] + use std::vec; + #[cfg(not(feature = "std"))] + use std::vec::Vec; + + use super::*; + + #[test] + #[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] + fn sse2_matches_scalar() { + #[cfg(target_arch = "x86")] + if !arch::is_x86_feature_detected!("sse2") { + return; + } + let sse42 = arch::is_x86_feature_detected!("sse4.2"); + let ssse3 = arch::is_x86_feature_detected!("ssse3"); + bolero::check!() + .with_iterations(4_096) + .with_test_time(Duration::from_millis(400)) + .with_type::<(Vec, Vec)>() + .for_each(|(left, right)| { + // SAFETY: runtime detection above establishes SSE2 support. + let token = unsafe { is_token(left) }; + assert_eq!(token, crate::scalar::is_token(left)); + // SAFETY: runtime detection above establishes SSE2 support. + let token68 = unsafe { is_token68_sse2(left) }; + assert_eq!(token68, crate::scalar::is_token68(left)); + assert_eq!( + sse42.then(|| { + // SAFETY: runtime detection above establishes SSE4.2 support. + unsafe { is_token68_sse42(left) } + }), + sse42.then(|| crate::scalar::is_token68(left)) + ); + // SAFETY: runtime detection above establishes SSE2 support. + let field_value = unsafe { is_field_value(left) }; + assert_eq!(field_value, crate::scalar::is_field_value(left)); + if left.len() == right.len() { + // SAFETY: runtime detection establishes SSE2 support; lengths are equal. + let equal = unsafe { eq_ignore_ascii_case(left, right) }; + assert_eq!(equal, crate::scalar::eq_ignore_ascii_case(left, right)); + } + // SAFETY: runtime detection above establishes SSE2 support. + let interesting = unsafe { find_interesting(left) }; + assert_eq!(interesting, crate::scalar::find_interesting(left)); + if left.len() >= WIDTH { + let expected = crate::base64::all_base64_alphabet(left); + assert_eq!( + sse42.then(|| { + // SAFETY: runtime detection establishes SSE4.2 support. + unsafe { all_base64_alphabet_sse42(left) } + }), + sse42.then_some(expected) + ); + // SAFETY: runtime detection establishes SSE2 support; the length is checked. + let base64 = unsafe { all_base64_alphabet_sse2(left) }; + assert_eq!(base64, expected); + // SAFETY: runtime detection establishes SSE2 support; the length is checked. + let simple = unsafe { is_simple_uri_tail_sse2(left) }; + assert_eq!(simple, crate::uri::simple_uri_tail(left, 0)); + let expected = crate::uri::simple_uri_tail(left, 0); + assert_eq!( + sse42.then(|| { + // SAFETY: runtime detection establishes SSE4.2 support. + unsafe { is_simple_uri_tail_sse42(left) } + }), + sse42.then_some(expected) + ); + assert_eq!( + ssse3.then(|| { + // SAFETY: runtime detection establishes SSSE3 support. + unsafe { is_simple_uri_tail_ssse3(left) } + }), + ssse3.then_some(expected) + ); + } + if left.len() >= WIDTH { + for empty in [EmptyMembers::Skip, EmptyMembers::Reject] { + let expected = crate::list::scan_token_list(left, empty); + // SAFETY: runtime detection establishes SSE2 support; the length is checked. + let list = unsafe { scan_token_list_sse2(left, empty) }; + assert_eq!(list, expected); + assert_eq!( + sse42.then(|| { + // SAFETY: runtime detection establishes SSE4.2 support. + unsafe { scan_token_list_sse42(left, empty) } + }), + sse42.then_some(expected) + ); + if left.len() == WIDTH { + // SAFETY: detection above, and the line is exactly one vector. + let short = unsafe { scan_short_token_list_sse2(left, empty, false) }; + assert_eq!(short, expected); + assert_eq!( + sse42.then(|| { + // SAFETY: detection above, and the line is exactly one vector. + unsafe { scan_short_token_list_sse42(left, empty, false) } + }), + sse42.then_some(expected) + ); + } else if left.len() <= 2 * WIDTH { + // SAFETY: detection above, and the line spans two vectors. + let short = unsafe { scan_short_token_list_sse2(left, empty, true) }; + assert_eq!(short, expected); + assert_eq!( + sse42.then(|| { + // SAFETY: detection above, and the line spans two vectors. + unsafe { scan_short_token_list_sse42(left, empty, true) } + }), + sse42.then_some(expected) + ); + } + } + } + }); + } + + /// Puts every byte value in every lane of a token-list block. + /// + /// Each kernel classifies a lane into one of four roles, so a single + /// misplaced range or compare shows up as one byte in one lane disagreeing + /// with the scalar state machine, which no random search reliably finds. + #[test] + fn token_list_lanes_match_scalar_for_every_byte() { + #[cfg(target_arch = "x86")] + if !arch::is_x86_feature_detected!("sse2") { + return; + } + let sse42 = arch::is_x86_feature_detected!("sse4.2"); + for lane in 0..2 * WIDTH { + for byte in u8::MIN..=u8::MAX { + let mut block = [b'a'; 2 * WIDTH]; + block[lane] = byte; + for empty in [EmptyMembers::Skip, EmptyMembers::Reject] { + let expected = crate::list::scan_token_list(&block, empty); + // SAFETY: runtime detection above establishes SSE2 support. + let list = unsafe { scan_token_list_sse2(&block, empty) }; + assert_eq!(list, expected, "byte {byte:#04x} lane {lane}"); + assert_eq!( + sse42.then(|| { + // SAFETY: runtime detection establishes SSE4.2 support. + unsafe { scan_token_list_sse42(&block, empty) } + }), + sse42.then_some(expected), + "byte {byte:#04x} lane {lane}" + ); + } + } + } + } + + #[test] + fn range_lanes_match_scalar_for_every_byte() { + #[cfg(target_arch = "x86")] + if !arch::is_x86_feature_detected!("sse2") { + return; + } + for lane in 0..WINDOW { + for byte in u8::MIN..=u8::MAX { + let mut window = [b'0'; WINDOW]; + window[lane] = byte; + // SAFETY: runtime detection above establishes SSE2 support. + let masks = unsafe { range_masks(&window) }; + assert_eq!(masks, RangeMasks::scalar(&window), "byte {byte:#04x} lane {lane}"); + } + } + } + + #[test] + fn base64_and_uri_lanes_match_every_x86_backend() { + #[cfg(target_arch = "x86")] + if !arch::is_x86_feature_detected!("sse2") { + return; + } + let sse42 = arch::is_x86_feature_detected!("sse4.2"); + let ssse3 = arch::is_x86_feature_detected!("ssse3"); + for lane in 0..2 * WIDTH { + for byte in u8::MIN..=u8::MAX { + let mut base64_block = [b'A'; 2 * WIDTH]; + base64_block[lane] = byte; + let expected_base64 = crate::base64::all_base64_alphabet(&base64_block); + assert_eq!( + // SAFETY: runtime detection above establishes SSE2 support. + unsafe { all_base64_alphabet_sse2(&base64_block) }, + expected_base64, + "base64 byte {byte:#04x} lane {lane}" + ); + assert_eq!( + sse42.then(|| { + // SAFETY: runtime detection establishes SSE4.2 support. + unsafe { all_base64_alphabet_sse42(&base64_block) } + }), + sse42.then_some(expected_base64), + "SSE4.2 base64 byte {byte:#04x} lane {lane}" + ); + + let mut uri_block = [b'a'; 2 * WIDTH]; + uri_block[lane] = byte; + let expected_uri = crate::uri::simple_uri_tail(&uri_block, 0); + assert_eq!( + // SAFETY: runtime detection above establishes SSE2 support. + unsafe { is_simple_uri_tail_sse2(&uri_block) }, + expected_uri, + "URI byte {byte:#04x} lane {lane}" + ); + assert_eq!( + ssse3.then(|| { + // SAFETY: runtime detection establishes SSSE3 support. + unsafe { is_simple_uri_tail_ssse3(&uri_block) } + }), + ssse3.then_some(expected_uri), + "SSSE3 URI byte {byte:#04x} lane {lane}" + ); + assert_eq!( + sse42.then(|| { + // SAFETY: runtime detection establishes SSE4.2 support. + unsafe { is_simple_uri_tail_sse42(&uri_block) } + }), + sse42.then_some(expected_uri), + "SSE4.2 URI byte {byte:#04x} lane {lane}" + ); + } + } + } + + #[test] + fn long_sse2_scanners_cover_full_blocks_and_overlapping_tails() { + #[cfg(target_arch = "x86")] + if !arch::is_x86_feature_detected!("sse2") { + return; + } + + let valid_token68 = [b'a'; 64]; + // SAFETY: runtime detection above establishes SSE2 support. + assert!(unsafe { is_token68_sse2(&valid_token68) }); + let mut invalid_token68 = valid_token68; + invalid_token68[20] = b':'; + // SAFETY: runtime detection above establishes SSE2 support. + assert!(!unsafe { is_token68_sse2(&invalid_token68) }); + let mut padded_token68 = valid_token68; + padded_token68[48..].fill(b'='); + // SAFETY: runtime detection above establishes SSE2 support. + assert!(unsafe { is_token68_sse2(&padded_token68) }); + let valid_token68_tail = [b'a'; 48]; + // SAFETY: runtime detection above establishes SSE2 support. + assert!(unsafe { is_token68_sse2(&valid_token68_tail) }); + let mut invalid_token68_tail = valid_token68_tail; + invalid_token68_tail[40] = b':'; + // SAFETY: runtime detection above establishes SSE2 support. + assert!(!unsafe { is_token68_sse2(&invalid_token68_tail) }); + + for length in [48, 49] { + let valid_base64 = vec![b'A'; length]; + // SAFETY: runtime detection above establishes SSE2 support. + assert!(unsafe { all_base64_alphabet_sse2(&valid_base64) }); + let mut invalid_base64 = valid_base64; + invalid_base64[length - 1] = b'='; + // SAFETY: runtime detection above establishes SSE2 support. + assert!(!unsafe { all_base64_alphabet_sse2(&invalid_base64) }); + } + + let exact_list = [b'a'; 48]; + assert_eq!( + // SAFETY: runtime detection above establishes SSE2 support. + unsafe { scan_token_list_sse2(&exact_list, EmptyMembers::Skip) }, + TokenListScan::Members + ); + let mut tailed_list = [b'a'; 49]; + tailed_list[47] = b','; + assert_eq!( + // SAFETY: runtime detection above establishes SSE2 support. + unsafe { scan_token_list_sse2(&tailed_list, EmptyMembers::Skip) }, + TokenListScan::Members + ); + tailed_list[48] = b' '; + assert_eq!( + // SAFETY: runtime detection above establishes SSE2 support. + unsafe { scan_token_list_sse2(&tailed_list, EmptyMembers::Reject) }, + TokenListScan::Rejected + ); + tailed_list[48] = b'('; + assert_eq!( + // SAFETY: runtime detection above establishes SSE2 support. + unsafe { scan_token_list_sse2(&tailed_list, EmptyMembers::Skip) }, + TokenListScan::Rejected + ); + } + + #[test] + fn long_uri_paths_cover_every_available_x86_backend() { + #[cfg(target_arch = "x86")] + if !arch::is_x86_feature_detected!("sse2") { + return; + } + + fn exercise(scan: unsafe fn(&[u8]) -> bool) { + let valid_exact = [b'a'; 48]; + // SAFETY: the caller supplies a scanner whose feature was detected. + assert!(unsafe { scan(&valid_exact) }); + + let valid_tail = [b'a'; 49]; + // SAFETY: the caller supplies a scanner whose feature was detected. + assert!(unsafe { scan(&valid_tail) }); + + let mut invalid_block = valid_exact; + invalid_block[17] = b'%'; + // SAFETY: the caller supplies a scanner whose feature was detected. + assert!(!unsafe { scan(&invalid_block) }); + + let mut repeated_fragment = valid_exact; + repeated_fragment[1] = b'#'; + repeated_fragment[33] = b'#'; + // SAFETY: the caller supplies a scanner whose feature was detected. + assert!(!unsafe { scan(&repeated_fragment) }); + + let mut invalid_tail = valid_tail; + invalid_tail[48] = b'%'; + // SAFETY: the caller supplies a scanner whose feature was detected. + assert!(!unsafe { scan(&invalid_tail) }); + + let mut fragment_tail = valid_tail; + fragment_tail[1] = b'#'; + fragment_tail[48] = b'#'; + // SAFETY: the caller supplies a scanner whose feature was detected. + assert!(!unsafe { scan(&fragment_tail) }); + } + + exercise(is_simple_uri_tail_sse2); + let _ = arch::is_x86_feature_detected!("ssse3").then(|| { + exercise(is_simple_uri_tail_ssse3); + }); + let _ = arch::is_x86_feature_detected!("sse4.2").then(|| { + exercise(is_simple_uri_tail_sse42); + }); + } + + #[test] + fn short_sse2_list_folds_and_lane_fallbacks_are_explicit() { + #[cfg(target_arch = "x86")] + if !arch::is_x86_feature_detected!("sse2") { + return; + } + let one = [b'a'; WIDTH]; + assert_eq!( + // SAFETY: runtime detection above establishes SSE2 support. + unsafe { scan_short_token_list_sse2(&one, EmptyMembers::Skip, false) }, + TokenListScan::Members + ); + let two = [b'a'; 2 * WIDTH]; + assert_eq!( + // SAFETY: runtime detection above establishes SSE2 support. + unsafe { scan_short_token_list_sse2(&two, EmptyMembers::Skip, true) }, + TokenListScan::Members + ); + #[cfg(target_pointer_width = "32")] + assert_eq!(lanes_of(usize::MAX), u32::MAX); + #[cfg(target_pointer_width = "64")] + assert_eq!(lanes_of(usize::MAX), WIDTH_LANES); + #[cfg(target_pointer_width = "64")] + assert_eq!( + shift_of(usize::try_from(u64::from(u32::MAX) + 1).expect("value fits a 64-bit usize"), 0), + WIDTH_LANES + ); + } +} diff --git a/crates/http_headers_simd/tests/__fuzz__/campaign.toml b/crates/http_headers_simd/tests/__fuzz__/campaign.toml new file mode 100644 index 000000000..32dec2934 --- /dev/null +++ b/crates/http_headers_simd/tests/__fuzz__/campaign.toml @@ -0,0 +1,14 @@ +# Copyright (c) Microsoft Corporation. +# Licensed under the MIT License. + +[campaign] +package = "http_headers_simd" +default_engine = "libfuzzer" +default_time = "60s" +default_max_input_length = 1024 +corpus_root = "crates/http_headers_simd/tests/__fuzz__" + +[[target]] +name = "public_scanners_match_scalar_oracles" +test = "fuzz" +description = "Compare the public scanner entry points with their scalar oracles over arbitrary bytes." diff --git a/crates/http_headers_simd/tests/__fuzz__/public_scanners_match_scalar_oracles/crashes/.gitkeep b/crates/http_headers_simd/tests/__fuzz__/public_scanners_match_scalar_oracles/crashes/.gitkeep new file mode 100644 index 000000000..e69de29bb diff --git a/crates/http_headers_simd/tests/bolero_fuzz.rs b/crates/http_headers_simd/tests/bolero_fuzz.rs new file mode 100644 index 000000000..534af5220 --- /dev/null +++ b/crates/http_headers_simd/tests/bolero_fuzz.rs @@ -0,0 +1,63 @@ +// Copyright (c) Microsoft Corporation. +// Licensed under the MIT License. + +//! Bounded public-API properties for the SIMD byte scanners. + +use std::time::Duration; + +use http_headers_simd::{eq_ignore_ascii_case, find_interesting, is_field_value, is_token, is_token68}; + +const BOUNDED_ITERATIONS: usize = 4_096; +const BOUNDED_TEST_TIME: Duration = Duration::from_millis(400); +const MAX_INPUT_LENGTH: usize = 1_024; + +#[test] +#[cfg_attr(miri, ignore = "Bolero corpus replay requires filesystem access unavailable under Miri isolation")] +fn public_scanners_match_scalar_oracles() { + bolero::check!() + .with_iterations(BOUNDED_ITERATIONS) + .with_test_time(BOUNDED_TEST_TIME) + .with_type::<(Vec, Vec)>() + .for_each(|(left, right)| { + let left = &left[..left.len().min(MAX_INPUT_LENGTH)]; + let right = &right[..right.len().min(MAX_INPUT_LENGTH)]; + + let token = !left.is_empty() + && left + .iter() + .copied() + .all(|byte| byte.is_ascii_alphanumeric() || b"!#$%&'*+-.^_`|~".contains(&byte)); + let token68_data_end = left.iter().position(|byte| *byte == b'=').unwrap_or(left.len()); + let token68 = token68_data_end != 0 + && left[..token68_data_end] + .iter() + .all(|byte| byte.is_ascii_alphanumeric() || b"-._~+/".contains(byte)) + && left[token68_data_end..].iter().all(|byte| *byte == b'='); + let field_value = left.iter().copied().all(|byte| byte == b'\t' || byte >= b' ' && byte != 0x7f); + let equal = left.len() == right.len() && left.iter().zip(right).all(|(left, right)| left.eq_ignore_ascii_case(right)); + let interesting = left.iter().position(|byte| b",;\"\\ \t".contains(byte)); + + assert_eq!(is_token(left), token); + assert_eq!(is_token68(left), token68); + assert_eq!(is_field_value(left), field_value); + assert_eq!(eq_ignore_ascii_case(left, right), equal); + assert_eq!(find_interesting(left), interesting); + }); +} + +#[test] +fn scanner_boundary_regression_seeds() { + for length in [0, 1, 15, 16, 17, 31, 32, 33, 63, 64, 65, 127, 128, 129, MAX_INPUT_LENGTH] { + let mut bytes = vec![b'a'; length]; + assert_eq!(is_token(&bytes), length != 0); + assert_eq!(is_token68(&bytes), length != 0); + assert!(is_field_value(&bytes)); + assert_eq!(find_interesting(&bytes), None); + + if let Some(last) = bytes.last_mut() { + *last = b';'; + assert!(!is_token(&bytes)); + assert_eq!(find_interesting(&bytes), Some(length - 1)); + } + } +} diff --git a/docs/design/README.md b/docs/design/README.md index 332a8c20a..55d4c1bfc 100644 --- a/docs/design/README.md +++ b/docs/design/README.md @@ -7,3 +7,73 @@ CI. Do not maintain parallel implementations of Anvil checks. Keep repository-specific automation only for capabilities outside Anvil's scope, such as release tooling and tests for that tooling. Generated Anvil files are updated with `cargo anvil`, not edited directly. + +The handwritten `repository-checks.yml` also verifies the +`http_headers_simd` `no_std` contract with package-isolated unit and integration +tests on AArch64 and x86-family Linux. The generated Anvil test groups can select the +`http_headers` facade alongside its SIMD dependency, which re-enables the +dependency's `std` feature even in their `--no-default-features` leg. The +isolated jobs cover otherwise-unselected configurations rather than duplicating +an Anvil test group. They use Anvil's selected stable toolchain and participate +in the stable `Required repository checks` fan-in. + +| Target | Host | Global CPU baseline | +| --- | --- | --- | +| `aarch64-unknown-linux-gnu` | Native AArch64 Linux | Target defaults | +| `x86_64-unknown-linux-gnu` | Native x86-64 Linux | `x86-64`, including SSE2 | +| `i686-unknown-linux-gnu` | x86-64 Linux, executing 32-bit binaries | `pentium4`, including SSE2 | +| `i586-unknown-linux-gnu` | x86-64 Linux, executing 32-bit binaries | `pentium`, without SSE2 | + +`just test-http-headers-simd-no-std-arm` retains the native AArch64 host guard +and requires the no_std NEON detector and differential test names before +executing the suite. `just test-http-headers-simd-no-std-x86 --target TARGET` +requires an x86-64 Linux host, both no_std CPUID-cache tests, and the independent +runtime-feature oracle test. The 32-bit targets additionally require the +no_std SSE2 detector test; unlike x86-64, their SSE2 availability depends on the +compile-time feature. CI installs the corresponding Rust target and 32-bit +linker/runtime where needed. + +The x86-family recipe replaces `CARGO_ENCODED_RUSTFLAGS` with the selected +`-Ctarget-cpu` and `-Ctarget-feature=-ssse3,-sse4.2`. Cargo gives these encoded +flags precedence over inherited `RUSTFLAGS`, target-specific flags, and build +flags, intentionally overriding the repository's `x86-64-v3` setting. Explicit +`--target` confines these flags to target compilation rather than host build +scripts and procedural macros. Required test names are checked before the +suite runs, so re-enabling `std` or either compile-time ISA feature cannot +silently omit its cache test and pass. + +These jobs do not mask the host's actual CPUID capabilities: they exercise +cached runtime detection, not necessarily the result for a processor lacking +SSSE3 or SSE4.2. Their runtime scope is Linux, not Windows or other operating +systems. These are ordinary correctness tests, not an additional coverage or +mutation requirement for no_std-only code. + +The `http_headers` and `http_headers_simd` crates are excluded from mutation +testing through `.cargo/mutants.toml`. Their ordinary tests, coverage checks, +and Miri checks remain enabled. + +Under Miri, `http_headers` samples repetitive generated inputs and byte +substitutions while retaining fixed regressions and numeric and vector +boundaries. Large WebSocket admission-limit cases retain the actual byte, +line, and item limits and both ownership paths, without repeating every +source representation and decoding mode. Native test runs keep the full +input corpora and source/mode combinations. + +Credential zeroization checks use smaller heap-backed buffers under Miri, +preserving base64 padding shapes, allocation growth, reuse, and full wipe +assertions. Differential checks compare complete errors without repeatedly +formatting them. + +The SIMD crate also reduces repeated byte-class combinations under Miri: +public scanner checks retain every byte in every lane at 33 bytes, and +byte-class boundaries at every other tested length. Direct architecture +checks keep their full byte/lane matrices. Range grammar enumeration keeps +short inputs and two-specification forms; fixed longer regressions remain. +Allocation-counter concurrency tests keep all four threads with fewer +updates per thread. Native runs retain the exhaustive input corpora. + +Bolero-generated properties run natively, but are ignored under Miri because +Bolero's corpus replay requires filesystem access blocked by Miri isolation. +Fixed regression seeds and deterministic scanner and parser checks still run +under Miri. The in-memory Axum example tests use a runtime without an OS I/O +driver, so they do not require Windows I/O completion port emulation. diff --git a/justfile b/justfile index b9785f6da..584a38c5f 100644 --- a/justfile +++ b/justfile @@ -25,6 +25,71 @@ test-scripts scope="": $ErrorActionPreference = "Stop" & ./scripts/tests/Pester/Run-Tests.ps1 -Path "{{ scope }}" +# Verify the package-isolated no_std contract on native AArch64 Linux. +test-http-headers-simd-no-std-arm: (_test-http-headers-simd-no-std "aarch64-unknown-linux-gnu") + +# Install the selected Rust target for baseline x86-family no_std tests. +[arg("target", long, pattern='(?:i586|i686|x86_64)-unknown-linux-gnu')] +[script("pwsh", "-NoProfile")] +setup-http-headers-simd-no-std-x86 target: anvil-toolchain-stable-install + $ErrorActionPreference = 'Stop' + $toolchainArgs = {{_anvil_stable_toolchain_args}} + if ($toolchainArgs.Count -gt 0) { + $env:RUSTUP_TOOLCHAIN = $toolchainArgs[0].Substring(1) + } + & rustup target add "{{ target }}" + exit $LASTEXITCODE + +# Verify baseline 64-bit or 32-bit x86 no_std dispatch on an x86-64 Linux host. +[arg("target", long, pattern='(?:i586|i686|x86_64)-unknown-linux-gnu')] +test-http-headers-simd-no-std-x86 target: (_test-http-headers-simd-no-std target) + +[private] +[arg("target", pattern='(?:aarch64|i586|i686|x86_64)-unknown-linux-gnu')] +[script("pwsh", "-NoProfile")] +_test-http-headers-simd-no-std target: anvil-tool-rustc-validate-prereqs + $ErrorActionPreference = 'Stop' + $target = '{{ target }}' + $hostArchitecture = if ($target -eq 'aarch64-unknown-linux-gnu') { 'Arm64' } else { 'X64' } + if (-not $IsLinux -or [System.Runtime.InteropServices.RuntimeInformation]::OSArchitecture -ne $hostArchitecture) { + throw "isolated no_std tests for $target require a native $hostArchitecture Linux host" + } + if ($target -eq 'aarch64-unknown-linux-gnu') { + $requiredTests = @( + 'dispatch::tests::no_std_neon_detection_matches_target_features: test', + 'arm::tests::neon_matches_scalar: test' + ) + } else { + $cpu = switch ($target) { + 'x86_64-unknown-linux-gnu' { 'x86-64' } + 'i686-unknown-linux-gnu' { 'pentium4' } + 'i586-unknown-linux-gnu' { 'pentium' } + } + # Encoded flags take precedence over inherited RUSTFLAGS and configured x86-64-v3. + $env:CARGO_ENCODED_RUSTFLAGS = @("-Ctarget-cpu=$cpu", '-Ctarget-feature=-ssse3,-sse4.2') -join [char]31 + $requiredTests = @( + 'dispatch::tests::x86_feature_detection_matches_compile_time_or_runtime_support: test', + 'dispatch::tests::no_std_ssse3_detection_is_stable_after_caching: test', + 'dispatch::tests::no_std_sse42_detection_is_stable_after_caching: test' + ) + if ($target -ne 'x86_64-unknown-linux-gnu') { + $requiredTests += 'dispatch::tests::no_std_sse2_detection_matches_target_features: test' + } + } + # Selecting the facade alongside this package would re-enable std. + $testArgs = @('--locked', '--package', 'http_headers_simd', '--no-default-features', '--target', $target, '--tests') + $tests = & cargo {{_anvil_stable_toolchain_args}} test @testArgs -- --list + if ($LASTEXITCODE -ne 0) { exit $LASTEXITCODE } + $tests | Write-Output + # Fail if feature unification or target selection hides either path. + foreach ($required in $requiredTests) { + if ($tests -notcontains $required) { + throw "required no_std test for $target is missing: $required" + } + } + & cargo {{_anvil_stable_toolchain_args}} test @testArgs + exit $LASTEXITCODE + # Run compile-fail tests while iterating on macro diagnostics. [arg("package", long)] [arg("filter", long)]