Skip to main content

sha3_literal/
lib.rs

1use std::mem::take;
2
3use proc_macro2::Span;
4use quote::{ToTokens, quote, quote_spanned};
5use sha3::Digest;
6use syn::parse::discouraged::Speculative;
7use syn::parse::{ParseStream, Parser};
8use syn::{ExprArray, LitByte, LitByteStr, LitInt, Token, bracketed};
9use syn::{ExprMacro, LitStr, Macro, parse::Parse, parse_macro_input, token::Token};
10macro_rules! literals {
11    ($a:ident => $t:ty) => {
12        paste::paste! {
13            #[proc_macro]
14        pub fn [< $a _literal>](a: proc_macro::TokenStream) -> proc_macro::TokenStream {
15            let a = parse_macro_input!(a as HashLiteral).emit::<$t>();
16            quote! {#a}.into()
17        }
18        #[proc_macro]
19        pub fn [< $a _hex_literal>](a: proc_macro::TokenStream) -> proc_macro::TokenStream {
20            let a = parse_macro_input!(a as HashLiteral).emit_hex::<$t>();
21            // let a = Sha3HexLiteral(a);
22            quote! {#a}.into()
23        }
24        }
25    };
26}
27literals!(sha3 => sha3::Sha3_256);
28literals!(sha3_512 => sha3::Sha3_512);
29struct HashLiteral {
30    lit: (Vec<u8>, Span),
31    cb: Option<Macro>,
32}
33impl Parse for HashLiteral {
34    fn parse(input: syn::parse::ParseStream) -> syn::Result<Self> {
35        let a = parse_bytes(input)?;
36        let mut this = Self { lit: a, cb: None };
37        if input.peek(Token![=>]) {
38            input.parse::<Token![=>]>()?;
39            this.cb = Some(input.parse()?);
40        }
41        Ok(this)
42    }
43}
44fn parse_bytes(input: ParseStream) -> syn::Result<(Vec<u8>, Span)> {
45    let fork = input.fork();
46    if let Ok(l) = fork.parse::<LitStr>() {
47        input.advance_to(&fork);
48        return Ok((l.value().into_bytes(), l.span()));
49    }
50    let fork = input.fork();
51    if let Ok(l) = fork.parse::<LitByteStr>() {
52        input.advance_to(&fork);
53        return Ok((l.value(), l.span()));
54    }
55    let fork = input.fork();
56    if let Ok(l) = fork.parse::<LitByte>() {
57        input.advance_to(&fork);
58        return Ok((vec![l.value()], l.span()));
59    }
60    let fork = input.fork();
61    if let Ok(l) = fork.parse::<LitInt>() {
62        if let Ok(v) = l.base10_parse() {
63            input.advance_to(&fork);
64            return Ok((vec![v], l.span()));
65        }
66    }
67    let fork = input.fork();
68    if let Ok(a) = fork.parse::<ExprArray>() {
69        if let Ok(x) = a
70            .elems
71            .iter()
72            .map(|a| parse_bytes.parse2(quote! {#a}))
73            .collect::<syn::Result<Vec<_>>>()
74        {
75            input.advance_to(&fork);
76            let (x, y) = x.into_iter().collect::<(Vec<_>, Vec<_>)>();
77            return Ok((
78                x.into_iter().flatten().collect(),
79                y.into_iter()
80                    .map(Some)
81                    .reduce(|a, b| a?.join(b?))
82                    .flatten()
83                    .unwrap_or_else(|| Span::call_site()),
84            ));
85        }
86    }
87    let fork = input.fork();
88    if let Ok(a) = fork.parse::<Macro>() {
89        if let Some(s) = a.path.get_ident().map(|i| i.to_string()) {
90            match s.as_str() {
91                "include_bytes" | "include_str" => {
92                    let l: LitStr = a.parse_body()?;
93                    input.advance_to(&fork);
94                    let r = match std::fs::read(l.value()) {
95                        Ok(r) => r,
96                        Err(e) => return Err(syn::Error::new(l.span(), e)),
97                    };
98                    return Ok((r, l.span()));
99                }
100                "include" => {
101                    let l: LitStr = a.parse_body()?;
102                    input.advance_to(&fork);
103                    let r = match std::fs::read_to_string(l.value()) {
104                        Ok(r) => r,
105                        Err(e) => return Err(syn::Error::new(l.span(), e)),
106                    };
107                    let r: HashLiteral = syn::parse_str(&r)?;
108                    return Ok((r.lit.0, l.span()));
109                }
110                "sha3_literal" => {
111                    let (b, c) = a.parse_body_with(parse_bytes)?;
112                    input.advance_to(&fork);
113                    let b = sha3::Sha3_256::digest(b).into_iter().collect();
114                    return Ok((b, c));
115                }
116                "sha3_hex_literal" => {
117                    let (b, c) = a.parse_body_with(parse_bytes)?;
118                    input.advance_to(&fork);
119                    let b = sha3::Sha3_256::digest(b);
120                    let b = hex::encode(b);
121                    let b = b.into_bytes();
122                    return Ok((b, c));
123                }
124                _ => {}
125            }
126        }
127    }
128    return Err(input.error("expected a hashable thing"));
129}
130impl HashLiteral {
131    fn emit<D: Digest>(&self) -> proc_macro2::TokenStream {
132        let s = D::digest(&self.lit.0);
133        let a = quote_spanned! { self.lit.1 =>
134            [#(#s),*]
135        };
136        return (match self.cb.clone() {
137            None => a,
138            Some(mut c) => {
139                let t = take(&mut c.tokens);
140                c.tokens = quote! {#a #t};
141                quote! {#c}
142            }
143        });
144    }
145    fn emit_hex<D: Digest>(&self) -> proc_macro2::TokenStream {
146        let s = D::digest(&self.lit.0);
147        let s = hex::encode(s);
148        let a = quote_spanned! { self.lit.1 =>
149            #s
150        };
151        return (match self.cb.clone() {
152            None => a,
153            Some(mut c) => {
154                let t = take(&mut c.tokens);
155                c.tokens = quote! {#a #t};
156                quote! {#c}
157            }
158        });
159    }
160}