Skip to main content

agent_client_protocol_derive/
lib.rs

1//! Derive macros for Agent Client Protocol JSON-RPC traits.
2//!
3//! This crate provides derive macros to reduce boilerplate when implementing
4//! custom JSON-RPC requests, notifications, and response types.
5//!
6//! # Example
7//!
8//! ```ignore
9//! use agent_client_protocol::{JsonRpcRequest, JsonRpcNotification, JsonRpcResponse};
10//!
11//! #[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)]
12//! #[request(method = "_hello", response = HelloResponse)]
13//! struct HelloRequest {
14//!     name: String,
15//! }
16//!
17//! #[derive(Debug, Clone, Serialize, Deserialize, JsonRpcResponse)]
18//! struct HelloResponse {
19//!     greeting: String,
20//! }
21//!
22//! #[derive(Debug, Clone, Serialize, Deserialize, JsonRpcNotification)]
23//! #[notification(method = "_ping")]
24//! struct PingNotification {
25//!     timestamp: u64,
26//! }
27//! ```
28//!
29//! # Using within the `agent_client_protocol` crate
30//!
31//! When using these derives within the `agent_client_protocol` crate itself, add `crate = crate`:
32//!
33//! ```ignore
34//! #[derive(JsonRpcRequest)]
35//! #[request(method = "_foo", response = FooResponse, crate = crate)]
36//! struct FooRequest { ... }
37//! ```
38
39use proc_macro::TokenStream;
40use quote::{format_ident, quote};
41use syn::{DeriveInput, GenericParam, Generics, Ident, LitStr, Path, Type, parse_macro_input};
42
43/// Derive macro for implementing `JsonRpcRequest` and `JsonRpcMessage` traits.
44///
45/// # Attributes
46///
47/// - `#[request(method = "method_name", response = ResponseType)]`, where `ResponseType` may be
48///   any Rust type, including generic types such as `Option<Response>`
49/// - `#[request(method = "method_name", response = ResponseType, crate = crate)]` - for use within the `agent_client_protocol` crate
50///
51/// # Example
52///
53/// ```ignore
54/// #[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)]
55/// #[request(method = "_hello", response = HelloResponse)]
56/// struct HelloRequest {
57///     name: String,
58/// }
59/// ```
60#[proc_macro_derive(JsonRpcRequest, attributes(request))]
61pub fn derive_json_rpc_request(input: TokenStream) -> TokenStream {
62    let input = parse_macro_input!(input as DeriveInput);
63    let name = &input.ident;
64    let method_arg = fresh_ident(&input, "method");
65    let params_arg = fresh_ident(&input, "params");
66
67    // Parse attributes
68    let (method, response_type, krate) = match parse_request_attrs(&input) {
69        Ok(attrs) => attrs,
70        Err(e) => return e.to_compile_error().into(),
71    };
72
73    let message_generics = message_generics(&input.generics, &krate);
74    let (message_impl_generics, type_generics, message_where_clause) =
75        message_generics.split_for_impl();
76    let request_generics = request_generics(&input.generics, &response_type, &krate);
77    let (request_impl_generics, _, request_where_clause) = request_generics.split_for_impl();
78
79    let expanded = quote! {
80        #[automatically_derived]
81        impl #message_impl_generics #krate::JsonRpcMessage for #name #type_generics #message_where_clause {
82            fn matches_method(#method_arg: &str) -> bool {
83                #method_arg == #method
84            }
85
86            fn method(&self) -> &str {
87                #method
88            }
89
90            fn to_untyped_message(&self) -> ::core::result::Result<#krate::UntypedMessage, #krate::Error> {
91                #krate::UntypedMessage::new(#method, self)
92            }
93
94            fn parse_message(
95                #method_arg: &str,
96                #params_arg: &impl #krate::__private::serde::Serialize,
97            ) -> ::core::result::Result<Self, #krate::Error> {
98                if #method_arg != #method {
99                    return ::core::result::Result::Err(#krate::Error::method_not_found());
100                }
101                #krate::util::json_cast_params(#params_arg)
102            }
103        }
104
105        #[automatically_derived]
106        impl #request_impl_generics #krate::JsonRpcRequest for #name #type_generics #request_where_clause {
107            type Response = #response_type;
108        }
109    };
110
111    TokenStream::from(expanded)
112}
113
114/// Derive macro for implementing `JsonRpcNotification` and `JsonRpcMessage` traits.
115///
116/// # Attributes
117///
118/// - `#[notification(method = "method_name")]`
119/// - `#[notification(method = "method_name", crate = crate)]` - for use within the `agent_client_protocol` crate
120///
121/// # Example
122///
123/// ```ignore
124/// #[derive(Debug, Clone, Serialize, Deserialize, JsonRpcNotification)]
125/// #[notification(method = "_ping")]
126/// struct PingNotification {
127///     timestamp: u64,
128/// }
129/// ```
130#[proc_macro_derive(JsonRpcNotification, attributes(notification))]
131pub fn derive_json_rpc_notification(input: TokenStream) -> TokenStream {
132    let input = parse_macro_input!(input as DeriveInput);
133    let name = &input.ident;
134    let method_arg = fresh_ident(&input, "method");
135    let params_arg = fresh_ident(&input, "params");
136
137    // Parse attributes
138    let (method, krate) = match parse_notification_attrs(&input) {
139        Ok(attrs) => attrs,
140        Err(e) => return e.to_compile_error().into(),
141    };
142
143    let message_generics = message_generics(&input.generics, &krate);
144    let (message_impl_generics, type_generics, message_where_clause) =
145        message_generics.split_for_impl();
146    let marker_generics = marker_generics(&input.generics, &krate);
147    let (marker_impl_generics, _, marker_where_clause) = marker_generics.split_for_impl();
148
149    let expanded = quote! {
150        #[automatically_derived]
151        impl #message_impl_generics #krate::JsonRpcMessage for #name #type_generics #message_where_clause {
152            fn matches_method(#method_arg: &str) -> bool {
153                #method_arg == #method
154            }
155
156            fn method(&self) -> &str {
157                #method
158            }
159
160            fn to_untyped_message(&self) -> ::core::result::Result<#krate::UntypedMessage, #krate::Error> {
161                #krate::UntypedMessage::new(#method, self)
162            }
163
164            fn parse_message(
165                #method_arg: &str,
166                #params_arg: &impl #krate::__private::serde::Serialize,
167            ) -> ::core::result::Result<Self, #krate::Error> {
168                if #method_arg != #method {
169                    return ::core::result::Result::Err(#krate::Error::method_not_found());
170                }
171                #krate::util::json_cast_params(#params_arg)
172            }
173        }
174
175        #[automatically_derived]
176        impl #marker_impl_generics #krate::JsonRpcNotification for #name #type_generics #marker_where_clause {}
177    };
178
179    TokenStream::from(expanded)
180}
181
182/// Derive macro for implementing `JsonRpcResponse` trait.
183///
184/// # Attributes
185///
186/// - `#[response(crate = crate)]` - for use within the `agent_client_protocol` crate
187///
188/// # Example
189///
190/// ```ignore
191/// #[derive(Debug, Clone, Serialize, Deserialize, JsonRpcResponse)]
192/// struct HelloResponse {
193///     greeting: String,
194/// }
195/// ```
196#[proc_macro_derive(JsonRpcResponse, attributes(response))]
197pub fn derive_json_rpc_response_payload(input: TokenStream) -> TokenStream {
198    let input = parse_macro_input!(input as DeriveInput);
199    let name = &input.ident;
200    let method_arg = fresh_ident(&input, "method");
201    let value_arg = fresh_ident(&input, "value");
202
203    let krate = match parse_response_attrs(&input) {
204        Ok(attrs) => attrs,
205        Err(e) => return e.to_compile_error().into(),
206    };
207
208    let response_generics = response_payload_generics(&input.generics, &krate);
209    let (impl_generics, type_generics, where_clause) = response_generics.split_for_impl();
210
211    let expanded = quote! {
212        #[automatically_derived]
213        impl #impl_generics #krate::JsonRpcResponse for #name #type_generics #where_clause {
214            fn into_json(self, #method_arg: &str) -> ::core::result::Result<#krate::__private::serde_json::Value, #krate::Error> {
215                #krate::__private::serde_json::to_value(self).map_err(#krate::Error::into_internal_error)
216            }
217
218            fn from_value(#method_arg: &str, #value_arg: #krate::__private::serde_json::Value) -> ::core::result::Result<Self, #krate::Error> {
219                #krate::util::json_cast(#value_arg)
220            }
221        }
222    };
223
224    TokenStream::from(expanded)
225}
226
227fn default_crate_path() -> Path {
228    syn::parse_quote!(agent_client_protocol)
229}
230
231fn fresh_ident(input: &DeriveInput, role: &str) -> Ident {
232    let mut suffix = 0;
233    loop {
234        let candidate = if suffix == 0 {
235            format_ident!("__acp_{role}")
236        } else {
237            format_ident!("__acp_{role}_{suffix}")
238        };
239        let collides = input.generics.params.iter().any(|param| match param {
240            GenericParam::Lifetime(param) => param.lifetime.ident == candidate,
241            GenericParam::Type(param) => param.ident == candidate,
242            GenericParam::Const(param) => param.ident == candidate,
243        });
244        if !collides {
245            return candidate;
246        }
247        suffix += 1;
248    }
249}
250
251fn message_generics(generics: &Generics, krate: &Path) -> Generics {
252    let mut generics = generics.clone();
253    generics
254        .make_where_clause()
255        .predicates
256        .push(syn::parse_quote! {
257            Self: ::core::fmt::Debug
258                + ::core::clone::Clone
259                + ::core::marker::Send
260                + #krate::__private::serde::Serialize
261                + #krate::__private::serde::de::DeserializeOwned
262                + 'static
263        });
264    generics
265}
266
267fn marker_generics(generics: &Generics, krate: &Path) -> Generics {
268    let mut generics = generics.clone();
269    generics
270        .make_where_clause()
271        .predicates
272        .push(syn::parse_quote!(Self: #krate::JsonRpcMessage));
273    generics
274}
275
276fn request_generics(generics: &Generics, response: &Type, krate: &Path) -> Generics {
277    let mut generics = marker_generics(generics, krate);
278    generics
279        .make_where_clause()
280        .predicates
281        .push(syn::parse_quote!(#response: #krate::JsonRpcResponse));
282    generics
283}
284
285fn response_payload_generics(generics: &Generics, krate: &Path) -> Generics {
286    let mut generics = generics.clone();
287    generics
288        .make_where_clause()
289        .predicates
290        .push(syn::parse_quote! {
291            Self: ::core::fmt::Debug
292                + ::core::clone::Clone
293                + ::core::marker::Send
294                + #krate::__private::serde::Serialize
295                + #krate::__private::serde::de::DeserializeOwned
296                + 'static
297        });
298    generics
299}
300
301fn parse_request_attrs(input: &DeriveInput) -> syn::Result<(LitStr, Type, Path)> {
302    let mut method: Option<LitStr> = None;
303    let mut response_type: Option<Type> = None;
304    let mut krate: Option<Path> = None;
305
306    for attr in &input.attrs {
307        if !attr.path().is_ident("request") {
308            continue;
309        }
310
311        attr.parse_nested_meta(|meta| {
312            if meta.path.is_ident("method") {
313                if method.is_some() {
314                    return Err(meta.error("duplicate `method` attribute"));
315                }
316                let value: LitStr = meta.value()?.parse()?;
317                method = Some(value);
318                return Ok(());
319            }
320
321            if meta.path.is_ident("response") {
322                if response_type.is_some() {
323                    return Err(meta.error("duplicate `response` attribute"));
324                }
325                response_type = Some(meta.value()?.parse()?);
326                return Ok(());
327            }
328
329            if meta.path.is_ident("crate") {
330                if krate.is_some() {
331                    return Err(meta.error("duplicate `crate` attribute"));
332                }
333                krate = Some(meta.value()?.parse()?);
334                return Ok(());
335            }
336
337            Err(meta.error("unknown attribute"))
338        })?;
339    }
340
341    let method = method.ok_or_else(|| {
342        syn::Error::new_spanned(
343            &input.ident,
344            "missing required attribute: #[request(method = \"...\")]",
345        )
346    })?;
347
348    let response_type = response_type.ok_or_else(|| {
349        syn::Error::new_spanned(
350            &input.ident,
351            "missing required attribute: #[request(response = ...)]",
352        )
353    })?;
354
355    Ok((
356        method,
357        response_type,
358        krate.unwrap_or_else(default_crate_path),
359    ))
360}
361
362fn parse_notification_attrs(input: &DeriveInput) -> syn::Result<(LitStr, Path)> {
363    let mut method: Option<LitStr> = None;
364    let mut krate: Option<Path> = None;
365
366    for attr in &input.attrs {
367        if !attr.path().is_ident("notification") {
368            continue;
369        }
370
371        attr.parse_nested_meta(|meta| {
372            if meta.path.is_ident("method") {
373                if method.is_some() {
374                    return Err(meta.error("duplicate `method` attribute"));
375                }
376                let value: LitStr = meta.value()?.parse()?;
377                method = Some(value);
378                return Ok(());
379            }
380
381            if meta.path.is_ident("crate") {
382                if krate.is_some() {
383                    return Err(meta.error("duplicate `crate` attribute"));
384                }
385                krate = Some(meta.value()?.parse()?);
386                return Ok(());
387            }
388
389            Err(meta.error("unknown attribute"))
390        })?;
391    }
392
393    let method = method.ok_or_else(|| {
394        syn::Error::new_spanned(
395            &input.ident,
396            "missing required attribute: #[notification(method = \"...\")]",
397        )
398    })?;
399
400    Ok((method, krate.unwrap_or_else(default_crate_path)))
401}
402
403fn parse_response_attrs(input: &DeriveInput) -> syn::Result<Path> {
404    let mut krate: Option<Path> = None;
405
406    for attr in &input.attrs {
407        if !attr.path().is_ident("response") {
408            continue;
409        }
410
411        attr.parse_nested_meta(|meta| {
412            if meta.path.is_ident("crate") {
413                if krate.is_some() {
414                    return Err(meta.error("duplicate `crate` attribute"));
415                }
416                krate = Some(meta.value()?.parse()?);
417                return Ok(());
418            }
419
420            Err(meta.error("unknown attribute"))
421        })?;
422    }
423
424    Ok(krate.unwrap_or_else(default_crate_path))
425}
426
427#[cfg(test)]
428mod tests {
429    use super::*;
430    use quote::quote;
431    use syn::parse_quote;
432
433    fn expect_error<T>(result: syn::Result<T>) -> syn::Error {
434        match result {
435            Ok(_) => panic!("expected attribute parsing to fail"),
436            Err(error) => error,
437        }
438    }
439
440    #[test]
441    fn request_attributes_accept_rust_types() {
442        let input = parse_quote! {
443            #[request(
444                method = "test/method",
445                response = Result<Option<Response>, Error>,
446                crate = crate::protocol
447            )]
448            struct Request;
449        };
450
451        let (method, response, krate) = parse_request_attrs(&input).unwrap();
452
453        assert_eq!(method.value(), "test/method");
454        assert_eq!(
455            quote!(#response).to_string(),
456            "Result < Option < Response > , Error >"
457        );
458        assert_eq!(quote!(#krate).to_string(), "crate :: protocol");
459    }
460
461    #[test]
462    fn request_attributes_use_a_relative_default_crate_path() {
463        let input = parse_quote! {
464            #[request(method = "test/method", response = Response)]
465            struct Request;
466        };
467
468        let (_, _, krate) = parse_request_attrs(&input).unwrap();
469
470        assert_eq!(quote!(#krate).to_string(), "agent_client_protocol");
471    }
472
473    #[test]
474    fn request_attributes_reject_duplicate_keys() {
475        let input = parse_quote! {
476            #[request(method = "first", method = "second", response = Response)]
477            struct Request;
478        };
479
480        let error = expect_error(parse_request_attrs(&input));
481
482        assert_eq!(error.to_string(), "duplicate `method` attribute");
483    }
484
485    #[test]
486    fn notification_attributes_reject_duplicate_keys_across_attributes() {
487        let input = parse_quote! {
488            #[notification(method = "test/method")]
489            #[notification(method = "test/other")]
490            struct Notification;
491        };
492
493        let error = expect_error(parse_notification_attrs(&input));
494
495        assert_eq!(error.to_string(), "duplicate `method` attribute");
496    }
497
498    #[test]
499    fn response_attributes_reject_duplicate_crate_paths() {
500        let input = parse_quote! {
501            #[response(crate = crate, crate = agent_client_protocol)]
502            struct Response;
503        };
504
505        let error = expect_error(parse_response_attrs(&input));
506
507        assert_eq!(error.to_string(), "duplicate `crate` attribute");
508    }
509}