use proc_macro::{TokenStream, TokenTree, Literal, Punct, Spacing, Ident, Span, Group, Delimiter};
#[proc_macro]
pub fn include_stdlib(input: TokenStream) -> TokenStream {
let mut iter = input.into_iter();
let path_token = match iter.next() {
Some(TokenTree::Literal(lit)) => lit,
_ => panic!("include_stdlib! expects a string literal path"),
};
let path_str = path_token.to_string();
let path = path_str.trim_matches('"');
let manifest_dir = std::env::var("CARGO_MANIFEST_DIR")
.expect("CARGO_MANIFEST_DIR not set");
let full_path = std::path::Path::new(&manifest_dir).join(path);
let content = std::fs::read_to_string(&full_path)
.unwrap_or_else(|e| panic!("Failed to read {}: {}", full_path.display(), e));
let stdlib_entries = parse_lisp_file(&content);
let mut tokens = Vec::new();
tokens.push(TokenTree::Ident(Ident::new("define_stdlib", Span::call_site())));
tokens.push(TokenTree::Punct(Punct::new('!', Spacing::Alone)));
let mut body_tokens = Vec::new();
for entry in stdlib_entries {
if let Some(doc) = &entry.doc {
body_tokens.push(TokenTree::Punct(Punct::new('#', Spacing::Alone)));
let mut attr_tokens = Vec::new();
attr_tokens.push(TokenTree::Ident(Ident::new("doc", Span::call_site())));
attr_tokens.push(TokenTree::Punct(Punct::new('=', Spacing::Alone)));
attr_tokens.push(TokenTree::Literal(Literal::string(doc)));
body_tokens.push(TokenTree::Group(Group::new(
Delimiter::Bracket,
attr_tokens.into_iter().collect(),
)));
}
body_tokens.push(TokenTree::Ident(Ident::new(&entry.variant_name, Span::call_site())));
let mut args = Vec::new();
args.push(TokenTree::Literal(Literal::string(&entry.name)));
args.push(TokenTree::Punct(Punct::new(',', Spacing::Alone)));
let mut params_tokens = Vec::new();
for (i, param) in entry.params.iter().enumerate() {
if i > 0 {
params_tokens.push(TokenTree::Punct(Punct::new(',', Spacing::Alone)));
}
params_tokens.push(TokenTree::Literal(Literal::string(param)));
}
args.push(TokenTree::Group(Group::new(
Delimiter::Bracket,
params_tokens.into_iter().collect(),
)));
args.push(TokenTree::Punct(Punct::new(',', Spacing::Alone)));
args.push(TokenTree::Literal(Literal::string(&entry.body)));
body_tokens.push(TokenTree::Group(Group::new(
Delimiter::Parenthesis,
args.into_iter().collect(),
)));
body_tokens.push(TokenTree::Punct(Punct::new(',', Spacing::Alone)));
}
tokens.push(TokenTree::Group(Group::new(
Delimiter::Brace,
body_tokens.into_iter().collect(),
)));
tokens.into_iter().collect()
}
struct StdlibEntry {
variant_name: String,
name: String,
params: Vec<String>,
body: String,
doc: Option<String>,
}
fn parse_lisp_file(content: &str) -> Vec<StdlibEntry> {
let mut entries = Vec::new();
let mut lines = content.lines().peekable();
while let Some(line) = lines.next() {
let trimmed = line.trim();
if trimmed.is_empty() {
continue;
}
let doc = if trimmed.starts_with(";;;") {
Some(trimmed.strip_prefix(";;;").unwrap_or("").trim().to_string())
} else {
None
};
let def_line = if doc.is_some() {
loop {
match lines.next() {
Some(l) if l.trim().is_empty() || l.trim().starts_with(";;") => continue,
Some(l) => break l,
None => break "",
}
}
} else if trimmed.starts_with(";;") {
continue;
} else {
line
};
let def_trimmed = def_line.trim();
if def_trimmed.starts_with("(define")
&& let Some(entry) = parse_define(def_trimmed, doc, &mut lines)
{
entries.push(entry);
}
}
entries
}
fn strip_inline_comment(line: &str) -> String {
let mut result = String::new();
let mut in_string = false;
let mut escape_next = false;
let mut chars = line.chars().peekable();
while let Some(c) = chars.next() {
if escape_next {
result.push(c);
escape_next = false;
continue;
}
match c {
'\\' if in_string => {
result.push(c);
escape_next = true;
}
'"' => {
in_string = !in_string;
result.push(c);
}
';' if !in_string => {
break;
}
_ => {
result.push(c);
}
}
}
result
}
fn parse_define<'a, I: Iterator<Item = &'a str>>(
first_line: &str,
doc: Option<String>,
remaining_lines: &mut I,
) -> Option<StdlibEntry> {
let mut full_def = strip_inline_comment(first_line);
let mut paren_count = 0;
for c in full_def.chars() {
match c {
'(' => paren_count += 1,
')' => paren_count -= 1,
_ => {}
}
}
while paren_count > 0 {
match remaining_lines.next() {
Some(line) => {
let stripped = strip_inline_comment(line.trim());
full_def.push(' ');
full_def.push_str(&stripped);
for c in stripped.chars() {
match c {
'(' => paren_count += 1,
')' => paren_count -= 1,
_ => {}
}
}
}
None => break,
}
}
let content = full_def.trim();
if !content.starts_with("(define") {
return None;
}
let after_define = content.strip_prefix("(define").unwrap_or("").trim_start();
if !after_define.starts_with('(') {
return None;
}
let sig_start = 1; let mut paren_depth = 1;
let mut sig_end = sig_start;
let after_define_chars: Vec<char> = after_define.chars().collect();
for (i, &c) in after_define_chars[sig_start..].iter().enumerate() {
match c {
'(' => paren_depth += 1,
')' => {
paren_depth -= 1;
if paren_depth == 0 {
sig_end = sig_start + i;
break;
}
}
_ => {}
}
}
let sig_content: String = after_define_chars[sig_start..sig_end].iter().collect();
let sig_parts: Vec<&str> = sig_content.split_whitespace().collect();
if sig_parts.is_empty() {
return None;
}
let name = sig_parts[0].to_string();
let params: Vec<String> = sig_parts[1..].iter().map(|s| s.to_string()).collect();
let body_start = sig_end + 1; let after_sig: String = after_define_chars[body_start..].iter().collect();
let body_trimmed = after_sig.trim();
let raw_body = body_trimmed.strip_suffix(')').unwrap_or(body_trimmed).trim();
let body = format!("(begin {})", raw_body);
let variant_name = to_pascal_case(&name);
Some(StdlibEntry {
variant_name,
name,
params,
body,
doc,
})
}
fn to_pascal_case(name: &str) -> String {
grift_util::to_pascal_case(name)
}