rtc_interceptor_derive/lib.rs
1#![warn(missing_docs)]
2//! Derive macros for RTC Interceptor trait.
3//!
4//! This crate provides two macros that work together:
5//!
6//! - `#[derive(Interceptor)]` - Marks a struct as an interceptor and identifies the next field
7//! - `#[interceptor]` - Attribute macro for impl blocks to generate trait implementations
8//!
9//! # Design Pattern
10//!
11//! The examples below are illustrative rather than compiled: the macros only expand to something
12//! meaningful in the presence of the `Interceptor` trait from
13//! [`rtc-interceptor`](https://docs.rs/rtc-interceptor), which depends on *this* crate — so a
14//! doctest here cannot import it. They are exercised for real by `rtc-interceptor`'s own
15//! documentation and tests.
16//!
17//! The design follows Rust's derive pattern (similar to `#[derive(Default)]` with `#[default]`):
18//!
19//! ```ignore
20//! use rtc_interceptor::{Interceptor, StreamInfo, TaggedPacket, interceptor};
21//! use rtc_shared::error::Error;
22//! use sansio::Protocol;
23//! use std::collections::VecDeque;
24//!
25//! #[derive(Interceptor)]
26//! pub struct MyInterceptor<P: Interceptor> {
27//! #[next]
28//! next: P, // The next interceptor in the chain (can use any field name)
29//! buffer: VecDeque<TaggedPacket>,
30//! }
31//!
32//! #[interceptor]
33//! impl<P: Interceptor> MyInterceptor<P> {
34//! #[overrides]
35//! fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
36//! // Custom logic here
37//! self.next.handle_read(msg)
38//! }
39//! }
40//! ```
41//!
42//! # Pure Delegation (No Custom Logic)
43//!
44//! For interceptors that just pass through without modification:
45//!
46//! ```ignore
47//! # use rtc_interceptor::{Interceptor, StreamInfo, TaggedPacket, interceptor};
48//! # use rtc_shared::error::Error;
49//! # use sansio::Protocol;
50//! # use std::collections::VecDeque;
51//! #[derive(Interceptor)]
52//! pub struct PassthroughInterceptor<P: Interceptor> {
53//! #[next]
54//! next: P,
55//! }
56//!
57//! #[interceptor]
58//! impl<P: Interceptor> PassthroughInterceptor<P> {}
59//! // Empty impl block - all methods are auto-generated
60//! ```
61//!
62//! # Required Imports
63//!
64//! The macros require certain types to be in scope:
65//!
66//! The generated code names `sansio::Protocol`, `Error`, `StreamInfo` and
67//! `TaggedPacket`, so all four must be in scope at the use site — not only the macros
68//! themselves:
69//!
70//! ```ignore
71//! use rtc_interceptor::{Interceptor, StreamInfo, TaggedPacket, interceptor};
72//! use rtc_shared::error::Error;
73//! use sansio::Protocol;
74//! ```
75//!
76//! Through the `rtc` umbrella crate the same imports are:
77//!
78//! ```ignore
79//! use rtc::interceptor::{Interceptor, StreamInfo, TaggedPacket, interceptor};
80//! use rtc::sansio::Protocol;
81//! use rtc::shared::error::Error;
82//! ```
83
84use proc_macro::TokenStream;
85use quote::quote;
86use syn::{Data, DeriveInput, Fields, Ident, ImplItem, ItemImpl, Type, parse_macro_input};
87
88/// Derive macro that marks a struct as an interceptor.
89///
90/// This macro validates the struct has a `#[next]` field and generates
91/// a hidden accessor method. It does NOT generate Protocol/Interceptor implementations -
92/// those are generated by the `#[interceptor]` attribute on the impl block.
93///
94/// # Attributes
95///
96/// - `#[next]` - Mark the field that contains the next interceptor in the chain (required)
97///
98/// # Examples
99///
100/// Pure delegation (no custom logic):
101/// ```ignore
102/// # use rtc_interceptor::{Interceptor, StreamInfo, TaggedPacket, interceptor};
103/// # use rtc_shared::error::Error;
104/// # use sansio::Protocol;
105/// # use std::collections::VecDeque;
106/// #[derive(Interceptor)]
107/// pub struct PassthroughInterceptor<P: Interceptor> {
108/// #[next]
109/// next: P,
110/// }
111///
112/// #[interceptor]
113/// impl<P: Interceptor> PassthroughInterceptor<P> {}
114/// ```
115///
116/// With custom logic:
117/// ```ignore
118/// # use rtc_interceptor::{Interceptor, StreamInfo, TaggedPacket, interceptor};
119/// # use rtc_shared::error::Error;
120/// # use sansio::Protocol;
121/// # use std::collections::VecDeque;
122/// #[derive(Interceptor)]
123/// pub struct MyInterceptor<P: Interceptor> {
124/// #[next]
125/// next: P,
126/// buffer: VecDeque<TaggedPacket>,
127/// }
128///
129/// #[interceptor]
130/// impl<P: Interceptor> MyInterceptor<P> {
131/// #[overrides]
132/// fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
133/// // Custom logic
134/// self.next.handle_read(msg)
135/// }
136/// }
137/// ```
138#[proc_macro_derive(Interceptor, attributes(next))]
139pub fn derive_interceptor(input: TokenStream) -> TokenStream {
140 let input = parse_macro_input!(input as DeriveInput);
141
142 // Find the next field marked with #[next] - validates it exists and gets its type
143 let (next_name, next_type) = match find_next_field(&input) {
144 Ok(field) => field,
145 Err(err) => return err.into_compile_error().into(),
146 };
147
148 let name = &input.ident;
149 let generics = &input.generics;
150 let (impl_generics, ty_generics, where_clause) = generics.split_for_impl();
151
152 // Generate hidden accessor method that #[interceptor] will use
153 // This allows #[interceptor] to work without knowing the field name
154 let expanded = quote! {
155 impl #impl_generics #name #ty_generics #where_clause {
156 /// Hidden accessor for the next interceptor (used by #[interceptor] macro)
157 #[doc(hidden)]
158 #[inline(always)]
159 fn __interceptor_inner_mut(&mut self) -> &mut #next_type {
160 &mut self.#next_name
161 }
162 }
163 };
164
165 TokenStream::from(expanded)
166}
167
168/// Attribute macro for impl blocks to generate Protocol and Interceptor implementations.
169///
170/// This macro generates the trait implementations, delegating non-overridden
171/// methods to the next interceptor field (identified by `#[next]` in the struct).
172///
173/// **Important:** The struct must have `#[derive(Interceptor)]` with a `#[next]` field.
174///
175/// # Attributes
176///
177/// - `#[overrides]` - Mark methods that provide custom implementations
178///
179/// # Examples
180///
181/// With custom logic:
182/// ```ignore
183/// # use rtc_interceptor::{Interceptor, StreamInfo, TaggedPacket, interceptor};
184/// # use rtc_shared::error::Error;
185/// # use sansio::Protocol;
186/// # use std::collections::VecDeque;
187/// #[derive(Interceptor)]
188/// pub struct MyInterceptor<P: Interceptor> {
189/// #[next]
190/// next: P, // Can use any field name
191/// buffer: VecDeque<TaggedPacket>,
192/// }
193///
194/// #[interceptor]
195/// impl<P: Interceptor> MyInterceptor<P> {
196/// #[overrides]
197/// fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
198/// // Custom logic
199/// self.next.handle_read(msg)
200/// }
201/// }
202/// ```
203///
204/// Pure delegation (no custom logic):
205/// ```ignore
206/// # use rtc_interceptor::{Interceptor, StreamInfo, TaggedPacket, interceptor};
207/// # use rtc_shared::error::Error;
208/// # use sansio::Protocol;
209/// # use std::collections::VecDeque;
210/// #[derive(Interceptor)]
211/// pub struct PassthroughInterceptor<P: Interceptor> {
212/// #[next]
213/// wrapped: P, // Can use any field name
214/// }
215///
216/// #[interceptor]
217/// impl<P: Interceptor> PassthroughInterceptor<P> {}
218/// // Empty impl - all methods delegate to wrapped field
219/// ```
220#[proc_macro_attribute]
221pub fn interceptor(_attr: TokenStream, item: TokenStream) -> TokenStream {
222 let mut input = parse_macro_input!(item as ItemImpl);
223
224 // Note: Field name is no longer needed here - we use __interceptor_inner_mut() accessor
225 // which is generated by #[derive(Interceptor)]
226
227 // Collect method names marked with #[overrides]
228 let mut override_methods: Vec<Ident> = Vec::new();
229
230 for item in &mut input.items {
231 if let ImplItem::Fn(method) = item {
232 // Check if method has #[overrides] attribute
233 let has_override = method
234 .attrs
235 .iter()
236 .any(|attr| attr.path().is_ident("overrides"));
237
238 if has_override {
239 override_methods.push(method.sig.ident.clone());
240 // Remove the #[overrides] attribute
241 method
242 .attrs
243 .retain(|attr| !attr.path().is_ident("overrides"));
244 }
245 }
246 }
247
248 // Get type name and generics from the impl
249 let self_ty = &input.self_ty;
250 let generics = &input.generics;
251 let where_clause = &generics.where_clause;
252 let (impl_generics, _, _) = generics.split_for_impl();
253
254 // Generate Protocol methods that are NOT overridden (using accessor method)
255 let protocol_methods = generate_protocol_methods(&override_methods);
256 let interceptor_methods = generate_interceptor_methods(&override_methods);
257
258 // Protocol method names
259 let protocol_method_names = [
260 "handle_read",
261 "poll_read",
262 "handle_write",
263 "poll_write",
264 "handle_event",
265 "poll_event",
266 "handle_timeout",
267 "poll_timeout",
268 "close",
269 ];
270
271 // Interceptor method names
272 let interceptor_method_names = [
273 "bind_local_stream",
274 "unbind_local_stream",
275 "bind_remote_stream",
276 "unbind_remote_stream",
277 ];
278
279 // Extract Protocol overridden methods
280 let protocol_override_items: Vec<_> = input
281 .items
282 .iter()
283 .filter(|item| {
284 if let ImplItem::Fn(method) = item {
285 let name = method.sig.ident.to_string();
286 override_methods.contains(&method.sig.ident)
287 && protocol_method_names.contains(&name.as_str())
288 } else {
289 false
290 }
291 })
292 .collect();
293
294 // Extract Interceptor overridden methods
295 let interceptor_override_items: Vec<_> = input
296 .items
297 .iter()
298 .filter(|item| {
299 if let ImplItem::Fn(method) = item {
300 let name = method.sig.ident.to_string();
301 override_methods.contains(&method.sig.ident)
302 && interceptor_method_names.contains(&name.as_str())
303 } else {
304 false
305 }
306 })
307 .collect();
308
309 let expanded = quote! {
310 impl #impl_generics sansio::Protocol<
311 TaggedPacket,
312 TaggedPacket,
313 ()
314 > for #self_ty #where_clause {
315 type Rout = TaggedPacket;
316 type Wout = TaggedPacket;
317 type Eout = ();
318 type Error = Error;
319 type Time = std::time::Instant;
320
321 #protocol_methods
322 #(#protocol_override_items)*
323 }
324
325 impl #impl_generics Interceptor for #self_ty #where_clause {
326 #interceptor_methods
327 #(#interceptor_override_items)*
328 }
329 };
330
331 TokenStream::from(expanded)
332}
333
334/// Find the field marked with #[next] attribute, returning both name and type
335fn find_next_field(input: &DeriveInput) -> syn::Result<(Ident, Type)> {
336 let fields = match &input.data {
337 Data::Struct(data) => &data.fields,
338 _ => {
339 return Err(syn::Error::new_spanned(
340 input,
341 "Interceptor can only be derived for structs",
342 ));
343 }
344 };
345
346 let named_fields = match fields {
347 Fields::Named(fields) => &fields.named,
348 _ => {
349 return Err(syn::Error::new_spanned(
350 input,
351 "Interceptor can only be derived for structs with named fields",
352 ));
353 }
354 };
355
356 for field in named_fields {
357 let has_next_attr = field.attrs.iter().any(|attr| attr.path().is_ident("next"));
358 if has_next_attr {
359 let ident = field
360 .ident
361 .clone()
362 .ok_or_else(|| syn::Error::new_spanned(field, "Field must have a name"))?;
363 let ty = field.ty.clone();
364 return Ok((ident, ty));
365 }
366 }
367
368 Err(syn::Error::new_spanned(
369 input,
370 "No field marked with #[next] attribute. Mark the next interceptor field with #[next].",
371 ))
372}
373
374/// Generate Protocol methods that delegate to inner, excluding overridden ones
375fn generate_protocol_methods(override_methods: &[Ident]) -> proc_macro2::TokenStream {
376 let mut methods = proc_macro2::TokenStream::new();
377
378 if !override_methods.iter().any(|m| m == "handle_read") {
379 methods.extend(quote! {
380 fn handle_read(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
381 self.__interceptor_inner_mut().handle_read(msg)
382 }
383 });
384 }
385
386 if !override_methods.iter().any(|m| m == "poll_read") {
387 methods.extend(quote! {
388 fn poll_read(&mut self) -> Option<Self::Rout> {
389 self.__interceptor_inner_mut().poll_read()
390 }
391 });
392 }
393
394 if !override_methods.iter().any(|m| m == "handle_write") {
395 methods.extend(quote! {
396 fn handle_write(&mut self, msg: TaggedPacket) -> Result<(), Self::Error> {
397 self.__interceptor_inner_mut().handle_write(msg)
398 }
399 });
400 }
401
402 if !override_methods.iter().any(|m| m == "poll_write") {
403 methods.extend(quote! {
404 fn poll_write(&mut self) -> Option<Self::Wout> {
405 self.__interceptor_inner_mut().poll_write()
406 }
407 });
408 }
409
410 if !override_methods.iter().any(|m| m == "handle_event") {
411 methods.extend(quote! {
412 fn handle_event(&mut self, evt: ()) -> Result<(), Self::Error> {
413 self.__interceptor_inner_mut().handle_event(evt)
414 }
415 });
416 }
417
418 if !override_methods.iter().any(|m| m == "poll_event") {
419 methods.extend(quote! {
420 fn poll_event(&mut self) -> Option<Self::Eout> {
421 self.__interceptor_inner_mut().poll_event()
422 }
423 });
424 }
425
426 if !override_methods.iter().any(|m| m == "handle_timeout") {
427 methods.extend(quote! {
428 fn handle_timeout(&mut self, now: Self::Time) -> Result<(), Self::Error> {
429 self.__interceptor_inner_mut().handle_timeout(now)
430 }
431 });
432 }
433
434 if !override_methods.iter().any(|m| m == "poll_timeout") {
435 methods.extend(quote! {
436 fn poll_timeout(&mut self) -> Option<Self::Time> {
437 self.__interceptor_inner_mut().poll_timeout()
438 }
439 });
440 }
441
442 if !override_methods.iter().any(|m| m == "close") {
443 methods.extend(quote! {
444 fn close(&mut self) -> Result<(), Self::Error> {
445 self.__interceptor_inner_mut().close()
446 }
447 });
448 }
449
450 methods
451}
452
453/// Generate Interceptor methods that delegate to inner, excluding overridden ones
454fn generate_interceptor_methods(override_methods: &[Ident]) -> proc_macro2::TokenStream {
455 let mut methods = proc_macro2::TokenStream::new();
456
457 if !override_methods.iter().any(|m| m == "bind_local_stream") {
458 methods.extend(quote! {
459 fn bind_local_stream(&mut self, info: &StreamInfo) {
460 self.__interceptor_inner_mut().bind_local_stream(info);
461 }
462 });
463 }
464
465 if !override_methods.iter().any(|m| m == "unbind_local_stream") {
466 methods.extend(quote! {
467 fn unbind_local_stream(&mut self, info: &StreamInfo) {
468 self.__interceptor_inner_mut().unbind_local_stream(info);
469 }
470 });
471 }
472
473 if !override_methods.iter().any(|m| m == "bind_remote_stream") {
474 methods.extend(quote! {
475 fn bind_remote_stream(&mut self, info: &StreamInfo) {
476 self.__interceptor_inner_mut().bind_remote_stream(info);
477 }
478 });
479 }
480
481 if !override_methods.iter().any(|m| m == "unbind_remote_stream") {
482 methods.extend(quote! {
483 fn unbind_remote_stream(&mut self, info: &StreamInfo) {
484 self.__interceptor_inner_mut().unbind_remote_stream(info);
485 }
486 });
487 }
488
489 methods
490}