1use proc_macro::TokenStream;
4use quote::{ToTokens, quote};
5use serde_json::{Map, Number, Value};
6use syn::{
7 Attribute, Error, FnArg, Ident, ImplItem, ItemImpl, Lit, LitStr, Result, Token, Type, braced,
8 bracketed,
9 ext::IdentExt,
10 parse::{Parse, ParseStream},
11 parse_macro_input,
12 punctuated::Punctuated,
13 spanned::Spanned,
14};
15
16#[proc_macro_attribute]
24pub fn endpoint(arguments: TokenStream, input: TokenStream) -> TokenStream {
25 if !arguments.is_empty() {
26 return Error::new(
27 proc_macro2::Span::call_site(),
28 "endpoint does not accept arguments",
29 )
30 .into_compile_error()
31 .into();
32 }
33
34 let implementation = parse_macro_input!(input as ItemImpl);
35 expand_endpoint(implementation)
36 .unwrap_or_else(Error::into_compile_error)
37 .into()
38}
39
40#[proc_macro]
45pub fn openapi_operation(input: TokenStream) -> TokenStream {
46 let operation = parse_macro_input!(input as OpenApiObject);
47 match operation_literal(&operation) {
48 Ok(operation) => quote!(#operation).into(),
49 Err(error) => error.into_compile_error().into(),
50 }
51}
52
53fn expand_endpoint(mut implementation: ItemImpl) -> Result<proc_macro2::TokenStream> {
54 if implementation.trait_.is_some() {
55 return Err(Error::new_spanned(
56 implementation.impl_token,
57 "endpoint can only be applied to an inherent impl block",
58 ));
59 }
60 if !implementation.generics.params.is_empty() {
61 return Err(Error::new_spanned(
62 &implementation.generics,
63 "endpoint does not support generic impl blocks",
64 ));
65 }
66
67 let provider = implementation.self_ty.clone();
68 let provider_middlewares = take_provider_middlewares(&mut implementation.attrs)?;
69 let mut routes = Vec::new();
70 for item in &mut implementation.items {
71 let ImplItem::Fn(method) = item else {
72 continue;
73 };
74 let metadata = take_handler_metadata(&mut method.attrs)?;
75 if let Some(mut route) = metadata.route {
76 route.openapi = metadata.openapi;
77 routes.push(Handler {
78 route,
79 middlewares: provider_middlewares
80 .iter()
81 .cloned()
82 .chain(metadata.middlewares)
83 .collect(),
84 method: method.sig.ident.clone(),
85 arguments: handler_arguments(method)?,
86 });
87 } else if !metadata.middlewares.is_empty() || metadata.openapi.is_some() {
88 return Err(Error::new_spanned(
89 &method.sig.ident,
90 "endpoint metadata can only be attached to an HTTP handler",
91 ));
92 }
93 }
94 if routes.is_empty() {
95 return Err(Error::new_spanned(
96 &implementation.self_ty,
97 "endpoint impl must declare at least one HTTP handler attribute",
98 ));
99 }
100
101 let const_routes = routes.iter().map(endpoint_route);
102 let implementation_routes = routes.iter().map(endpoint_route);
103 let dispatch_arms = routes.iter().map(dispatch_arm);
104
105 Ok(quote! {
106 #implementation
107
108 const _: () = {
109 const ROUTES: &[::lenso_capability_http_endpoint::EndpointRoute] = &[
110 #(#const_routes)*
111 ];
112 ::lenso_capability_http_endpoint::__private::validate_endpoint_routes(ROUTES);
113 };
114
115 impl ::lenso_capability_http_endpoint::HttpEndpoint for #provider {
116 const ROUTES: &'static [::lenso_capability_http_endpoint::EndpointRoute] = &[
117 #(#implementation_routes)*
118 ];
119
120 fn dispatch(
121 &self,
122 context: ::lenso_capability_http_endpoint::__private::InvocationContext,
123 request: ::lenso_capability_http_endpoint::HandleRequest,
124 ) -> ::lenso_capability_http_endpoint::EndpointFuture {
125 let provider = self.clone();
126 Box::pin(async move {
127 let route_id = request.route_id.clone();
128 match route_id.as_str() {
129 #(#dispatch_arms,)*
130 _ => Ok(Err(::lenso_capability_http_endpoint::HandleError::Rejected)),
131 }
132 })
133 }
134 }
135 })
136}
137
138fn endpoint_route(handler: &Handler) -> proc_macro2::TokenStream {
139 let route_id = &handler.route.id;
140 let method = &handler.route.method;
141 let path = &handler.route.path;
142 let openapi = handler
143 .route
144 .openapi
145 .as_ref()
146 .map(|operation| quote!(.with_openapi(#operation)));
147 quote! {
148 ::lenso_capability_http_endpoint::EndpointRoute::new(
149 #route_id,
150 #method,
151 #path,
152 ) #openapi,
153 }
154}
155
156fn take_provider_middlewares(attributes: &mut Vec<Attribute>) -> Result<Vec<Ident>> {
157 let mut middlewares = Vec::new();
158 let mut retained = Vec::with_capacity(attributes.len());
159 for attribute in attributes.drain(..) {
160 if attribute.path().is_ident("middleware") {
161 let arguments =
162 attribute.parse_args_with(Punctuated::<Ident, Token![,]>::parse_terminated)?;
163 if arguments.is_empty() {
164 return Err(Error::new_spanned(
165 attribute,
166 "middleware requires at least one provider method",
167 ));
168 }
169 middlewares.extend(arguments);
170 } else {
171 retained.push(attribute);
172 }
173 }
174 *attributes = retained;
175 Ok(middlewares)
176}
177
178fn take_handler_metadata(attributes: &mut Vec<Attribute>) -> Result<HandlerMetadata> {
179 let mut route = None;
180 let mut middlewares = Vec::new();
181 let mut openapi = None;
182 let mut retained = Vec::with_capacity(attributes.len());
183 for attribute in attributes.drain(..) {
184 if attribute.path().is_ident("middleware") {
185 let arguments =
186 attribute.parse_args_with(Punctuated::<Ident, Token![,]>::parse_terminated)?;
187 if arguments.is_empty() {
188 return Err(Error::new_spanned(
189 attribute,
190 "middleware requires at least one provider method",
191 ));
192 }
193 middlewares.extend(arguments);
194 continue;
195 }
196 if attribute.path().is_ident("openapi") {
197 if openapi.is_some() {
198 return Err(Error::new_spanned(
199 attribute,
200 "an endpoint handler may declare only one OpenAPI Operation Object",
201 ));
202 }
203 let operation = attribute.parse_args::<OpenApiOperation>()?.into_literal()?;
204 openapi = Some(operation);
205 continue;
206 }
207 let Some(http_method) = http_method(&attribute) else {
208 retained.push(attribute);
209 continue;
210 };
211 if route.is_some() {
212 return Err(Error::new_spanned(
213 attribute,
214 "an endpoint handler may declare only one HTTP method",
215 ));
216 }
217 let arguments = attribute.parse_args::<RouteArguments>()?;
218 route = Some(Route {
219 method: LitStr::new(http_method, attribute.path().span()),
220 id: arguments.route_id,
221 path: arguments.path,
222 openapi: None,
223 });
224 }
225 *attributes = retained;
226 Ok(HandlerMetadata {
227 route,
228 middlewares,
229 openapi,
230 })
231}
232
233enum OpenApiOperation {
234 Json(LitStr),
235 Object(OpenApiObject),
236}
237
238impl OpenApiOperation {
239 fn into_literal(self) -> Result<LitStr> {
240 match self {
241 Self::Json(operation) => {
242 validate_openapi_operation(&operation)?;
243 Ok(operation)
244 }
245 Self::Object(operation) => operation_literal(&operation),
246 }
247 }
248}
249
250impl Parse for OpenApiOperation {
251 fn parse(input: ParseStream<'_>) -> Result<Self> {
252 if input.peek(LitStr) {
253 return input.parse().map(Self::Json);
254 }
255 input.parse().map(Self::Object)
256 }
257}
258
259struct OpenApiObject {
260 value: Map<String, Value>,
261 span: proc_macro2::Span,
262}
263
264impl Parse for OpenApiObject {
265 fn parse(input: ParseStream<'_>) -> Result<Self> {
266 let content;
267 let brace = braced!(content in input);
268 Ok(Self {
269 value: parse_object_entries(&content)?,
270 span: brace.span.join(),
271 })
272 }
273}
274
275fn parse_object_entries(input: ParseStream<'_>) -> Result<Map<String, Value>> {
276 let mut object = Map::new();
277 while !input.is_empty() {
278 let (key, span) = parse_object_key(input)?;
279 input.parse::<Token![:]>()?;
280 let value = parse_json_value(input)?;
281 if object.insert(key.clone(), value).is_some() {
282 return Err(Error::new(
283 span,
284 format!("duplicate OpenAPI object key `{key}`"),
285 ));
286 }
287 if input.is_empty() {
288 break;
289 }
290 input.parse::<Token![,]>()?;
291 }
292 Ok(object)
293}
294
295fn parse_object_key(input: ParseStream<'_>) -> Result<(String, proc_macro2::Span)> {
296 if input.peek(LitStr) {
297 let key = input.parse::<LitStr>()?;
298 return Ok((key.value(), key.span()));
299 }
300 let key = Ident::parse_any(input)?;
301 Ok((key.unraw().to_string(), key.span()))
302}
303
304fn parse_json_value(input: ParseStream<'_>) -> Result<Value> {
305 if input.peek(syn::token::Brace) {
306 let content;
307 braced!(content in input);
308 return parse_object_entries(&content).map(Value::Object);
309 }
310 if input.peek(syn::token::Bracket) {
311 let content;
312 bracketed!(content in input);
313 let values = Punctuated::<JsonValue, Token![,]>::parse_terminated(&content)?;
314 return Ok(Value::Array(
315 values.into_iter().map(|value| value.0).collect(),
316 ));
317 }
318 if input.peek(Token![-]) {
319 input.parse::<Token![-]>()?;
320 let literal = input.parse::<Lit>()?;
321 return parse_number_literal(&literal, true);
322 }
323 if input.peek(syn::LitBool) {
324 let literal = input.parse::<Lit>()?;
325 return parse_literal(&literal);
326 }
327 if input.peek(Ident::peek_any) {
328 let ident = Ident::parse_any(input)?;
329 if ident == "null" {
330 return Ok(Value::Null);
331 }
332 return Err(Error::new(ident.span(), "expected a JSON value"));
333 }
334 let literal = input.parse::<Lit>()?;
335 parse_literal(&literal)
336}
337
338struct JsonValue(Value);
339
340impl Parse for JsonValue {
341 fn parse(input: ParseStream<'_>) -> Result<Self> {
342 parse_json_value(input).map(Self)
343 }
344}
345
346fn parse_literal(literal: &Lit) -> Result<Value> {
347 match literal {
348 Lit::Str(value) => Ok(Value::String(value.value())),
349 Lit::Bool(value) => Ok(Value::Bool(value.value)),
350 Lit::Int(_) | Lit::Float(_) => parse_number_literal(literal, false),
351 _ => Err(Error::new(
352 literal.span(),
353 "OpenAPI metadata supports JSON string, number, boolean, null, array, and object values",
354 )),
355 }
356}
357
358fn parse_number_literal(literal: &Lit, negative: bool) -> Result<Value> {
359 let mut source = literal.to_token_stream().to_string().replace(' ', "");
360 if negative {
361 source.insert(0, '-');
362 }
363 let number = source.parse::<Number>().map_err(|_| {
364 Error::new(
365 literal.span(),
366 "OpenAPI numeric values must be unsuffixed JSON numbers",
367 )
368 })?;
369 Ok(Value::Number(number))
370}
371
372fn operation_literal(operation: &OpenApiObject) -> Result<LitStr> {
373 if operation.value.contains_key("operationId") {
374 return Err(Error::new(
375 operation.span,
376 "OpenAPI operationId is generated from the stable route ID",
377 ));
378 }
379 let json = serde_json::to_string(&operation.value).map_err(|error| {
380 Error::new(
381 operation.span,
382 format!("could not encode OpenAPI Operation Object: {error}"),
383 )
384 })?;
385 Ok(LitStr::new(&json, operation.span))
386}
387
388fn validate_openapi_operation(operation: &LitStr) -> Result<()> {
389 let value = serde_json::from_str::<serde_json::Value>(&operation.value()).map_err(|error| {
390 Error::new(
391 operation.span(),
392 format!("OpenAPI Operation Object is not valid JSON: {error}"),
393 )
394 })?;
395 let Some(object) = value.as_object() else {
396 return Err(Error::new(
397 operation.span(),
398 "OpenAPI operation metadata must be a JSON object",
399 ));
400 };
401 if object.contains_key("operationId") {
402 return Err(Error::new(
403 operation.span(),
404 "OpenAPI operationId is generated from the stable route ID",
405 ));
406 }
407 Ok(())
408}
409
410fn handler_arguments(method: &syn::ImplItemFn) -> Result<Vec<HandlerArgument>> {
411 let mut arguments = Vec::new();
412 let mut context_count = 0;
413 let mut request_count = 0;
414 for (index, input) in method.sig.inputs.iter().enumerate() {
415 let FnArg::Typed(argument) = input else {
416 continue;
417 };
418 let ty = (*argument.ty).clone();
419 let kind = match final_type_ident(&ty).map(Ident::to_string).as_deref() {
420 Some("InvocationContext") => {
421 context_count += 1;
422 ArgumentKind::Context
423 }
424 Some("HandleRequest") => {
425 request_count += 1;
426 ArgumentKind::Request
427 }
428 _ => ArgumentKind::Extractor(Ident::new(
429 &format!("__lenso_extracted_{index}"),
430 argument.span(),
431 )),
432 };
433 arguments.push(HandlerArgument { ty, kind });
434 }
435 if context_count > 1 || request_count > 1 {
436 return Err(Error::new_spanned(
437 &method.sig.inputs,
438 "an endpoint handler accepts at most one InvocationContext and one HandleRequest",
439 ));
440 }
441 Ok(arguments)
442}
443
444fn final_type_ident(ty: &Type) -> Option<&Ident> {
445 let Type::Path(path) = ty else {
446 return None;
447 };
448 path.path.segments.last().map(|segment| &segment.ident)
449}
450
451fn dispatch_arm(handler: &Handler) -> proc_macro2::TokenStream {
452 let route_id = &handler.route.id;
453 let method = &handler.method;
454 let middleware_steps = handler.middlewares.iter().map(|middleware| {
455 quote! {
456 let (context, request) = match provider.#middleware(context, request).await {
457 Ok(outcome) => match outcome.into_result() {
458 Ok(next) => next,
459 Err(response) => return Ok(Ok(response)),
460 },
461 Err(::lenso_capability_http_endpoint::EndpointHandleInvocationError::Domain(
462 error,
463 )) => return Ok(Err(error)),
464 Err(::lenso_capability_http_endpoint::EndpointHandleInvocationError::Runtime(
465 error,
466 )) => return Err(error),
467 };
468 }
469 });
470 let extractor_steps = handler.arguments.iter().filter_map(|argument| {
471 let ArgumentKind::Extractor(binding) = &argument.kind else {
472 return None;
473 };
474 let ty = &argument.ty;
475 Some(quote! {
476 let #binding: #ty = match <#ty as
477 ::lenso_capability_http_endpoint::__private::FromRequest<Self>
478 >::from_request(&provider, &mut context, &request).await {
479 Ok(value) => value,
480 Err(::lenso_capability_http_endpoint::__private::ExtractorRejection::Response(
481 response,
482 )) => return Ok(Ok(response)),
483 Err(::lenso_capability_http_endpoint::__private::ExtractorRejection::Invocation(
484 ::lenso_capability_http_endpoint::EndpointHandleInvocationError::Domain(
485 error,
486 ),
487 )) => return Ok(Err(error)),
488 Err(::lenso_capability_http_endpoint::__private::ExtractorRejection::Invocation(
489 ::lenso_capability_http_endpoint::EndpointHandleInvocationError::Runtime(
490 error,
491 ),
492 )) => return Err(error),
493 };
494 })
495 });
496 let mutable_context = handler
497 .arguments
498 .iter()
499 .any(|argument| matches!(&argument.kind, ArgumentKind::Extractor(_)))
500 .then(|| quote!(let mut context = context;));
501 let arguments = handler
502 .arguments
503 .iter()
504 .map(|argument| match &argument.kind {
505 ArgumentKind::Context => quote!(context),
506 ArgumentKind::Request => quote!(request),
507 ArgumentKind::Extractor(binding) => quote!(#binding),
508 });
509
510 quote! {
511 #route_id => {
512 #(#middleware_steps)*
513 #mutable_context
514 #(#extractor_steps)*
515 let _ = (&context, &request);
516 match provider.#method(#(#arguments),*).await {
517 Ok(response) => Ok(Ok(response)),
518 Err(::lenso_capability_http_endpoint::EndpointHandleInvocationError::Domain(
519 error,
520 )) => Ok(Err(error)),
521 Err(::lenso_capability_http_endpoint::EndpointHandleInvocationError::Runtime(
522 error,
523 )) => Err(error),
524 }
525 }
526 }
527}
528
529fn http_method(attribute: &Attribute) -> Option<&'static str> {
530 let path = attribute.path();
531 [
532 ("get", "GET"),
533 ("post", "POST"),
534 ("put", "PUT"),
535 ("patch", "PATCH"),
536 ("delete", "DELETE"),
537 ("head", "HEAD"),
538 ("options", "OPTIONS"),
539 ]
540 .into_iter()
541 .find_map(|(attribute, method)| path.is_ident(attribute).then_some(method))
542}
543
544struct RouteArguments {
545 route_id: LitStr,
546 path: LitStr,
547}
548
549impl Parse for RouteArguments {
550 fn parse(input: ParseStream<'_>) -> Result<Self> {
551 let arguments = Punctuated::<LitStr, Token![,]>::parse_terminated(input)?;
552 if arguments.len() != 2 {
553 return Err(Error::new(
554 input.span(),
555 "HTTP handler attributes require a route ID and path",
556 ));
557 }
558 let mut arguments = arguments.into_iter();
559 Ok(Self {
560 route_id: arguments.next().expect("length was checked"),
561 path: arguments.next().expect("length was checked"),
562 })
563 }
564}
565
566struct Route {
567 method: LitStr,
568 id: LitStr,
569 path: LitStr,
570 openapi: Option<LitStr>,
571}
572
573struct HandlerMetadata {
574 route: Option<Route>,
575 middlewares: Vec<Ident>,
576 openapi: Option<LitStr>,
577}
578
579struct Handler {
580 route: Route,
581 middlewares: Vec<Ident>,
582 method: Ident,
583 arguments: Vec<HandlerArgument>,
584}
585
586struct HandlerArgument {
587 ty: Type,
588 kind: ArgumentKind,
589}
590
591enum ArgumentKind {
592 Context,
593 Request,
594 Extractor(Ident),
595}
596
597#[cfg(test)]
598mod tests {
599 use super::expand_endpoint;
600 use syn::parse_quote;
601
602 #[test]
603 fn expands_handler_attributes_into_the_static_route_table() {
604 let expanded = expand_endpoint(parse_quote! {
605 impl OrdersHttp {
606 #[get("orders.read", "/orders/{order_id}")]
607 #[openapi({
608 summary: "Read an order",
609 tags: ["orders"],
610 responses: {
611 "200": { description: "Order" }
612 }
613 })]
614 async fn read(&self) {}
615 }
616 })
617 .unwrap()
618 .to_string();
619
620 assert!(expanded.contains("HttpEndpoint"));
621 assert!(expanded.contains("orders.read"));
622 assert!(expanded.contains("GET"));
623 assert!(expanded.contains("/orders/{order_id}"));
624 assert!(expanded.contains("with_openapi"));
625 assert!(!expanded.contains("# [get"));
626 assert!(!expanded.contains("# [openapi"));
627 assert!(expanded.contains("Read an order"));
628 assert!(expanded.contains(r#"\"200\""#));
629 }
630
631 #[test]
632 fn accepts_legacy_json_string_metadata() {
633 let expanded = expand_endpoint(parse_quote! {
634 impl OrdersHttp {
635 #[get("orders.read", "/orders/{order_id}")]
636 #[openapi(r#"{"summary":"Read an order"}"#)]
637 async fn read(&self) {}
638 }
639 })
640 .unwrap()
641 .to_string();
642
643 assert!(expanded.contains("Read an order"));
644 }
645
646 #[test]
647 fn rejects_an_operation_id_in_structured_metadata() {
648 let error = expand_endpoint(parse_quote! {
649 impl OrdersHttp {
650 #[get("orders.read", "/orders/{order_id}")]
651 #[openapi({ operationId: "another.id" })]
652 async fn read(&self) {}
653 }
654 })
655 .unwrap_err();
656
657 assert!(
658 error
659 .to_string()
660 .contains("generated from the stable route ID")
661 );
662 }
663
664 #[test]
665 fn rejects_duplicate_structured_metadata_keys() {
666 let error = expand_endpoint(parse_quote! {
667 impl OrdersHttp {
668 #[get("orders.read", "/orders/{order_id}")]
669 #[openapi({ summary: "Read", summary: "Read again" })]
670 async fn read(&self) {}
671 }
672 })
673 .unwrap_err();
674
675 assert!(error.to_string().contains("duplicate OpenAPI object key"));
676 }
677
678 #[test]
679 fn rejects_an_openapi_operation_id_that_can_drift_from_the_route_id() {
680 let error = expand_endpoint(parse_quote! {
681 impl OrdersHttp {
682 #[get("orders.read", "/orders/{order_id}")]
683 #[openapi(r#"{"operationId":"another.id"}"#)]
684 async fn read(&self) {}
685 }
686 })
687 .unwrap_err();
688
689 assert!(
690 error
691 .to_string()
692 .contains("generated from the stable route ID")
693 );
694 }
695
696 #[test]
697 fn rejects_invalid_openapi_json() {
698 let error = expand_endpoint(parse_quote! {
699 impl OrdersHttp {
700 #[get("orders.read", "/orders/{order_id}")]
701 #[openapi("not-json")]
702 async fn read(&self) {}
703 }
704 })
705 .unwrap_err();
706
707 assert!(error.to_string().contains("not valid JSON"));
708 }
709
710 #[test]
711 fn expands_middleware_and_typed_extractors_before_the_handler() {
712 let expanded = expand_endpoint(parse_quote! {
713 impl OrdersHttp {
714 #[middleware(authenticate)]
715 #[get("orders.read", "/orders/{order_id}")]
716 async fn read(
717 &self,
718 context: InvocationContext,
719 Path(path): Path<OrderPath>,
720 ) {}
721 }
722 })
723 .unwrap()
724 .to_string();
725
726 assert!(expanded.contains("provider . authenticate"));
727 assert!(expanded.contains("into_result"));
728 assert!(expanded.contains("FromRequest"));
729 assert!(expanded.contains("provider . read (context , __lenso_extracted_2)"));
730 assert!(!expanded.contains("# [middleware"));
731 }
732
733 #[test]
734 fn applies_provider_middleware_before_route_middleware() {
735 let expanded = expand_endpoint(parse_quote! {
736 #[middleware(trace_all)]
737 impl OrdersHttp {
738 #[middleware(authorize_read)]
739 #[get("orders.read", "/orders/{order_id}")]
740 async fn read(&self) {}
741
742 #[get("orders.list", "/orders")]
743 async fn list(&self) {}
744 }
745 })
746 .unwrap()
747 .to_string();
748
749 assert_eq!(expanded.matches("provider . trace_all").count(), 2);
750 assert_eq!(expanded.matches("provider . authorize_read").count(), 1);
751 assert!(
752 expanded.find("provider . trace_all").unwrap()
753 < expanded.find("provider . authorize_read").unwrap()
754 );
755 assert!(!expanded.contains("# [middleware"));
756 }
757
758 #[test]
759 fn rejects_handlers_with_multiple_http_methods() {
760 let error = expand_endpoint(parse_quote! {
761 impl OrdersHttp {
762 #[get("orders.read", "/orders/{order_id}")]
763 #[post("orders.read", "/orders/{order_id}")]
764 async fn read(&self) {}
765 }
766 })
767 .unwrap_err();
768
769 assert!(error.to_string().contains("only one HTTP method"));
770 }
771
772 #[test]
773 fn rejects_impls_without_handlers() {
774 let error = expand_endpoint(parse_quote! {
775 impl OrdersHttp {
776 fn helper(&self) {}
777 }
778 })
779 .unwrap_err();
780
781 assert!(error.to_string().contains("at least one HTTP handler"));
782 }
783}