use proc_macro2::TokenStream;
use quote::quote;
use std::collections::HashSet;
use syn::Token;
use syn::parse::{Parse, ParseStream};
pub fn expand_container(input: TokenStream) -> TokenStream {
if input.is_empty() {
return quote! {
injectable_rs::Container::builder().build()
};
}
let entries = match syn::parse2::<ContainerInput>(input) {
Ok(input) => input.entries,
Err(err) => return err.to_compile_error(),
};
match validate_graph(&entries) {
Ok(()) => {
generate_container_code(&entries)
}
Err(errors) => {
let error_tokens: Vec<TokenStream> = errors
.iter()
.map(|err| {
let msg = err.to_string();
quote! { compile_error!(#msg); }
})
.collect();
quote! { #(#error_tokens)* }
}
}
}
struct ContainerInput {
entries: Vec<TypeEntry>,
}
struct TypeEntry {
name: syn::Ident,
name_str: String,
dependencies: Vec<String>,
scope: String,
}
impl Parse for ContainerInput {
fn parse(input: ParseStream) -> syn::Result<Self> {
let mut entries = Vec::new();
while !input.is_empty() {
entries.push(input.parse::<TypeEntry>()?);
let _ = input.parse::<Token![,]>();
}
Ok(ContainerInput { entries })
}
}
impl Parse for TypeEntry {
fn parse(input: ParseStream) -> syn::Result<Self> {
let name: syn::Ident = input.parse()?;
let name_str = name.to_string();
let mut dependencies = Vec::new();
let mut scope = "singleton".to_string();
if input.peek(syn::token::Brace) {
let content;
syn::braced!(content in input);
while !content.is_empty() {
let key: syn::Ident = content.parse()?;
match key.to_string().as_str() {
"deps" => {
content.parse::<Token![:]>()?;
let dep_list;
syn::bracketed!(dep_list in content);
while !dep_list.is_empty() {
let dep: syn::Ident = dep_list.parse()?;
dependencies.push(dep.to_string());
if dep_list.peek(Token![,]) {
dep_list.parse::<Token![,]>()?;
}
}
}
"scope" => {
content.parse::<Token![:]>()?;
let scope_lit: syn::LitStr = content.parse()?;
scope = scope_lit.value();
}
other => {
return Err(syn::Error::new(
key.span(),
format!(
"unknown container entry attribute: `{other}`; expected `deps` or `scope`"
),
));
}
}
if content.peek(Token![,]) {
content.parse::<Token![,]>()?;
}
}
}
Ok(TypeEntry {
name,
name_str,
dependencies,
scope,
})
}
}
fn generate_container_code(entries: &[TypeEntry]) -> TokenStream {
let type_assertions: Vec<TokenStream> = entries
.iter()
.map(|entry| {
let type_name = &entry.name;
quote! {
let _ = || {
fn _assert_injectable<T: injectable_rs_runtime::Injectable>() {}
_assert_injectable::<#type_name>();
};
}
})
.collect();
quote! {
{
#(#type_assertions)*
injectable_rs::Container::builder().build()
}
}
}
fn validate_graph(entries: &[TypeEntry]) -> Result<(), Vec<CompileValidationError>> {
let mut errors = Vec::new();
check_duplicates(entries, &mut errors);
check_missing(entries, &mut errors);
check_cycles(entries, &mut errors);
check_scope_mismatches(entries, &mut errors);
if errors.is_empty() {
Ok(())
} else {
Err(errors)
}
}
#[derive(Debug, Clone)]
enum CompileValidationError {
CircularDependency { chain: Vec<String> },
MissingDependency { source: String, missing: String },
DuplicateNode { name: String },
ScopeMismatch {
source: String,
source_scope: String,
dependency: String,
dependency_scope: String,
},
}
impl std::fmt::Display for CompileValidationError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::CircularDependency { chain } => {
write!(f, "circular dependency detected: ")?;
for (i, t) in chain.iter().enumerate() {
if i > 0 {
write!(f, " -> ")?;
}
write!(f, "{t}")?;
}
Ok(())
}
Self::MissingDependency { source, missing } => {
write!(
f,
"`{source}` depends on `{missing}`, which is not registered in the container"
)
}
Self::DuplicateNode { name } => {
write!(f, "duplicate type `{name}` registered in the container")
}
Self::ScopeMismatch {
source,
source_scope,
dependency,
dependency_scope,
} => {
write!(
f,
"scope mismatch: `{source}` ({source_scope}) depends on `{dependency}` ({dependency_scope}); \
wider-scope types cannot depend on narrower-scope types"
)
}
}
}
}
fn check_duplicates(entries: &[TypeEntry], errors: &mut Vec<CompileValidationError>) {
let mut seen = HashSet::new();
for entry in entries {
if !seen.insert(&entry.name_str) {
errors.push(CompileValidationError::DuplicateNode {
name: entry.name_str.clone(),
});
}
}
}
fn check_missing(entries: &[TypeEntry], errors: &mut Vec<CompileValidationError>) {
let names: HashSet<&str> = entries.iter().map(|e| e.name_str.as_str()).collect();
for entry in entries {
for dep in &entry.dependencies {
if !names.contains(dep.as_str()) {
errors.push(CompileValidationError::MissingDependency {
source: entry.name_str.clone(),
missing: dep.clone(),
});
}
}
}
}
fn check_cycles(entries: &[TypeEntry], errors: &mut Vec<CompileValidationError>) {
let mut visited = HashSet::new();
let mut in_stack = HashSet::new();
let mut path = Vec::new();
for entry in entries {
if !visited.contains(entry.name_str.as_str()) {
dfs_cycle(
entry,
entries,
&mut visited,
&mut in_stack,
&mut path,
errors,
);
}
}
}
fn dfs_cycle<'a>(
current: &'a TypeEntry,
entries: &'a [TypeEntry],
visited: &mut HashSet<&'a str>,
in_stack: &mut HashSet<&'a str>,
path: &mut Vec<&'a str>,
errors: &mut Vec<CompileValidationError>,
) {
visited.insert(¤t.name_str);
in_stack.insert(¤t.name_str);
path.push(¤t.name_str);
for dep_name in ¤t.dependencies {
let dep_entry = entries.iter().find(|e| e.name_str == *dep_name);
if let Some(dep) = dep_entry {
if !visited.contains(dep.name_str.as_str()) {
dfs_cycle(dep, entries, visited, in_stack, path, errors);
} else if in_stack.contains(dep.name_str.as_str()) {
let cycle_start = path
.iter()
.position(|n| *n == dep_name.as_str())
.unwrap_or(0);
let chain: Vec<String> = path
.get(cycle_start..)
.unwrap_or(&[])
.iter()
.map(|s| s.to_string())
.chain(std::iter::once(dep_name.clone()))
.collect();
errors.push(CompileValidationError::CircularDependency { chain });
}
}
}
path.pop();
in_stack.remove(current.name_str.as_str());
}
fn check_scope_mismatches(entries: &[TypeEntry], errors: &mut Vec<CompileValidationError>) {
for entry in entries {
for dep_name in &entry.dependencies {
if let Some(dep) = entries.iter().find(|e| e.name_str == *dep_name) {
if is_wider_scope(&entry.scope, &dep.scope) {
errors.push(CompileValidationError::ScopeMismatch {
source: entry.name_str.clone(),
source_scope: entry.scope.clone(),
dependency: dep_name.clone(),
dependency_scope: dep.scope.clone(),
});
}
}
}
}
}
fn is_wider_scope(source_scope: &str, dep_scope: &str) -> bool {
matches!((source_scope, dep_scope), ("singleton", "transient"))
}