use proc_macro2::TokenStream;
use quote::quote;
use syn::{
Expr, ItemFn,
parse::{Parse, ParseStream},
};
pub struct ContractArgs {
pub path_expr: Option<Expr>,
}
impl Parse for ContractArgs {
fn parse(input: ParseStream) -> syn::Result<Self> {
if input.is_empty() {
return Ok(ContractArgs { path_expr: None });
}
let expr: Expr = input.parse()?;
Ok(ContractArgs {
path_expr: Some(expr),
})
}
}
fn normalize_path(path: &std::path::Path) -> std::path::PathBuf {
use std::path::Component;
let mut stack: Vec<std::ffi::OsString> = Vec::new();
for component in path.components() {
match component {
Component::Prefix(_) => {
stack.clear();
stack.push(component.as_os_str().to_owned());
}
Component::RootDir => {
stack.push(component.as_os_str().to_owned());
}
Component::CurDir => {
}
Component::ParentDir => {
let last_is_normal = stack
.last()
.map(|s| {
let p = std::path::Path::new(s);
matches!(p.components().next(), Some(Component::Normal(_)))
})
.unwrap_or(false);
if last_is_normal {
stack.pop();
}
}
Component::Normal(name) => {
stack.push(name.to_owned());
}
}
}
stack.iter().collect()
}
fn read_metadata_client_path() -> Option<String> {
let manifest_dir = std::env::var("CARGO_MANIFEST_DIR").ok()?;
let cargo_toml_path = std::path::Path::new(&manifest_dir).join("Cargo.toml");
let content = std::fs::read_to_string(cargo_toml_path).ok()?;
let manifest: toml::Value = toml::from_str(&content).ok()?;
let client_path = manifest
.get("package")
.and_then(|p| p.get("metadata"))
.and_then(|m| m.get("rorpc"))
.and_then(|r| r.get("client_path"))
.and_then(|v| v.as_str())?;
let joined = std::path::Path::new(&manifest_dir).join(client_path);
let absolute = normalize_path(&joined).to_string_lossy().into_owned();
Some(absolute)
}
pub fn expand_contract(args: ContractArgs, func: ItemFn) -> TokenStream {
let ItemFn {
attrs,
vis,
sig,
block,
..
} = func;
let original_body = &block.stmts;
let path_tokens: TokenStream = if let Some(expr) = args.path_expr {
quote! { #expr }
} else if let Some(path) = read_metadata_client_path() {
quote! { #path }
} else {
quote! { env!("RORPC_CLIENT_PATH") }
};
quote! {
#(#attrs)*
#vis #sig {
#[cfg(debug_assertions)]
{
::rorpc::generate_contract()
.output(#path_tokens)
.expect("contract generation failed");
}
#(#original_body)*
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use quote::quote;
#[test]
fn parse_empty_args() {
let args: ContractArgs = syn::parse2(quote! {}).expect("parse failed");
assert!(args.path_expr.is_none());
}
#[test]
fn parse_string_literal() {
let args: ContractArgs =
syn::parse2(quote! { "../client/bindings.ts" }).expect("parse failed");
assert!(args.path_expr.is_some());
}
#[test]
fn parse_env_macro() {
let args: ContractArgs =
syn::parse2(quote! { env!("RORPC_CLIENT_PATH") }).expect("parse failed");
assert!(args.path_expr.is_some());
}
#[test]
fn parse_concat_macro() {
let args: ContractArgs = syn::parse2(quote! {
concat!(env!("CARGO_MANIFEST_DIR"), "/../client/src/rpc/bindings.ts")
})
.expect("parse failed");
assert!(args.path_expr.is_some());
}
#[test]
fn parse_constant() {
let args: ContractArgs = syn::parse2(quote! { CLIENT_PATH }).expect("parse failed");
assert!(args.path_expr.is_some());
}
#[test]
fn expand_with_string_literal() {
let func: ItemFn = syn::parse2(quote! {
fn main() { println!("Hello"); }
})
.expect("parse failed");
let args: ContractArgs =
syn::parse2(quote! { "../client/bindings.ts" }).expect("parse failed");
let expanded = expand_contract(args, func);
let s = expanded.to_string();
assert!(s.contains("\"../client/bindings.ts\""));
assert!(s.contains("rorpc :: generate_contract"));
assert!(s.contains("# [cfg (debug_assertions)]") || s.contains("#[cfg(debug_assertions)]"));
}
#[test]
fn expand_preserves_attributes() {
let func: ItemFn = syn::parse2(quote! {
#[tokio::main]
async fn main() { println!("Hello"); }
})
.expect("parse failed");
let args: ContractArgs =
syn::parse2(quote! { "../client/bindings.ts" }).expect("parse failed");
let expanded = expand_contract(args, func);
let s = expanded.to_string();
assert!(s.contains("# [tokio :: main]") || s.contains("#[tokio::main]"));
assert!(s.contains("async fn main"));
}
#[test]
fn normalize_simple_parent_traversal() {
let p = std::path::Path::new("/repo/server/crate").join("../../out.ts");
assert_eq!(normalize_path(&p), std::path::Path::new("/repo/out.ts"));
}
#[test]
fn normalize_sibling_dir() {
let p = std::path::Path::new("/repo/server").join("../client/src/bindings.ts");
assert_eq!(
normalize_path(&p),
std::path::Path::new("/repo/client/src/bindings.ts")
);
}
#[test]
fn normalize_deep_traversal_stops_at_root() {
let p = std::path::Path::new("/a/b").join("../../../../out.ts");
assert_eq!(normalize_path(&p), std::path::Path::new("/out.ts"));
}
#[test]
fn normalize_curdirs_are_skipped() {
let p = std::path::Path::new("/repo/./server/./crate").join("./out.ts");
assert_eq!(
normalize_path(&p),
std::path::Path::new("/repo/server/crate/out.ts")
);
}
#[test]
fn normalize_already_clean_path_unchanged() {
let p = std::path::Path::new("/repo/client/src/bindings.ts");
assert_eq!(normalize_path(p), p);
}
#[cfg(windows)]
#[test]
fn normalize_windows_preserves_drive_letter_shallow() {
let base = std::path::Path::new(
r"D:\programming\Rust\rust-orpc\examples\axum-react\better-auth-integration",
);
let p = base.join("../client/src/rpc/bindings.ts");
assert_eq!(
normalize_path(&p),
std::path::Path::new(
r"D:\programming\Rust\rust-orpc\examples\axum-react\client\src\rpc\bindings.ts"
),
);
}
#[cfg(windows)]
#[test]
fn normalize_windows_deep_traversal_to_near_root() {
let base = std::path::Path::new(
r"D:\programming\Rust\rust-orpc\examples\axum-react\better-auth-integration",
);
let p = base.join("../../../../../out.ts");
assert_eq!(
normalize_path(&p),
std::path::Path::new(r"D:\programming\out.ts"),
);
}
#[cfg(windows)]
#[test]
fn normalize_windows_excessive_traversal_stops_at_root() {
let base = std::path::Path::new(r"D:\a\b");
let p = base.join("../../../../../out.ts");
assert_eq!(normalize_path(&p), std::path::Path::new(r"D:\out.ts"));
}
#[cfg(windows)]
#[test]
fn normalize_windows_no_traversal_unchanged() {
let p = std::path::Path::new(r"D:\programming\Rust\client\src\bindings.ts");
assert_eq!(normalize_path(p), p);
}
}