Skip to main content

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}