File
Blob: src/rust/kj/tests/lib.rs
| 1 | use std::pin::Pin; |
| 2 | |
| 3 | use kj::Result; |
| 4 | use kj::http::ConnectResponse; |
| 5 | use kj::http::ConnectSettings; |
| 6 | use kj::http::CustomHeaderId; |
| 7 | use kj::http::CxxService; |
| 8 | use kj::http::DynHttpService; |
| 9 | use kj::http::HeadersRef; |
| 10 | use kj::http::Method; |
| 11 | use kj::http::Service; |
| 12 | use kj::http::ServiceResponse; |
| 13 | use kj::io::AsyncInputStream; |
| 14 | use kj::io::AsyncIoStream; |
| 15 | use kj_rs::KjMaybe; |
| 16 | use kj_rs::KjOwn; |
| 17 | |
| 18 | #[cxx::bridge(namespace = "kj::rust::tests")] |
| 19 | pub mod ffi { |
| 20 | #[namespace = "kj::rust"] |
| 21 | unsafe extern "C++" { |
| 22 | include!("workerd/rust/kj/ffi.h"); |
| 23 | type HttpService = kj::http::ffi::HttpService; |
| 24 | type HttpHeaders = kj::http::ffi::HttpHeaders; |
| 25 | type HttpHeaderId = kj::http::ffi::HttpHeaderId; |
| 26 | } |
| 27 | |
| 28 | #[namespace = "kj::rust"] |
| 29 | extern "Rust" { |
| 30 | type DynHttpService = kj::http::DynHttpService; |
| 31 | } |
| 32 | |
| 33 | extern "Rust" { |
| 34 | type ProxyHttpService; |
| 35 | |
| 36 | #[expect(clippy::unnecessary_box_returns)] |
| 37 | fn new_proxy_http_service(service: KjOwn<HttpService>) -> Box<DynHttpService>; |
| 38 | |
| 39 | /// Look up a header value by HttpHeaderId, returning the value if present. |
| 40 | /// This exercises the C++ -> Rust -> C++ round-trip for HttpHeaderId. |
| 41 | unsafe fn get_header_value_via_id<'a>( |
| 42 | headers: &'a HttpHeaders, |
| 43 | id: &HttpHeaderId, |
| 44 | ) -> KjMaybe<&'a [u8]>; |
| 45 | |
| 46 | /// Receive an array of HttpHeaderIdpointers, convert to &[HttpHeaderIdRef] via |
| 47 | /// from_ptr_slice, look up each header, and assert all are present. |
| 48 | /// This exercises passing a kj::ArrayPtr<const kj::HttpHeaderId> into Rust. |
| 49 | unsafe fn assert_header_ids_present(headers: &HttpHeaders, ids: &[*const HttpHeaderId]); |
| 50 | } |
| 51 | } |
| 52 | |
| 53 | struct ProxyHttpService { |
| 54 | target: CxxService<'static>, |
| 55 | } |
| 56 | |
| 57 | #[async_trait::async_trait(?Send)] |
| 58 | impl Service for ProxyHttpService { |
| 59 | async fn request<'a>( |
| 60 | &'a mut self, |
| 61 | method: Method, |
| 62 | url: &'a [u8], |
| 63 | headers: HeadersRef<'a>, |
| 64 | request_body: Pin<&'a mut AsyncInputStream>, |
| 65 | response: ServiceResponse<'a>, |
| 66 | ) -> Result<()> { |
| 67 | self.target |
| 68 | .request(method, url, headers, request_body, response) |
| 69 | .await?; |
| 70 | Ok(()) |
| 71 | } |
| 72 | |
| 73 | fn connect<'a, 'b>( |
| 74 | &'a mut self, |
| 75 | host: &'a [u8], |
| 76 | headers: HeadersRef<'a>, |
| 77 | connection: Pin<&'a mut AsyncIoStream>, |
| 78 | response: ConnectResponse<'a>, |
| 79 | settings: ConnectSettings<'a>, |
| 80 | ) -> ::core::pin::Pin<Box<dyn ::core::future::Future<Output = Result<()>> + 'b>> |
| 81 | where |
| 82 | 'a: 'b, |
| 83 | Self: 'b, |
| 84 | { |
| 85 | Box::pin( |
| 86 | self.target |
| 87 | .connect(host, headers, connection, response, settings), |
| 88 | ) |
| 89 | } |
| 90 | } |
| 91 | |
| 92 | #[expect(clippy::unnecessary_box_returns)] |
| 93 | fn new_proxy_http_service(service: KjOwn<ffi::HttpService>) -> Box<DynHttpService> { |
| 94 | ProxyHttpService { |
| 95 | target: service.into(), |
| 96 | } |
| 97 | .into_ffi() |
| 98 | } |
| 99 | |
| 100 | fn get_header_value_via_id<'a>( |
| 101 | headers: &'a ffi::HttpHeaders, |
| 102 | id: &ffi::HttpHeaderId, |
| 103 | ) -> KjMaybe<&'a [u8]> { |
| 104 | // SAFETY: headers is a valid HttpHeaders ref and id is a valid HttpHeaderId from C++. |
| 105 | unsafe { kj::http::ffi::get_header_by_id(headers, id) } |
| 106 | } |
| 107 | |
| 108 | /// # Safety |
| 109 | /// |
| 110 | /// Each pointer in `ids` must be non-null and point to a valid, live `HttpHeaderId`. |
| 111 | unsafe fn assert_header_ids_present(headers: &ffi::HttpHeaders, ids: &[*const ffi::HttpHeaderId]) { |
| 112 | let headers_ref = HeadersRef::from(headers); |
| 113 | // SAFETY: All pointers in ids are valid, as guaranteed by the unsafe fn contract. |
| 114 | let id_refs = unsafe { CustomHeaderId::from_ptr_slice(ids) }; |
| 115 | for (i, &id_ref) in id_refs.iter().enumerate() { |
| 116 | let value = headers_ref.get_by_id(id_ref); |
| 117 | assert!( |
| 118 | value.is_some(), |
| 119 | "expected header at index {i} to be present, but got None" |
| 120 | ); |
| 121 | } |
| 122 | } |