Skip to main content

rig_derive/
lib.rs

1extern crate proc_macro;
2
3use proc_macro::TokenStream;
4use syn::{DeriveInput, parse_macro_input};
5
6mod context;
7mod embed;
8mod resolve;
9mod tool;
10
11/// Derives `rig_core::embeddings::Embed` using fields marked with `#[embed]`.
12///
13/// ```
14/// use rig_derive::Embed;
15///
16/// #[derive(Embed)]
17/// struct Document {
18///     #[embed]
19///     description: String,
20/// }
21/// ```
22#[proc_macro_derive(Embed, attributes(embed))]
23pub fn derive_embedding_trait(item: TokenStream) -> TokenStream {
24    let mut input = parse_macro_input!(item as DeriveInput);
25
26    embed::expand_derive_embedding(&mut input)
27        .unwrap_or_else(syn::Error::into_compile_error)
28        .into()
29}
30
31/// Derives `rig_core::tool::ContextValue` with the type name as its default key.
32/// `#[context(key = "...")]` overrides the key. The type must support serde
33/// serialization and deserialization.
34///
35/// ```
36/// use rig_derive::ContextValue;
37///
38/// #[derive(serde::Serialize, serde::Deserialize, ContextValue)]
39/// #[context(key = "session.id")]
40/// struct SessionId(String);
41/// ```
42#[proc_macro_derive(ContextValue, attributes(context))]
43pub fn derive_context_value(item: TokenStream) -> TokenStream {
44    let input = parse_macro_input!(item as DeriveInput);
45    context::expand_derive_context_value(&input)
46        .unwrap_or_else(syn::Error::into_compile_error)
47        .into()
48}
49
50/// Generates a tool type, parameter struct, and uppercase static from a function
51/// returning `Result<T, E>`. Functions with a mutable context implement
52/// `rig_core::tool::Tool`; others implement `rig_core::tool::PortableTool`.
53/// Imported context names need `#[rig(context)]`; qualified Rig paths are
54/// recognized directly. The context is excluded from model arguments.
55///
56/// Accepts `name`, `description`, `params(name = "description")`, and
57/// `required(name, ...)`. Explicit names must start with an ASCII letter or `_`,
58/// contain only ASCII letters, digits, `_`, or `-`, and be at most 64 bytes.
59/// Description attributes override doc comments. Unknown or duplicate options
60/// and parameter names are rejected.
61///
62/// Non-`Option` arguments are required by default. An explicit required list
63/// makes omitted arguments default through serde, requiring `Default`.
64/// Listing an `Option` argument as required is rejected.
65///
66/// ```
67/// use rig_derive::rig_tool;
68///
69/// #[rig_tool(description = "Add integers", params(a = "First operand"), required(a))]
70/// fn add(a: i32, b: i32) -> Result<i32, rig_core::tool::ToolExecutionError> {
71///     Ok(a + b)
72/// }
73/// ```
74#[proc_macro_attribute]
75pub fn rig_tool(args: TokenStream, input: TokenStream) -> TokenStream {
76    let args = parse_macro_input!(args as tool::args::MacroArgs);
77    let input_fn = parse_macro_input!(input as syn::ItemFn);
78
79    tool::expand::expand_rig_tool(&args, &input_fn)
80        .unwrap_or_else(syn::Error::into_compile_error)
81        .into()
82}