1use std::error::Error;
2use std::fmt::{self, Display, Formatter};
3
4use http::header::{ACCEPT, CONTENT_ENCODING, CONTENT_TYPE};
5use http::{HeaderMap, Method, Request, Response, StatusCode};
6use wip_protocol::{
7 CallOperationRequest, CallOperationResponse, FetchInterfaceRequest, FetchInterfaceResponse,
8 INTERFACE_FORMAT_V1, InterfaceDescriptor, InterfaceTarget, ObserveRequest, ObserveResponse,
9 ProtocolError, ProtocolInteraction, Target,
10};
11
12use crate::response::{self, DecodedResponse};
13use crate::{
14 BodyKind, ClientResponseError, CodecError, Endpoint, InvalidResponseKind, JSON_CONTENT_TYPE,
15 Limits, Route, json,
16};
17
18fn request(
19 endpoint: &Endpoint,
20 route: Route,
21 body: Vec<u8>,
22) -> Result<Request<Vec<u8>>, http::Error> {
23 Request::builder()
24 .method(Method::POST)
25 .uri(endpoint.route(route).as_str())
26 .header(CONTENT_TYPE, JSON_CONTENT_TYPE)
27 .header(ACCEPT, "application/json")
28 .body(body)
29}
30
31fn response(status: StatusCode, body: Vec<u8>) -> Result<Response<Vec<u8>>, http::Error> {
32 Response::builder()
33 .status(status)
34 .header(CONTENT_TYPE, JSON_CONTENT_TYPE)
35 .body(body)
36}
37
38fn require_response_content_type<B>(response: &Response<B>) -> Result<(), ClientResponseError> {
39 if !has_supported_content_coding(response.headers()) {
40 return Err(response::unsupported_content_coding(response.status()));
41 }
42 let valid = response
43 .headers()
44 .get(CONTENT_TYPE)
45 .and_then(|value| value.to_str().ok())
46 .is_some_and(is_json_content_type);
47 if valid {
48 Ok(())
49 } else if response.status().is_success() {
50 Err(ClientResponseError::invalid(
51 InvalidResponseKind::InvalidSuccessBody,
52 "missing or non-UTF-8 application/json Content-Type",
53 ))
54 } else {
55 Err(response::transport_failure(response.status()))
56 }
57}
58
59fn is_utf8_charset(value: &str) -> bool {
60 let value = value.trim();
61 value.eq_ignore_ascii_case("utf-8")
62 || value
63 .strip_prefix('"')
64 .and_then(|value| value.strip_suffix('"'))
65 .is_some_and(|value| value.eq_ignore_ascii_case("utf-8"))
66}
67
68#[must_use]
74pub fn is_json_content_type(value: &str) -> bool {
75 let mut parts = value.split(';');
76 if !parts
77 .next()
78 .is_some_and(|media| media.trim().eq_ignore_ascii_case("application/json"))
79 {
80 return false;
81 }
82 let mut saw_charset = false;
83 for parameter in parts {
84 let Some((name, value)) = parameter.trim().split_once('=') else {
85 return false;
86 };
87 if !name.trim().eq_ignore_ascii_case("charset") || !is_utf8_charset(value) || saw_charset {
88 return false;
89 }
90 saw_charset = true;
91 }
92 true
93}
94
95fn accepts_json(headers: &HeaderMap) -> bool {
96 let mut present = false;
97 let mut best_match: Option<((u8, u8), f32)> = None;
98 for value in headers.get_all(ACCEPT) {
99 present = true;
100 let Ok(value) = value.to_str() else {
101 continue;
102 };
103 for range in value.split(',') {
104 let mut parts = range.split(';');
105 let media_range = parts.next().unwrap_or_default().trim();
106 let mut quality = 1.0_f32;
107 let mut saw_quality = false;
108 let mut saw_charset = false;
109 let mut media_parameter_count = 0;
110 let mut valid = true;
111 for parameter in parts {
112 let parameter = parameter.trim();
113 let (name, value) = parameter
114 .split_once('=')
115 .map_or((parameter, None), |(name, value)| {
116 (name.trim(), Some(value))
117 });
118 if name.is_empty() {
119 valid = false;
120 break;
121 }
122 if saw_quality {
123 continue;
126 }
127 if name.eq_ignore_ascii_case("q") {
128 let Some(value) = value else {
129 valid = false;
130 break;
131 };
132 saw_quality = true;
133 match value.trim().parse::<f32>() {
134 Ok(value) if (0.0..=1.0).contains(&value) => quality = value,
135 _ => {
136 valid = false;
137 break;
138 }
139 }
140 } else {
141 let Some(value) = value else {
142 valid = false;
143 break;
144 };
145 if !name.eq_ignore_ascii_case("charset")
146 || !is_utf8_charset(value)
147 || saw_charset
148 {
149 valid = false;
150 break;
151 }
152 saw_charset = true;
153 media_parameter_count += 1;
154 }
155 }
156 let media_specificity = if media_range.eq_ignore_ascii_case("application/json") {
157 Some(2)
158 } else if media_range.eq_ignore_ascii_case("application/*") {
159 Some(1)
160 } else if media_range == "*/*" {
161 Some(0)
162 } else {
163 None
164 };
165 if valid && let Some(media_specificity) = media_specificity {
166 let specificity = (media_specificity, media_parameter_count);
167 match best_match {
168 Some((best_specificity, _)) if best_specificity > specificity => {}
169 Some((best_specificity, best_quality)) if best_specificity == specificity => {
170 best_match = Some((specificity, best_quality.max(quality)));
171 }
172 _ => best_match = Some((specificity, quality)),
173 }
174 }
175 }
176 }
177 if present {
178 best_match.is_some_and(|(_, quality)| quality > 0.0)
179 } else {
180 true
181 }
182}
183
184#[derive(Debug)]
186pub enum RequestMetadataError {
187 NotFound,
189 MethodNotAllowed,
191 UnsupportedMediaType,
193 NotAcceptable,
195}
196
197impl RequestMetadataError {
198 #[must_use]
200 pub const fn status(&self) -> StatusCode {
201 match self {
202 Self::NotFound => StatusCode::NOT_FOUND,
203 Self::MethodNotAllowed => StatusCode::METHOD_NOT_ALLOWED,
204 Self::UnsupportedMediaType => StatusCode::UNSUPPORTED_MEDIA_TYPE,
205 Self::NotAcceptable => StatusCode::NOT_ACCEPTABLE,
206 }
207 }
208
209 #[must_use]
211 pub const fn allow(&self) -> Option<&'static str> {
212 match self {
213 Self::MethodNotAllowed => Some("POST"),
214 Self::NotFound | Self::UnsupportedMediaType | Self::NotAcceptable => None,
215 }
216 }
217}
218
219impl Display for RequestMetadataError {
220 fn fmt(&self, formatter: &mut Formatter<'_>) -> fmt::Result {
221 match self {
222 Self::NotFound => formatter.write_str("request URI is not a standard WIP route"),
223 Self::MethodNotAllowed => formatter.write_str("standard WIP routes require POST"),
224 Self::UnsupportedMediaType => {
225 formatter.write_str("expected identity-coded UTF-8 application/json Content-Type")
226 }
227 Self::NotAcceptable => formatter.write_str("request does not accept application/json"),
228 }
229 }
230}
231
232impl Error for RequestMetadataError {}
233
234pub(crate) fn has_supported_content_coding(headers: &HeaderMap) -> bool {
235 headers.get_all(CONTENT_ENCODING).iter().all(|value| {
236 value.to_str().is_ok_and(|value| {
237 let mut codings = value.split(',').map(str::trim);
238 let first = codings.next();
239 first.is_some_and(|coding| coding.eq_ignore_ascii_case("identity"))
240 && codings.all(|coding| coding.eq_ignore_ascii_case("identity"))
241 })
242 })
243}
244
245pub fn validate_http_request<B>(
252 endpoint: &Endpoint,
253 route: Route,
254 value: &Request<B>,
255) -> Result<(), RequestMetadataError> {
256 let expected = endpoint.route(route);
257 if value.uri().path() != expected.path() || value.uri().query().is_some() {
258 return Err(RequestMetadataError::NotFound);
259 }
260 if value.method() != Method::POST {
261 return Err(RequestMetadataError::MethodNotAllowed);
262 }
263 let valid_content_type = value
264 .headers()
265 .get(CONTENT_TYPE)
266 .and_then(|value| value.to_str().ok())
267 .is_some_and(is_json_content_type);
268 if !valid_content_type || !has_supported_content_coding(value.headers()) {
269 return Err(RequestMetadataError::UnsupportedMediaType);
270 }
271 if !accepts_json(value.headers()) {
272 return Err(RequestMetadataError::NotAcceptable);
273 }
274 Ok(())
275}
276
277pub fn encode_observe_request(
279 endpoint: &Endpoint,
280 value: &ObserveRequest,
281 limits: Limits,
282) -> Result<Request<Vec<u8>>, EncodeHttpError> {
283 value.validate().map_err(CodecError::from)?;
284 let body = json::serialize(
285 &json::observe_request_to_json(value),
286 BodyKind::Request,
287 limits,
288 )?;
289 Ok(request(endpoint, Route::Observe, body)?)
290}
291
292pub fn decode_observe_request(body: &[u8], limits: Limits) -> Result<ObserveRequest, CodecError> {
294 let value = json::parse(body, BodyKind::Request, limits)?;
295 json::observe_request_from_json(&value)
296}
297
298pub fn encode_fetch_interface_request(
302 endpoint: &Endpoint,
303 value: &FetchInterfaceRequest,
304 limits: Limits,
305) -> Result<Request<Vec<u8>>, EncodeHttpError> {
306 value.interface.validate().map_err(CodecError::from)?;
307 let body = json::serialize(
308 &json::fetch_interface_request_to_json(value),
309 BodyKind::Request,
310 limits,
311 )?;
312 Ok(request(endpoint, Route::FetchInterface, body)?)
313}
314
315pub fn decode_fetch_interface_request(
317 body: &[u8],
318 limits: Limits,
319) -> Result<FetchInterfaceRequest, CodecError> {
320 let value = json::parse(body, BodyKind::Request, limits)?;
321 json::fetch_interface_request_from_json(&value)
322}
323
324#[derive(Debug, Clone, PartialEq, Eq)]
331pub struct CallOperationMetadata {
332 pub target: Target,
334 pub interface: InterfaceTarget,
336 pub operation: String,
338}
339
340impl CallOperationMetadata {
341 pub fn decode_request(
344 &self,
345 body: &[u8],
346 descriptor: &InterfaceDescriptor,
347 limits: Limits,
348 ) -> Result<CallOperationRequest, CodecError> {
349 let request = decode_call_operation_request(body, descriptor, limits)?;
350 if request.target != self.target
351 || request.interface != self.interface
352 || request.operation != self.operation
353 {
354 return Err(CodecError::InvalidField {
355 field: "request".into(),
356 reason: "call metadata changed after descriptor resolution".into(),
357 });
358 }
359 Ok(request)
360 }
361}
362
363pub fn decode_call_operation_metadata(
369 body: &[u8],
370 limits: Limits,
371) -> Result<CallOperationMetadata, CodecError> {
372 let value = json::parse(body, BodyKind::Request, limits)?;
373 let (target, interface, operation) = json::call_request_metadata_from_json(&value)?;
374 Ok(CallOperationMetadata {
375 target,
376 interface,
377 operation,
378 })
379}
380
381pub fn encode_call_operation_request(
383 endpoint: &Endpoint,
384 value: &CallOperationRequest,
385 descriptor: &InterfaceDescriptor,
386 limits: Limits,
387) -> Result<Request<Vec<u8>>, EncodeHttpError> {
388 let value = json::call_request_to_json(value, descriptor)?;
389 let body = json::serialize(&value, BodyKind::Request, limits)?;
390 Ok(request(endpoint, Route::CallOperation, body)?)
391}
392
393pub fn decode_call_operation_request(
395 body: &[u8],
396 descriptor: &InterfaceDescriptor,
397 limits: Limits,
398) -> Result<CallOperationRequest, CodecError> {
399 let value = json::parse(body, BodyKind::Request, limits)?;
400 json::call_request_from_json(&value, descriptor)
401}
402
403pub fn encode_observe_response(
405 request_value: &ObserveRequest,
406 value: &ObserveResponse,
407 limits: Limits,
408) -> Result<Response<Vec<u8>>, EncodeHttpError> {
409 request_value
410 .validate_response(value)
411 .map_err(CodecError::from)?;
412 let body = json::serialize(
413 &json::observe_response_to_json(value),
414 BodyKind::Response,
415 limits,
416 )?;
417 Ok(response(StatusCode::OK, body)?)
418}
419
420pub fn decode_observe_response<B: AsRef<[u8]>>(
422 request_value: &ObserveRequest,
423 value: &Response<B>,
424 limits: Limits,
425) -> Result<DecodedResponse<ObserveResponse>, ClientResponseError> {
426 require_response_content_type(value)?;
427 response::decode(
428 value.status(),
429 value.body().as_ref(),
430 ProtocolInteraction::Observe,
431 limits,
432 |json| {
433 let observation = json::observe_response_from_json(json)?;
434 request_value.validate_response(&observation)?;
435 Ok(observation)
436 },
437 )
438}
439
440pub fn encode_fetch_interface_response(
442 request_value: &FetchInterfaceRequest,
443 value: &FetchInterfaceResponse,
444 limits: Limits,
445) -> Result<Response<Vec<u8>>, EncodeHttpError> {
446 value
447 .validate_for(request_value)
448 .map_err(CodecError::from)?;
449 let body = json::serialize(
450 &json::fetch_interface_response_to_json(value),
451 BodyKind::Response,
452 limits,
453 )?;
454 Ok(response(StatusCode::OK, body)?)
455}
456
457pub fn decode_fetch_interface_response<B: AsRef<[u8]>>(
459 request_value: &FetchInterfaceRequest,
460 value: &Response<B>,
461 limits: Limits,
462) -> Result<DecodedResponse<FetchInterfaceResponse>, ClientResponseError> {
463 require_response_content_type(value)?;
464 let decoded = response::decode(
465 value.status(),
466 value.body().as_ref(),
467 ProtocolInteraction::FetchInterface,
468 limits,
469 |json| {
470 let interface = json::fetch_interface_response_from_json(json)?;
471 interface.validate_for(request_value)?;
472 Ok(interface)
473 },
474 )?;
475 if let DecodedResponse::Success(success) = &decoded
476 && success.descriptor.format != INTERFACE_FORMAT_V1
477 {
478 return Err(ClientResponseError::UnsupportedDescriptorFormat {
479 format: success.descriptor.format.clone(),
480 });
481 }
482 Ok(decoded)
483}
484
485pub fn encode_call_operation_response(
487 request_value: &CallOperationRequest,
488 descriptor: &InterfaceDescriptor,
489 value: &CallOperationResponse,
490 limits: Limits,
491) -> Result<Response<Vec<u8>>, EncodeHttpError> {
492 let json = json::call_response_to_json(value, descriptor, &request_value.operation)?;
493 let body = json::serialize(&json, BodyKind::Response, limits)?;
494 Ok(response(StatusCode::OK, body)?)
495}
496
497pub fn decode_call_operation_response<B: AsRef<[u8]>>(
499 request_value: &CallOperationRequest,
500 descriptor: &InterfaceDescriptor,
501 value: &Response<B>,
502 limits: Limits,
503) -> Result<DecodedResponse<CallOperationResponse>, ClientResponseError> {
504 require_response_content_type(value)?;
505 response::decode(
506 value.status(),
507 value.body().as_ref(),
508 ProtocolInteraction::CallOperation,
509 limits,
510 |json| json::call_response_from_json(json, descriptor, &request_value.operation),
511 )
512}
513
514fn validate_protocol_error(
515 interaction: ProtocolInteraction,
516 error: &ProtocolError,
517) -> Result<(), CodecError> {
518 if !error.code.is_allowed_for(interaction) {
519 return Err(CodecError::InvalidField {
520 field: "error.code".into(),
521 reason: format!(
522 "code `{}` is not allowed for {interaction:?}",
523 response::error_code_name(error.code)
524 ),
525 });
526 }
527 Ok(())
528}
529
530fn encode_protocol_error_with_status(
531 interaction: ProtocolInteraction,
532 error: &ProtocolError,
533 status: StatusCode,
534 limits: Limits,
535) -> Result<Response<Vec<u8>>, EncodeHttpError> {
536 validate_protocol_error(interaction, error)?;
537 let body = json::serialize(&response::error_to_json(error), BodyKind::Response, limits)?;
538 Ok(response(status, body)?)
539}
540
541pub fn encode_protocol_error_response(
549 interaction: ProtocolInteraction,
550 error: &ProtocolError,
551 limits: Limits,
552) -> Result<Response<Vec<u8>>, EncodeHttpError> {
553 encode_protocol_error_with_status(
554 interaction,
555 error,
556 response::status_for_error(error.code),
557 limits,
558 )
559}
560
561#[derive(Debug)]
563pub enum EncodeHttpError {
564 Codec(CodecError),
566 Http(http::Error),
568}
569
570impl std::fmt::Display for EncodeHttpError {
571 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
572 match self {
573 Self::Codec(error) => error.fmt(formatter),
574 Self::Http(error) => error.fmt(formatter),
575 }
576 }
577}
578
579impl std::error::Error for EncodeHttpError {
580 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
581 match self {
582 Self::Codec(error) => Some(error),
583 Self::Http(error) => Some(error),
584 }
585 }
586}
587
588impl From<CodecError> for EncodeHttpError {
589 fn from(value: CodecError) -> Self {
590 Self::Codec(value)
591 }
592}
593
594impl From<http::Error> for EncodeHttpError {
595 fn from(value: http::Error) -> Self {
596 Self::Http(value)
597 }
598}