use std::{
collections::{BTreeMap, BTreeSet},
sync::{Mutex, OnceLock},
};
use miden_note_schema::{NoteStorageSchema, SchemaCase, SchemaType, SchemaTypeKind};
use miden_note_schema_codegen::generated_type_ident;
use proc_macro2::Span;
#[derive(Clone, Debug, Eq, PartialEq)]
pub(crate) struct CodecRegistration {
pub(crate) fqn: String,
pub(crate) rust_type: String,
}
#[derive(Clone, Debug)]
struct RegisteredCodec {
registration: CodecRegistration,
location: ExpansionLocation,
}
#[derive(Default)]
struct Registry {
schema: Option<(String, ExpansionLocation)>,
types: BTreeMap<String, String>,
codecs: BTreeMap<String, RegisteredCodec>,
}
static REGISTRY: OnceLock<Mutex<BTreeMap<String, Registry>>> = OnceLock::new();
type ExpansionLocation = (String, usize, usize);
fn expansion_location(span: Span) -> ExpansionLocation {
let start = span.start();
(span.file(), start.line, start.column)
}
fn macro_invocation_crate_key() -> String {
let package = std::env::var("CARGO_MANIFEST_DIR")
.or_else(|_| std::env::var("CARGO_PKG_NAME"))
.unwrap_or_default();
let crate_name = std::env::var("CARGO_CRATE_NAME").unwrap_or_default();
format!("{package}\u{1f}{crate_name}")
}
fn registry() -> &'static Mutex<BTreeMap<String, Registry>> {
REGISTRY.get_or_init(|| Mutex::new(BTreeMap::new()))
}
pub(crate) fn register_schema(schema: &NoteStorageSchema, span: Span) -> syn::Result<()> {
let mut bindings = BTreeMap::new();
collect_type_bindings(schema.root(), &mut BTreeSet::new(), &mut bindings)?;
let bindings = index_by_rust_type_name(&bindings, span)?;
let mut registries = registry()
.lock()
.map_err(|_| syn::Error::new(span, "note codec registry mutex is poisoned"))?;
let registry = registries.entry(macro_invocation_crate_key()).or_default();
let location = expansion_location(span);
match ®istry.schema {
Some((existing, _)) if existing == schema.wit_text() => return Ok(()),
Some((_, existing_location)) if *existing_location == location => {
registry.codecs.clear();
}
Some(_) => {
return Err(syn::Error::new(
span,
"miden-note-codec supports one note schema per crate; remove the second distinct \
from_project!, from_package!, or from_wit_text! invocation. If this error \
appears in your IDE after an edit, restart the rust-analyzer proc-macro server",
));
}
None => {}
}
registry.schema = Some((schema.wit_text().to_owned(), location));
registry.types = bindings;
Ok(())
}
pub(crate) fn register_codec(rust_name: &str, rust_type: String, span: Span) -> syn::Result<()> {
let mut registries = registry()
.lock()
.map_err(|_| syn::Error::new(span, "note codec registry mutex is poisoned"))?;
let registry = registries.entry(macro_invocation_crate_key()).or_default();
let fqn = registry.types.get(rust_name).ok_or_else(|| {
syn::Error::new(
span,
format!(
"type `{rust_name}` is not part of a registered note schema; invoke \
miden_note_codec::from_project! or from_package! before #[note_codec]"
),
)
})?;
let fqn = fqn.clone();
let location = expansion_location(span);
let registration = CodecRegistration {
fqn: fqn.clone(),
rust_type,
};
if let Some(existing) = registry.codecs.get(&fqn)
&& existing.registration != registration
&& existing.location != location
{
return Err(syn::Error::new(
span,
format!(
"WIT type `{fqn}` already has a different #[note_codec] implementation. If this \
error appears in your IDE after an edit, restart the rust-analyzer proc-macro \
server"
),
));
}
registry.codecs.insert(
fqn,
RegisteredCodec {
registration,
location,
},
);
Ok(())
}
pub(crate) fn registered_codecs(span: Span) -> syn::Result<Vec<CodecRegistration>> {
let mut registries = registry()
.lock()
.map_err(|_| syn::Error::new(span, "note codec registry mutex is poisoned"))?;
let registry = registries.entry(macro_invocation_crate_key()).or_default();
if registry.codecs.is_empty() {
return Err(syn::Error::new(
span,
"export_codecs! found no registered codecs; place it after from_project! or \
from_package! and after every #[note_codec] implementation because procedural macros \
register as they expand; supported types are exported in canonical FQN order",
));
}
Ok(registry.codecs.values().map(|codec| codec.registration.clone()).collect())
}
fn index_by_rust_type_name(
bindings: &BTreeMap<String, String>,
span: Span,
) -> syn::Result<BTreeMap<String, String>> {
let mut index = BTreeMap::new();
for (fqn, rust_name) in bindings {
if let Some(existing) = index.insert(rust_name.clone(), fqn.clone()) {
return Err(syn::Error::new(
span,
format!("WIT types `{existing}` and `{fqn}` both map to Rust type `{rust_name}`"),
));
}
}
Ok(index)
}
fn collect_type_bindings(
ty: &SchemaType,
seen: &mut BTreeSet<String>,
bindings: &mut BTreeMap<String, String>,
) -> syn::Result<()> {
if ty.standard_leaf().is_some() {
return Ok(());
}
if matches!(ty.kind(), SchemaTypeKind::Record(_) | SchemaTypeKind::Variant(_)) {
let fqn = ty.fqn().ok_or_else(|| {
syn::Error::new(Span::call_site(), "a generated note codec type has no WIT FQN")
})?;
let name = ty.name().ok_or_else(|| {
syn::Error::new(
Span::call_site(),
format!("generated WIT type `{fqn}` has no local name"),
)
})?;
if !seen.insert(fqn.to_owned()) {
return Ok(());
}
bindings.insert(fqn.to_owned(), generated_type_ident(name));
}
match ty.kind() {
SchemaTypeKind::Record(fields) => {
for field in fields {
collect_type_bindings(field.ty(), seen, bindings)?;
}
}
SchemaTypeKind::Option(payload) => collect_type_bindings(payload, seen, bindings)?,
SchemaTypeKind::Variant(cases) => {
for payload in cases.iter().filter_map(SchemaCase::payload) {
collect_type_bindings(payload, seen, bindings)?;
}
}
SchemaTypeKind::Felt | SchemaTypeKind::Primitive(_) => {}
}
Ok(())
}
#[cfg(test)]
pub(crate) fn reset_for_tests() {
if let Some(registry) = REGISTRY.get() {
registry.lock().expect("mutex poisoned").clear();
}
}