From 6ded6e4aaa00cd48f604b40286f842575e2b2666 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Tue, 26 May 2026 17:00:48 -0700 Subject: [PATCH 01/21] Checkpoint. --- Cargo.lock | 10 + Cargo.toml | 2 + diskann-benchmark-runner-derive/Cargo.toml | 19 ++ diskann-benchmark-runner-derive/src/lib.rs | 106 ++++++++ diskann-benchmark-runner/Cargo.toml | 1 + diskann-benchmark-runner/src/lib.rs | 5 + diskann-benchmark-runner/src/reflect.rs | 269 +++++++++++++++++++++ 7 files changed, 412 insertions(+) create mode 100644 diskann-benchmark-runner-derive/Cargo.toml create mode 100644 diskann-benchmark-runner-derive/src/lib.rs create mode 100644 diskann-benchmark-runner/src/reflect.rs diff --git a/Cargo.lock b/Cargo.lock index 64a8c6b399..61a26655d4 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -703,6 +703,7 @@ version = "0.52.0" dependencies = [ "anyhow", "clap", + "diskann-benchmark-runner-derive", "half", "indicatif", "serde", @@ -711,6 +712,15 @@ dependencies = [ "thiserror 2.0.17", ] +[[package]] +name = "diskann-benchmark-runner-derive" +version = "0.52.0" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.113", +] + [[package]] name = "diskann-benchmark-simd" version = "0.52.0" diff --git a/Cargo.toml b/Cargo.toml index c9f1fb8b23..5f39abcfd2 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -25,6 +25,7 @@ members = [ "diskann-tools", "vectorset", "diskann-bftree", + "diskann-benchmark-runner-derive", ] default-members = [ @@ -63,6 +64,7 @@ diskann-disk = { path = "diskann-disk", version = "0.52.0" } diskann-label-filter = { path = "diskann-label-filter", version = "0.52.0" } # Infra diskann-benchmark-runner = { path = "diskann-benchmark-runner", version = "0.52.0" } +diskann-benchmark-runner-derive = { path = "diskann-benchmark-runner-derive", version = "0.52.0" } diskann-benchmark-core = { path = "diskann-benchmark-core", version = "0.52.0" } diskann-tools = { path = "diskann-tools", version = "0.52.0" } diff --git a/diskann-benchmark-runner-derive/Cargo.toml b/diskann-benchmark-runner-derive/Cargo.toml new file mode 100644 index 0000000000..c9a4525f08 --- /dev/null +++ b/diskann-benchmark-runner-derive/Cargo.toml @@ -0,0 +1,19 @@ +[package] +name = "diskann-benchmark-runner-derive" +version.workspace = true +description.workspace = true +authors.workspace = true +documentation.workspace = true +license.workspace = true +edition = "2024" + +[lib] +proc-macro = true + +[dependencies] +syn = { version = "2", features = ["full"] } +quote = "1" +proc-macro2 = "1" + +[lints] +workspace = true diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs new file mode 100644 index 0000000000..a27b21137e --- /dev/null +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -0,0 +1,106 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use proc_macro::TokenStream; +use quote::{quote, quote_spanned}; +use syn::{Data, DeriveInput, Fields, parse_macro_input, spanned::Spanned}; + +fn crate_name() -> syn::Path { + syn::parse_quote!(::diskann_benchmark_runner::reflect) +} + +/// Derive macro for the `Reflect` trait. +/// +/// Supports named structs, tuple structs, and enums. Doc comments on the item +/// and its fields/variants are captured as reflection metadata. +/// +/// # Example +/// +/// ```ignore +/// use diskann_benchmark_runner::reflect::Reflect; +/// +/// /// A test aggregate. +/// #[derive(Reflect)] +/// struct MyInput { +/// /// The number of threads. +/// threads: usize, +/// } +/// ``` +#[proc_macro_derive(Reflect, attributes(reflect))] +pub fn derive_reflect(input: TokenStream) -> TokenStream { + let input = parse_macro_input!(input as DeriveInput); + let output = match &input.data { + Data::Struct(s) => process_struct(&input, s), + _ => todo!("need to figure this out"), + }; + output.into() +} + +fn process_struct(input: &DeriveInput, s: &syn::DataStruct) -> proc_macro2::TokenStream { + let type_name = &input.ident; + let type_name_str = type_name.to_string(); + let path = crate_name(); + let doc = format_docstrings(&input.attrs); + + match &s.fields { + Fields::Named(named) => { + let fields = named.named.iter().map(|f| { + let ty = &f.ty; + let ident = f.ident.as_ref().unwrap().to_string(); + let doc = format_docstrings(&f.attrs); + + quote_spanned! { ty.span()=> #path::Field::new::<#ty>(#ident, #doc) } + }); + + quote! { + impl #path::Reflect for #type_name { + fn reflect() -> #path::Type { + #path::Type::aggregate( + #type_name_str, + [#(#fields),*], + #doc, + ) + } + } + } + } + _ => todo!("more todos"), + } +} + +fn format_docstrings(attributes: &[syn::Attribute]) -> proc_macro2::TokenStream { + match extract_docs(attributes) { + None => quote!{ ::std::option::Option::None }, + Some(docs) => quote! { ::std::option::Option::Some(#docs.into()) }, + } +} + +fn extract_docs(attributes: &[syn::Attribute]) -> Option { + let docstrings = attributes + .iter() + .filter_map(|a| { + if a.path().is_ident("doc") + && let syn::Meta::NameValue(name) = &a.meta + && let syn::Expr::Lit(literal) = &name.value + && let syn::Lit::Str(s) = &literal.lit + { + let value = s.value(); + let processed = match value.strip_prefix(" ") { + Some(stripped) => stripped.to_owned(), + None => value, + }; + Some(processed) + } else { + None + } + }) + .collect::>(); + + if docstrings.is_empty() { + None + } else { + Some(docstrings.join("\n")) + } +} diff --git a/diskann-benchmark-runner/Cargo.toml b/diskann-benchmark-runner/Cargo.toml index 33cb63d0d3..2ee9b0775c 100644 --- a/diskann-benchmark-runner/Cargo.toml +++ b/diskann-benchmark-runner/Cargo.toml @@ -10,6 +10,7 @@ edition.workspace = true [dependencies] anyhow = { workspace = true } clap = { workspace = true, features = ["derive"] } +diskann-benchmark-runner-derive = { workspace = true } half = { workspace = true } indicatif = "0.18.3" serde = { workspace = true, features = ["derive"] } diff --git a/diskann-benchmark-runner/src/lib.rs b/diskann-benchmark-runner/src/lib.rs index 724a827f66..effb270a8e 100644 --- a/diskann-benchmark-runner/src/lib.rs +++ b/diskann-benchmark-runner/src/lib.rs @@ -5,6 +5,9 @@ //! A moderately functional utility for making simple benchmarking CLI applications. +#[doc(hidden)] +extern crate self as diskann_benchmark_runner; + pub mod benchmark; mod checker; mod internal; @@ -15,6 +18,7 @@ pub mod app; pub mod files; pub mod input; pub mod output; +pub mod reflect; pub mod registry; pub mod utils; @@ -23,6 +27,7 @@ pub use benchmark::Benchmark; pub use checker::Checker; pub use input::Input; pub use output::Output; +pub use reflect::Reflect; pub use registry::{Registry, RegistryError}; pub use result::Checkpoint; diff --git a/diskann-benchmark-runner/src/reflect.rs b/diskann-benchmark-runner/src/reflect.rs new file mode 100644 index 0000000000..c9a375ba8d --- /dev/null +++ b/diskann-benchmark-runner/src/reflect.rs @@ -0,0 +1,269 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::borrow::Cow; + +pub use diskann_benchmark_runner_derive::Reflect; + +pub trait Reflect: 'static { + fn reflect() -> Type; +} + +pub fn reflect() -> Reflection +where + T: Reflect, +{ + Reflection::new::() +} + +#[derive(Clone, Copy)] +pub struct Reflection { + reflection: &'static dyn internal::Reflect, +} + +impl Reflection { + pub const fn new() -> Self + where + T: Reflect, + { + Self { + reflection: &internal::Wrapper::::INSTANCE, + } + } + + pub fn reflect(&self) -> Type { + self.reflection.reflect() + } +} + +pub type Doc = Cow<'static, str>; + +pub enum Type { + Primitive(Primitive), + Aggregate(Aggregate), + Enum(Enum), +} + +impl Type { + pub fn primitive(type_name: &'static str, doc: &'static str) -> Self { + Self::from(Primitive::new(type_name, doc)) + } + + pub fn aggregate(type_name: &'static str, fields: Itr, doc: Option) -> Self + where + Itr: IntoIterator, + { + Self::from(Aggregate::new(type_name, fields, doc)) + } + + pub(crate) fn walk(&self, name: &str) -> Result { + match self { + Self::Primitive(_) => Err(WalkError), + Self::Aggregate(aggregate) => aggregate.walk(name), + Self::Enum(variant) => variant.walk(name), + } + } +} + +impl From for Type { + fn from(primitive: Primitive) -> Self { + Self::Primitive(primitive) + } +} + +impl From for Type { + fn from(aggergate: Aggregate) -> Self { + Self::Aggregate(aggergate) + } +} + +impl From for Type { + fn from(e: Enum) -> Self { + Self::Enum(e) + } +} + +pub struct Primitive { + type_name: &'static str, + doc: &'static str, +} + +impl Primitive { + pub fn new(type_name: &'static str, doc: &'static str) -> Self { + Self { type_name, doc } + } +} + +pub struct Aggregate { + type_name: &'static str, + fields: Vec, + doc: Option, +} + +impl Aggregate { + pub fn new(type_name: &'static str, fields: Itr, doc: Option) -> Self + where + Itr: IntoIterator, + { + Self { + type_name, + fields: fields.into_iter().collect(), + doc, + } + } + + fn walk(&self, field: &str) -> Result { + match self.fields.iter().find(|f| f.name == field) { + Some(f) => Ok(f.field), + None => Err(WalkError), + } + } +} + +pub struct Field { + name: &'static str, + field: Reflection, + doc: Option, +} + +impl Field { + pub fn new(name: &'static str, doc: Option) -> Self + where + T: Reflect, + { + Self { + name, + field: reflect::(), + doc, + } + } +} + +pub struct Enum { + type_name: &'static str, + variants: Vec<(&'static str, Variant)>, + doc: Option, +} + +impl Enum { + pub fn new(type_name: &'static str, variants: Itr, doc: Option) -> Self + where + Itr: IntoIterator, + { + Self { + type_name, + variants: variants.into_iter().collect(), + doc, + } + } + + fn walk(&self, variant: &str) -> Result { + match self.variants.iter().find(|(v, _)| *v == variant) { + Some(v) => Ok(v.1.variant), + None => Err(WalkError), + } + } +} + +pub struct Variant { + variant: Option, + doc: Option, +} + +impl Variant { + pub fn new(doc: Option) -> Self + where + T: Reflect, + { + Self { + variant: Some(reflect::()), + doc, + } + } + + pub fn aggregate(doc: Option) -> Self { + + } +} + +//////////////// +// Algorithms // +//////////////// + +pub fn walk<'a, I>(reflection: Reflection, paths: I) -> Result +where + I: IntoIterator, +{ + let mut current = reflection; + for p in paths { + current = current.reflect().walk(p)?; + } + Ok(current) +} + +#[derive(Debug, Clone, Copy)] +pub struct WalkError; + +/////////////// +// Bootstrap // +/////////////// + +impl Reflect for usize { + fn reflect() -> Type { + Type::primitive("usize", "An system dependent unsigned integer") + } +} + +/// This is a test! +/// +/// Hello world! +#[derive(Reflect)] +struct Test { + /// This field affects this value. + a: usize, + + /// This field does something else. + b: usize, +} + +// impl Reflect for Test { +// fn reflect() -> Type { +// Type::aggregate( +// "Test", +// [ +// Field::new::("a", Some("this fields does a thing".into())), +// Field::new::("b", Some("this fields does another thing".into())), +// ], +// None, +// ) +// } +// } + +pub(crate) mod internal { + use std::marker::PhantomData; + + pub(crate) trait Reflect { + fn reflect(&self) -> super::Type; + } + + pub(crate) struct Wrapper(PhantomData); + + impl Wrapper { + pub(crate) const INSTANCE: Self = Self::new(); + + pub(crate) const fn new() -> Self { + Self(PhantomData) + } + } + + impl Reflect for Wrapper + where + T: super::Reflect, + { + fn reflect(&self) -> super::Type { + ::reflect() + } + } +} From 7a16b930b6fa84153ad8deb066436efdd15575f0 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Thu, 11 Jun 2026 13:47:02 -0700 Subject: [PATCH 02/21] Checkpoint. --- diskann-benchmark-runner-derive/src/lib.rs | 6 +- diskann-benchmark-runner/src/reflect.rs | 131 ++++++++++++--------- 2 files changed, 76 insertions(+), 61 deletions(-) diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index a27b21137e..cefdce2928 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -51,7 +51,7 @@ fn process_struct(input: &DeriveInput, s: &syn::DataStruct) -> proc_macro2::Toke let ident = f.ident.as_ref().unwrap().to_string(); let doc = format_docstrings(&f.attrs); - quote_spanned! { ty.span()=> #path::Field::new::<#ty>(#ident, #doc) } + quote_spanned! { ty.span()=> #path::NamedField::new::<#ty>(#ident, #doc) } }); quote! { @@ -59,7 +59,7 @@ fn process_struct(input: &DeriveInput, s: &syn::DataStruct) -> proc_macro2::Toke fn reflect() -> #path::Type { #path::Type::aggregate( #type_name_str, - [#(#fields),*], + #path::Fields::Named(vec![#(#fields),*]), #doc, ) } @@ -72,7 +72,7 @@ fn process_struct(input: &DeriveInput, s: &syn::DataStruct) -> proc_macro2::Toke fn format_docstrings(attributes: &[syn::Attribute]) -> proc_macro2::TokenStream { match extract_docs(attributes) { - None => quote!{ ::std::option::Option::None }, + None => quote! { ::std::option::Option::None }, Some(docs) => quote! { ::std::option::Option::Some(#docs.into()) }, } } diff --git a/diskann-benchmark-runner/src/reflect.rs b/diskann-benchmark-runner/src/reflect.rs index c9a375ba8d..d575218039 100644 --- a/diskann-benchmark-runner/src/reflect.rs +++ b/diskann-benchmark-runner/src/reflect.rs @@ -51,20 +51,24 @@ impl Type { Self::from(Primitive::new(type_name, doc)) } - pub fn aggregate(type_name: &'static str, fields: Itr, doc: Option) -> Self - where - Itr: IntoIterator, - { + pub fn aggregate(type_name: &'static str, fields: Fields, doc: Option) -> Self { Self::from(Aggregate::new(type_name, fields, doc)) } - pub(crate) fn walk(&self, name: &str) -> Result { - match self { - Self::Primitive(_) => Err(WalkError), - Self::Aggregate(aggregate) => aggregate.walk(name), - Self::Enum(variant) => variant.walk(name), - } + pub fn enum_(type_name: &'static str, variants: Itr, doc: Option) -> Self + where + Itr: IntoIterator, + { + Self::from(Enum::new(type_name, variants, doc)) } + + // pub(crate) fn walk(&self, name: &str) -> Result { + // match self { + // Self::Primitive(_) => Err(WalkError), + // Self::Aggregate(aggregate) => aggregate.walk(name), + // Self::Enum(variant) => variant.walk(name), + // } + // } } impl From for Type { @@ -98,37 +102,40 @@ impl Primitive { pub struct Aggregate { type_name: &'static str, - fields: Vec, + fields: Fields, doc: Option, } impl Aggregate { - pub fn new(type_name: &'static str, fields: Itr, doc: Option) -> Self - where - Itr: IntoIterator, - { + pub fn new(type_name: &'static str, fields: Fields, doc: Option) -> Self { Self { type_name, - fields: fields.into_iter().collect(), + fields, doc, } } - fn walk(&self, field: &str) -> Result { - match self.fields.iter().find(|f| f.name == field) { - Some(f) => Ok(f.field), - None => Err(WalkError), - } - } + // fn walk(&self, field: &str) -> Result { + // match self.fields.iter().find(|f| f.name == field) { + // Some(f) => Ok(f.field), + // None => Err(WalkError), + // } + // } +} + +pub enum Fields { + Named(Vec), + Unnamed(Vec), + Unit, } -pub struct Field { +pub struct NamedField { name: &'static str, field: Reflection, doc: Option, } -impl Field { +impl NamedField { pub fn new(name: &'static str, doc: Option) -> Self where T: Reflect, @@ -141,16 +148,33 @@ impl Field { } } +pub struct UnnamedField { + field: Reflection, + doc: Option, +} + +impl UnnamedField { + pub fn new(doc: Option) -> Self + where + T: Reflect, + { + Self { + field: reflect::(), + doc, + } + } +} + pub struct Enum { type_name: &'static str, - variants: Vec<(&'static str, Variant)>, + variants: Vec, doc: Option, } impl Enum { pub fn new(type_name: &'static str, variants: Itr, doc: Option) -> Self where - Itr: IntoIterator, + Itr: IntoIterator, { Self { type_name, @@ -159,32 +183,23 @@ impl Enum { } } - fn walk(&self, variant: &str) -> Result { - match self.variants.iter().find(|(v, _)| *v == variant) { - Some(v) => Ok(v.1.variant), - None => Err(WalkError), - } - } + // fn walk(&self, variant: &str) -> Result { + // match self.variants.iter().find(|(v, _)| *v == variant) { + // Some(v) => Ok(v.1.variant), + // None => Err(WalkError), + // } + // } } pub struct Variant { - variant: Option, + name: &'static str, + fields: Fields, doc: Option, } impl Variant { - pub fn new(doc: Option) -> Self - where - T: Reflect, - { - Self { - variant: Some(reflect::()), - doc, - } - } - - pub fn aggregate(doc: Option) -> Self { - + pub fn new(name: &'static str, fields: Fields, doc: Option) -> Self { + Self { name, fields, doc } } } @@ -192,19 +207,19 @@ impl Variant { // Algorithms // //////////////// -pub fn walk<'a, I>(reflection: Reflection, paths: I) -> Result -where - I: IntoIterator, -{ - let mut current = reflection; - for p in paths { - current = current.reflect().walk(p)?; - } - Ok(current) -} - -#[derive(Debug, Clone, Copy)] -pub struct WalkError; +// pub fn walk<'a, I>(reflection: Reflection, paths: I) -> Result +// where +// I: IntoIterator, +// { +// let mut current = reflection; +// for p in paths { +// current = current.reflect().walk(p)?; +// } +// Ok(current) +// } +// +// #[derive(Debug, Clone, Copy)] +// pub struct WalkError; /////////////// // Bootstrap // From 6f735df33bcae86c27c24807da92285b705d0f4a Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Fri, 25 Sep 2026 10:29:50 -0700 Subject: [PATCH 03/21] Checkpoint. --- diskann-benchmark-runner-derive/src/lib.rs | 5 +- diskann-benchmark-runner/src/reflect.rs | 369 ++++++++++++++------- diskann-benchmark-runner/temp/main.rs | 5 +- 3 files changed, 266 insertions(+), 113 deletions(-) diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index cd3118e42a..022bc7c545 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -77,12 +77,15 @@ fn process_struct(input: &DeriveInput, s: &syn::DataStruct) -> proc_macro2::Toke impl #impl_generics #path::Reflect for #type_name #ty_generics #where_clause { fn reflect() -> #path::Type { #path::Type::aggregate( - #type_name_str, #type_id, #path::Fields::Named(vec![#(#fields),*]), #doc, ) } + + fn type_name(f: &mut dyn ::std::fmt::Write) -> ::std::fmt::Result { + f.write_str(#type_name_str) + } } } } diff --git a/diskann-benchmark-runner/src/reflect.rs b/diskann-benchmark-runner/src/reflect.rs index af84387afd..824a95cda1 100644 --- a/diskann-benchmark-runner/src/reflect.rs +++ b/diskann-benchmark-runner/src/reflect.rs @@ -17,6 +17,7 @@ const INDENT: usize = 2; pub trait Reflect: 'static { fn reflect() -> Type; + fn type_name(f: &mut dyn Write) -> fmt::Result; } pub fn reflect() -> Reflection @@ -44,57 +45,57 @@ impl Reflection { pub fn reflect(&self) -> Type { self.reflection.reflect() } + + pub fn type_name(&self) -> TypeName { + TypeName(*self) + } + + pub fn render(&self) -> Render { + Render(*self) + } } pub type Doc = Cow<'static, str>; -#[derive(Debug)] -pub enum Query<'a> { - Field(Cow<'a, str>), - Index(usize), +pub struct TypeName(Reflection); + +impl std::fmt::Display for TypeName { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.reflection.type_name(f) + } +} + +pub struct Render(Reflection); + +impl std::fmt::Display for Render { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut r = Renderer::new(f, 2); + r.render_reflection(self.0) + } } pub enum Type { Primitive(Primitive), Aggregate(Aggregate), Enum(Enum), + Sequence(Sequence), } impl Type { - pub fn primitive(type_name: &'static str, type_id: TypeId, doc: &'static str) -> Self { - Self::from(Primitive::new(type_name, type_id, doc)) + pub fn primitive(type_id: TypeId, doc: Option) -> Self { + Self::from(Primitive::new(type_id, doc)) } - pub fn aggregate( - type_name: &'static str, - type_id: TypeId, - fields: Fields, - doc: Option, - ) -> Self { - Self::from(Aggregate::new(type_name, type_id, fields, doc)) + pub fn aggregate(type_id: TypeId, fields: Fields, doc: Option) -> Self { + Self::from(Aggregate::new(type_id, fields, doc)) } pub fn enum_( - type_name: &'static str, type_id: TypeId, variants: impl IntoIterator, doc: Option, ) -> Self { - Self::from(Enum::new(type_name, type_id, variants, doc)) - } - - fn format_into(&self, f: &mut dyn Write) -> fmt::Result { - match self { - Self::Primitive(p) => p.format_into(f), - Self::Aggregate(a) => a.format_into(f), - Self::Enum(e) => e.format_into(f), - } - } -} - -impl fmt::Display for Type { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.format_into(f) + Self::from(Enum::new(type_id, variants, doc)) } } @@ -116,52 +117,52 @@ impl From for Type { } } +impl From for Type { + fn from(s: Sequence) -> Self { + Self::Sequence(s) + } +} + +//-----------// +// Primitive // +//-----------// + pub struct Primitive { - type_name: &'static str, type_id: TypeId, - doc: &'static str, + doc: Option, } impl Primitive { - pub fn new(type_name: &'static str, type_id: TypeId, doc: &'static str) -> Self { - Self { - type_name, - type_id, - doc, - } + pub fn new(type_id: TypeId, doc: Option) -> Self { + Self { type_id, doc } } - fn format_into(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result { - write!(f, "{}: {}", self.type_name, self.doc) + fn doc(&self) -> Option<&str> { + self.doc.as_deref() } } +//-----------// +// Aggregate // +//-----------// + pub struct Aggregate { - type_name: &'static str, type_id: TypeId, fields: Fields, doc: Option, } impl Aggregate { - pub fn new(type_name: &'static str, type_id: TypeId, fields: Fields, doc: Option) -> Self { + pub fn new(type_id: TypeId, fields: Fields, doc: Option) -> Self { Self { - type_name, type_id, fields, doc, } } - fn format_into(&self, f: &mut dyn Write) -> fmt::Result { - f.write_str(self.type_name)?; - if let Some(doc) = &self.doc { - write!(f, "\n{}\n", Indent::new(&doc, INDENT))?; - } - - let mut scratch = String::new(); - self.fields.format_into(&mut scratch)?; - write!(f, "{}", Indent::new(&scratch, INDENT)) + fn doc(&self) -> Option<&str> { + self.doc.as_deref() } } @@ -171,30 +172,6 @@ pub enum Fields { Unit, } -impl Fields { - fn format_into(&self, f: &mut dyn Write) -> fmt::Result { - match self { - Self::Named(fields) => { - for field in fields.iter() { - field.format_into(f)?; - f.write_str("\n\n")?; - } - } - Self::Unnamed(fields) => { - let mut buf = String::new(); - for (i, field) in fields.iter().enumerate() { - write!(f, "{}", i)?; - buf.clear(); - field.format_into(&mut buf)?; - write!(f, "{}", Indent::new(&buf, INDENT))?; - } - } - Self::Unit => {} - } - Ok(()) - } -} - pub struct NamedField { name: &'static str, field: Reflection, @@ -213,15 +190,8 @@ impl NamedField { } } - fn format_into(&self, f: &mut dyn Write) -> fmt::Result { - write!(f, "{}", Quote(self.name))?; - if let Some(doc) = &self.doc { - write!(f, "\n{}\n", Indent::new(doc, INDENT))?; - } - - let mut buf = String::new(); - self.field.reflect().format_into(&mut buf)?; - write!(f, "{}", Indent::new(&buf, INDENT)) + fn doc(&self) -> Option<&str> { + self.doc.as_deref() } } @@ -240,14 +210,13 @@ impl UnnamedField { doc, } } - - fn format_into(&self, f: &mut dyn Write) -> fmt::Result { - self.field.reflect().format_into(f) - } } +//------// +// Enum // +//------// + pub struct Enum { - type_name: &'static str, type_id: TypeId, variants: Vec, doc: Option, @@ -255,33 +224,16 @@ pub struct Enum { impl Enum { pub fn new( - type_name: &'static str, type_id: TypeId, variants: impl IntoIterator, doc: Option, ) -> Self { Self { - type_name, type_id: type_id, variants: variants.into_iter().collect(), doc, } } - - fn format_into(&self, f: &mut dyn Write) -> fmt::Result { - f.write_str(self.type_name)?; - if let Some(doc) = &self.doc { - write!(f, "{}", Indent::new(&doc, INDENT))?; - } - - let mut buf = String::new(); - for variant in self.variants.iter() { - variant.format_into(&mut buf)?; - write!(f, "{}", Indent::new(&buf, INDENT))?; - } - - Ok(()) - } } pub struct Variant { @@ -294,10 +246,179 @@ impl Variant { pub fn new(name: &'static str, fields: Fields, doc: Option) -> Self { Self { name, fields, doc } } +} + +//----------// +// Sequence // +//----------// + +pub struct Sequence { + type_id: TypeId, + element: Reflection, + doc: Option, +} + +impl Sequence { + pub fn new(type_id: TypeId, doc: Option) -> Self + where + T: Reflect, + { + Self { + type_id, + element: reflect::(), + doc, + } + } +} + +////////////// +// Renderer // +////////////// + +struct Renderer<'a> { + output: &'a mut dyn Write, + indent: usize, + depth: usize, + max_depth: usize, +} + +impl<'a> Renderer<'a> { + fn new(output: &'a mut dyn Write, max_depth: usize) -> Self { + Self { + output, + indent: 0, + depth: 0, + max_depth, + } + } + + fn at_top(&self) -> bool { + self.indent == 0 + } + + fn at_bottom(&self) -> bool { + self.depth == self.max_depth + } + + fn line(&mut self, display: D) -> fmt::Result + where + D: std::fmt::Display, + { + let indent = INDENT * self.indent; + write!(self.output, "{: >indent$}{}\n", "", display) + } + + fn blank(&mut self) -> fmt::Result { + self.output.write_char('\n') + } + + fn maybe_indent(&mut self, indent: bool, f: F) -> fmt::Result + where + F: FnOnce(&mut Self) -> fmt::Result, + { + if indent { + self.indent += 1; + } + let result = f(self); + if indent { + self.indent -= 1; + } + result + } + + fn indent(&mut self, f: F) -> fmt::Result + where + F: FnOnce(&mut Self) -> fmt::Result, + { + self.maybe_indent(true, f) + } + + fn next(&mut self, f: F) -> fmt::Result + where + F: FnOnce(&mut Self) -> fmt::Result, + { + if self.at_bottom() { + Ok(()) + } else { + self.depth += 1; + let result = self.indent(f); + self.depth -= 1; + result + } + } + + fn render_doc(&mut self, s: Option<&str>) -> fmt::Result { + if let Some(s) = s { + for ln in s.lines() { + if ln.is_empty() { + self.blank()?; + } else { + self.line(ln)?; + } + } + } + + Ok(()) + } + + //-------// + // Types // + //-------// + + fn render_reflection(&mut self, reflection: Reflection) -> fmt::Result { + // Render the type name if this is the first item in the stack. + if self.at_top() { + self.line(reflection.type_name())?; + } + + match reflection.reflect() { + Type::Primitive(primitive) => self.render_primitive(&primitive), + Type::Aggregate(aggregate) => self.render_aggregate(&aggregate), + Type::Enum(enum_) => todo!(), + Type::Sequence(sequence) => todo!(), + } + } + + fn render_primitive(&mut self, primitive: &Primitive) -> fmt::Result { + if self.at_top() { + self.indent(|r| r.render_doc(primitive.doc()))?; + } + + Ok(()) + } - fn format_into(&self, f: &mut dyn Write) -> fmt::Result { - f.write_str(self.name)?; - self.fields.format_into(f) + fn render_aggregate(&mut self, aggregate: &Aggregate) -> fmt::Result { + self.maybe_indent(self.at_top(), |r| { + r.render_doc(aggregate.doc())?; + r.render_fields(&aggregate.fields) + }) + } + + //--------// + // Fields // + //--------// + + fn render_fields(&mut self, fields: &Fields) -> fmt::Result { + match fields { + Fields::Named(named) => { + for field in named.iter() { + self.render_named_field(field)?; + } + } + Fields::Unnamed(_) => todo!(), + Fields::Unit => todo!(), + } + + Ok(()) + } + + fn render_named_field(&mut self, field: &NamedField) -> fmt::Result { + let f = field.field; + + self.blank()?; + self.line(format_args!("{}: {}", Quote(field.name), TypeName(f)))?; + self.indent(|r| r.render_doc(field.doc()))?; + self.next(|r| r.render_reflection(f)) } } @@ -326,13 +447,25 @@ impl Variant { impl Reflect for usize { fn reflect() -> Type { Type::primitive( - "usize", TypeId::of::(), - "An system dependent unsigned integer", + Some("A system dependent unsigned integer".into()), ) } + + fn type_name(f: &mut dyn Write) -> fmt::Result { + f.write_str("usize") + } } +// impl Reflect for Vec +// where +// T: Reflect, +// { +// fn reflect() -> Type { +// +// } +// } + /// This is a test! /// /// Hello world! @@ -345,6 +478,15 @@ pub struct Test { b: usize, } +/// This is a nother test! +#[derive(Reflect)] +pub struct Test2 { + /// This field affects this value. + a: usize, + + other: Test, +} + #[derive(Reflect)] struct Wrapper { /// Inner @@ -356,6 +498,7 @@ pub(crate) mod internal { pub(crate) trait Reflect { fn reflect(&self) -> super::Type; + fn type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result; } pub(crate) struct Wrapper(PhantomData); @@ -375,5 +518,9 @@ pub(crate) mod internal { fn reflect(&self) -> super::Type { ::reflect() } + + fn type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result { + ::type_name(f) + } } } diff --git a/diskann-benchmark-runner/temp/main.rs b/diskann-benchmark-runner/temp/main.rs index 2d167b9d00..79186060f8 100644 --- a/diskann-benchmark-runner/temp/main.rs +++ b/diskann-benchmark-runner/temp/main.rs @@ -8,6 +8,9 @@ use diskann_benchmark_runner::reflect; fn main() -> anyhow::Result<()> { - println!("{}", reflect::reflect::().reflect()); + println!("{}", reflect::reflect::().render()); + + println!("{}", reflect::reflect::().render()); + Ok(()) } From 502e1547f2da0e8cf06fab4d3847af2bad2f4e9a Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Fri, 25 Sep 2026 13:08:28 -0700 Subject: [PATCH 04/21] Enum? --- diskann-benchmark-runner-derive/src/lib.rs | 5 - diskann-benchmark-runner/src/reflect.rs | 544 +++++++++++++++++---- diskann-benchmark-runner/temp/main.rs | 2 + 3 files changed, 444 insertions(+), 107 deletions(-) diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index 022bc7c545..b2cb2eb32b 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -69,15 +69,10 @@ fn process_struct(input: &DeriveInput, s: &syn::DataStruct) -> proc_macro2::Toke let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); - let type_id = quote_spanned! { - input.span()=> ::std::any::TypeId::of::<#type_name #ty_generics>() - }; - quote! { impl #impl_generics #path::Reflect for #type_name #ty_generics #where_clause { fn reflect() -> #path::Type { #path::Type::aggregate( - #type_id, #path::Fields::Named(vec![#(#fields),*]), #doc, ) diff --git a/diskann-benchmark-runner/src/reflect.rs b/diskann-benchmark-runner/src/reflect.rs index 824a95cda1..ec2cc85c55 100644 --- a/diskann-benchmark-runner/src/reflect.rs +++ b/diskann-benchmark-runner/src/reflect.rs @@ -11,7 +11,7 @@ use std::{ pub use diskann_benchmark_runner_derive::Reflect; -use crate::utils::fmt::{Indent, Quote}; +use crate::utils::fmt::Quote; const INDENT: usize = 2; @@ -50,18 +50,42 @@ impl Reflection { TypeName(*self) } + pub fn type_id(&self) -> TypeId { + self.reflection.type_id() + } + pub fn render(&self) -> Render { Render(*self) } } +impl fmt::Debug for Reflection { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("Reflection") + .field("type_name", &self.type_name()) + .finish_non_exhaustive() + } +} + pub type Doc = Cow<'static, str>; pub struct TypeName(Reflection); +impl TypeName { + fn format_into(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.reflection.type_name(f) + } +} + +impl std::fmt::Debug for TypeName { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.format_into(f) + } +} + impl std::fmt::Display for TypeName { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.reflection.type_name(f) + self.format_into(f) } } @@ -69,11 +93,12 @@ pub struct Render(Reflection); impl std::fmt::Display for Render { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - let mut r = Renderer::new(f, 2); - r.render_reflection(self.0) + let mut r = Renderer::new(f, 3); + r.render_subject(self.0) } } +#[derive(Debug)] pub enum Type { Primitive(Primitive), Aggregate(Aggregate), @@ -82,20 +107,45 @@ pub enum Type { } impl Type { - pub fn primitive(type_id: TypeId, doc: Option) -> Self { - Self::from(Primitive::new(type_id, doc)) + pub fn primitive(doc: Option) -> Self { + Self::from(Primitive::new(doc)) } - pub fn aggregate(type_id: TypeId, fields: Fields, doc: Option) -> Self { - Self::from(Aggregate::new(type_id, fields, doc)) + pub fn aggregate(fields: Fields, doc: Option) -> Self { + Self::from(Aggregate::new(fields, doc)) } pub fn enum_( - type_id: TypeId, + repr: EnumRepr, variants: impl IntoIterator, doc: Option, ) -> Self { - Self::from(Enum::new(type_id, variants, doc)) + Self::from(Enum::new(repr, variants, doc)) + } + + pub fn sequence(doc: Option) -> Self + where + T: Reflect, + { + Self::from(Sequence::new::(doc)) + } + + pub fn doc(&self) -> Option<&str> { + match self { + Self::Primitive(p) => p.doc(), + Self::Aggregate(a) => a.doc(), + Self::Enum(e) => e.doc(), + Self::Sequence(s) => s.doc(), + } + } + + fn has_body(&self) -> bool { + match self { + Self::Primitive(_) => false, + Self::Aggregate(a) => a.has_body(), + Self::Enum(e) => e.has_body(), + Self::Sequence(_) => true, + } } } @@ -127,14 +177,14 @@ impl From for Type { // Primitive // //-----------// +#[derive(Debug)] pub struct Primitive { - type_id: TypeId, doc: Option, } impl Primitive { - pub fn new(type_id: TypeId, doc: Option) -> Self { - Self { type_id, doc } + pub fn new(doc: Option) -> Self { + Self { doc } } fn doc(&self) -> Option<&str> { @@ -146,32 +196,52 @@ impl Primitive { // Aggregate // //-----------// +#[derive(Debug)] pub struct Aggregate { - type_id: TypeId, fields: Fields, doc: Option, } impl Aggregate { - pub fn new(type_id: TypeId, fields: Fields, doc: Option) -> Self { - Self { - type_id, - fields, - doc, - } + pub fn new(fields: Fields, doc: Option) -> Self { + Self { fields, doc } } fn doc(&self) -> Option<&str> { self.doc.as_deref() } + + fn has_body(&self) -> bool { + self.fields.has_body() + } } +#[derive(Debug)] pub enum Fields { Named(Vec), Unnamed(Vec), Unit, } +impl Fields { + fn has_body(&self) -> bool { + match self { + Self::Named(fields) => !fields.is_empty(), + Self::Unnamed(fields) => !fields.is_empty(), + Self::Unit => false, + } + } + + pub fn named(itr: impl IntoIterator) -> Self { + Self::Named(itr.into_iter().collect()) + } + + pub fn unnamed(itr: impl IntoIterator) -> Self { + Self::Unnamed(itr.into_iter().collect()) + } +} + +#[derive(Debug)] pub struct NamedField { name: &'static str, field: Reflection, @@ -195,6 +265,7 @@ impl NamedField { } } +#[derive(Debug)] pub struct UnnamedField { field: Reflection, doc: Option, @@ -210,32 +281,87 @@ impl UnnamedField { doc, } } + + fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } } //------// // Enum // //------// +#[derive(Debug)] pub struct Enum { - type_id: TypeId, + repr: EnumRepr, variants: Vec, doc: Option, } impl Enum { pub fn new( - type_id: TypeId, + repr: EnumRepr, variants: impl IntoIterator, doc: Option, ) -> Self { Self { - type_id: type_id, + repr, variants: variants.into_iter().collect(), doc, } } -} + pub fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } + + fn has_body(&self) -> bool { + !self.variants.is_empty() + } +} + +#[derive(Debug)] +pub enum EnumRepr { + /// Enums are tagged as the key in a collection. + External, + + /// Enums are tagged by an internal field. + /// + /// This contains the name of the field that is used as the tag. + /// + /// ```json + /// { + /// "tag": "some-enum-tag", + /// "value": 10, + /// "members": [ + /// 1, + /// "world", + /// ] + /// } + /// ``` + Internal { tag: &'static str }, + + /// Enums have a separate tag and content payload. + /// + /// ```json + /// { + /// "tag": "some-enum-tag", + /// "content": { + /// "value": 10, + /// "members": [ + /// 1, + /// "world", + /// ] + /// } + /// } + /// ``` + Adjacent { + tag: &'static str, + content: &'static str, + }, +} + +#[derive(Debug)] pub struct Variant { name: &'static str, fields: Fields, @@ -246,35 +372,65 @@ impl Variant { pub fn new(name: &'static str, fields: Fields, doc: Option) -> Self { Self { name, fields, doc } } + + fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } } //----------// // Sequence // //----------// +#[derive(Debug)] pub struct Sequence { - type_id: TypeId, element: Reflection, doc: Option, } impl Sequence { - pub fn new(type_id: TypeId, doc: Option) -> Self + pub fn new(doc: Option) -> Self where T: Reflect, { Self { - type_id, element: reflect::(), doc, } } + + pub fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } } ////////////// // Renderer // ////////////// +#[derive(Debug)] +struct Tagged { + ty: Type, + reflection: Reflection, +} + +impl Tagged { + fn new(reflection: Reflection) -> Self { + Self { + ty: reflection.reflect(), + reflection, + } + } + + fn ty(&self) -> &Type { + &self.ty + } + + fn reflection(&self) -> Reflection { + self.reflection + } +} + struct Renderer<'a> { output: &'a mut dyn Write, indent: usize, @@ -292,10 +448,6 @@ impl<'a> Renderer<'a> { } } - fn at_top(&self) -> bool { - self.indent == 0 - } - fn at_bottom(&self) -> bool { self.depth == self.max_depth } @@ -312,9 +464,9 @@ impl<'a> Renderer<'a> { self.output.write_char('\n') } - fn maybe_indent(&mut self, indent: bool, f: F) -> fmt::Result + fn maybe_indent(&mut self, indent: bool, f: F) -> Result where - F: FnOnce(&mut Self) -> fmt::Result, + F: FnOnce(&mut Self) -> Result, { if indent { self.indent += 1; @@ -326,74 +478,129 @@ impl<'a> Renderer<'a> { result } - fn indent(&mut self, f: F) -> fmt::Result + fn indent(&mut self, f: F) -> Result where - F: FnOnce(&mut Self) -> fmt::Result, + F: FnOnce(&mut Self) -> Result, { self.maybe_indent(true, f) } - fn next(&mut self, f: F) -> fmt::Result - where - F: FnOnce(&mut Self) -> fmt::Result, - { + fn next_with( + &mut self, + pre: impl FnOnce(&mut Self) -> fmt::Result, + body: impl FnOnce(&mut Self) -> fmt::Result, + post: impl FnOnce(&mut Self) -> fmt::Result, + ) -> fmt::Result { if self.at_bottom() { Ok(()) } else { + pre(self)?; self.depth += 1; - let result = self.indent(f); + let result = self.indent(body); self.depth -= 1; + post(self)?; result } } - fn render_doc(&mut self, s: Option<&str>) -> fmt::Result { + fn next(&mut self, f: F) -> fmt::Result + where + F: FnOnce(&mut Self) -> fmt::Result, + { + self.next_with(|_| Ok(()), f, |_| Ok(())) + } + + fn render_doc(&mut self, s: Option<&str>) -> Result { if let Some(s) = s { + let mut rendered = false; for ln in s.lines() { if ln.is_empty() { self.blank()?; } else { self.line(ln)?; } + + rendered = true; } + Ok(rendered) + } else { + Ok(false) } + } - Ok(()) + /// Return `true` if a nested type will be rendered. + fn will_render(&self, ty: &Type) -> bool { + !self.at_bottom() && ty.has_body() } //-------// // Types // //-------// - fn render_reflection(&mut self, reflection: Reflection) -> fmt::Result { - // Render the type name if this is the first item in the stack. - if self.at_top() { - self.line(reflection.type_name())?; - } + fn render_subject(&mut self, reflection: Reflection) -> fmt::Result { + self.line(reflection.type_name())?; + self.indent(|r| { + let tagged = Tagged::new(reflection); + let wrote_doc = r.render_doc(tagged.ty().doc())?; - match reflection.reflect() { - Type::Primitive(primitive) => self.render_primitive(&primitive), + if wrote_doc && r.will_render(tagged.ty()) { + r.blank()?; + } + + r.render_body(&tagged) + }) + } + + fn render_body(&mut self, tagged: &Tagged) -> fmt::Result { + match tagged.ty() { + Type::Primitive(_) => Ok(()), Type::Aggregate(aggregate) => self.render_aggregate(&aggregate), - Type::Enum(enum_) => todo!(), - Type::Sequence(sequence) => todo!(), + Type::Enum(enum_) => self.render_enum(enum_), + Type::Sequence(sequence) => self.render_sequence(&sequence), } } - fn render_primitive(&mut self, primitive: &Primitive) -> fmt::Result { - if self.at_top() { - self.indent(|r| r.render_doc(primitive.doc()))?; + fn render_aggregate(&mut self, aggregate: &Aggregate) -> fmt::Result { + self.render_fields(&aggregate.fields) + } + + fn render_enum(&mut self, enum_: &Enum) -> fmt::Result { + match enum_.repr { + EnumRepr::External => self.line("Representation: string")?, + EnumRepr::Internal { tag } => { + self.line(format_args!("Discriminant field: {}", Quote(tag)))? + } + EnumRepr::Adjacent { tag, content } => { + self.line(format_args!("Discriminant field: {}", Quote(tag)))?; + self.line(format_args!("Content field: {}", Quote(content)))?; + } } - Ok(()) - } + self.blank()?; + self.line("Options:")?; + self.indent(|r| { + let mut first = true; + for variant in enum_.variants.iter() { + if !first { + r.blank()?; + } - fn render_aggregate(&mut self, aggregate: &Aggregate) -> fmt::Result { - self.maybe_indent(self.at_top(), |r| { - r.render_doc(aggregate.doc())?; - r.render_fields(&aggregate.fields) + r.render_variant(variant)?; + first = false; + } + + Ok(()) }) } + fn render_sequence(&mut self, sequence: &Sequence) -> fmt::Result { + self.next_with( + |r| r.line(format_args!("Elements: {}", sequence.element.type_name())), + |r| r.render_body(&Tagged::new(sequence.element)), + |_| Ok(()), + ) + } + //--------// // Fields // //--------// @@ -401,71 +608,126 @@ impl<'a> Renderer<'a> { fn render_fields(&mut self, fields: &Fields) -> fmt::Result { match fields { Fields::Named(named) => { + let mut first = true; for field in named.iter() { + if !first { + self.blank()?; + } self.render_named_field(field)?; + first = false; + } + } + Fields::Unnamed(unnamed) => { + for (i, field) in unnamed.iter().enumerate() { + if i != 0 { + self.blank()?; + } + self.render_unnamed_field(i, field)?; } } - Fields::Unnamed(_) => todo!(), - Fields::Unit => todo!(), + Fields::Unit => {} } Ok(()) } fn render_named_field(&mut self, field: &NamedField) -> fmt::Result { - let f = field.field; + let tagged = Tagged::new(field.field); + let will_render_body = self.will_render(tagged.ty()); + + self.line(format_args!( + "{}: {}", + Quote(field.name), + tagged.reflection().type_name() + ))?; + + let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; + if rendered_doc && will_render_body { + self.blank()?; + } - self.blank()?; - self.line(format_args!("{}: {}", Quote(field.name), TypeName(f)))?; - self.indent(|r| r.render_doc(field.doc()))?; - self.next(|r| r.render_reflection(f)) - } -} - -// //////////////// -// // Algorithms // -// //////////////// -// -// pub fn walk<'a, I>(reflection: Reflection, paths: I) -> Result -// where -// I: IntoIterator, -// { -// let mut current = reflection; -// for p in paths { -// current = current.reflect().walk(p)?; -// } -// Ok(current) -// } -// -// #[derive(Debug, Clone, Copy)] -// pub struct WalkError; + if will_render_body { + self.next(|r| r.render_body(&tagged))?; + } + Ok(()) + } + + fn render_unnamed_field(&mut self, index: usize, field: &UnnamedField) -> fmt::Result { + let tagged = Tagged::new(field.field); + let will_render_body = self.will_render(tagged.ty()); + + self.line(format_args!( + "{}: {}", + index, + tagged.reflection().type_name() + ))?; + + let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; + if rendered_doc && will_render_body { + self.blank()?; + } + + if will_render_body { + self.next(|r| r.render_body(&tagged))?; + } + Ok(()) + } + + //---------// + // Variant // + //---------// + + fn render_variant(&mut self, variant: &Variant) -> fmt::Result { + self.line(Quote(variant.name))?; + + self.indent(|r| { + let rendered_doc = r.render_doc(variant.doc())?; + + if rendered_doc && variant.fields.has_body() { + r.blank()?; + } + + r.render_fields(&variant.fields) + }) + } +} /////////////// // Bootstrap // /////////////// -impl Reflect for usize { +macro_rules! primitive { + ($T:ty, $doc:literal, $type_name:literal) => { + impl Reflect for $T { + fn reflect() -> Type { + Type::primitive(Some($doc.into())) + } + + fn type_name(f: &mut dyn Write) -> fmt::Result { + f.write_str($type_name) + } + } + } +} + +primitive!(usize, "A system dependent unsigned integer", "usize"); +primitive!(u32, "A 32-bit unsigned integer", "u32"); + +primitive!(String, "A string", "string"); + +impl Reflect for Vec +where + T: Reflect, +{ fn reflect() -> Type { - Type::primitive( - TypeId::of::(), - Some("A system dependent unsigned integer".into()), - ) + Type::sequence::(Some("An ordered collection of elements".into())) } fn type_name(f: &mut dyn Write) -> fmt::Result { - f.write_str("usize") + write!(f, "Vec<{}>", reflect::().type_name()) } } -// impl Reflect for Vec -// where -// T: Reflect, -// { -// fn reflect() -> Type { -// -// } -// } - /// This is a test! /// /// Hello world! @@ -485,6 +747,12 @@ pub struct Test2 { a: usize, other: Test, + + /// How are we going to compute distances? + metric: AdjacentEnum, + + /// These control a bunch of parameters. + seq: Vec, } #[derive(Reflect)] @@ -493,12 +761,80 @@ struct Wrapper { a: T, } +/// An enum with no payloads. +#[derive(Debug, Clone, Copy)] +pub enum Metric { + SquaredL2, + InnerProduct, + Cosine, +} + +impl Reflect for Metric { + fn reflect() -> Type { + Type::enum_( + EnumRepr::External, + [ + Variant::new("squared-l2", Fields::Unit, Some("Squared Euclidean".into())), + Variant::new("inner-product", Fields::Unit, Some("Inner Product".into())), + Variant::new("cosine", Fields::Unit, Some("Cosine Similarity".into())), + ], + Some("The similarity measure to use".into()), + ) + } + + fn type_name(f: &mut dyn Write) -> fmt::Result { + f.write_str("Metric") + } +} + +/// An enum with no payloads. +#[derive(Debug)] +pub enum AdjacentEnum { + SquaredL2, + InnerProduct(u32), + Cosine { test: String }, +} + +impl Reflect for AdjacentEnum { + fn reflect() -> Type { + Type::enum_( + EnumRepr::Adjacent { + tag: "enum-type", + content: "content", + }, + [ + Variant::new("squared-l2", Fields::Unit, None), + Variant::new( + "inner-product", + Fields::unnamed([UnnamedField::new::(Some("testing".into()))]), + Some("Inner Product with some payload".into()), + ), + Variant::new( + "cosine", + Fields::named([NamedField::new::("test", None)]), + Some("Cosine Similarity".into()), + ), + ], + Some("The similarity measure to use".into()), + ) + } + + fn type_name(f: &mut dyn Write) -> fmt::Result { + f.write_str("Metric") + } +} + +////////////// +// Internal // +////////////// + pub(crate) mod internal { use std::marker::PhantomData; pub(crate) trait Reflect { fn reflect(&self) -> super::Type; fn type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result; + fn type_id(&self) -> std::any::TypeId; } pub(crate) struct Wrapper(PhantomData); @@ -522,5 +858,9 @@ pub(crate) mod internal { fn type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result { ::type_name(f) } + + fn type_id(&self) -> std::any::TypeId { + std::any::TypeId::of::() + } } } diff --git a/diskann-benchmark-runner/temp/main.rs b/diskann-benchmark-runner/temp/main.rs index 79186060f8..8a66f3ddb2 100644 --- a/diskann-benchmark-runner/temp/main.rs +++ b/diskann-benchmark-runner/temp/main.rs @@ -12,5 +12,7 @@ fn main() -> anyhow::Result<()> { println!("{}", reflect::reflect::().render()); + println!("{}", reflect::reflect::().render()); + Ok(()) } From 57c5c33e6c81209b1814dac1bca1db9208a8a86c Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Fri, 25 Sep 2026 16:19:15 -0700 Subject: [PATCH 05/21] More derives! --- diskann-benchmark-runner-derive/src/lib.rs | 323 +++++++++++-- diskann-benchmark-runner/src/lib.rs | 2 +- diskann-benchmark-runner/src/reflect.rs | 510 +++++++++++++++++---- diskann-benchmark-runner/temp/main.rs | 8 +- 4 files changed, 715 insertions(+), 128 deletions(-) diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index b2cb2eb32b..7c2d789817 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -3,7 +3,7 @@ * Licensed under the MIT license. */ -use proc_macro::TokenStream; +use proc_macro2::TokenStream; use quote::{quote, quote_spanned}; use syn::{Data, DeriveInput, Fields, parse_macro_input, parse_quote, spanned::Spanned}; @@ -29,66 +29,311 @@ fn crate_name() -> syn::Path { /// } /// ``` #[proc_macro_derive(Reflect, attributes(reflect))] -pub fn derive_reflect(input: TokenStream) -> TokenStream { +pub fn derive_reflect(input: proc_macro::TokenStream) -> proc_macro::TokenStream { let input = parse_macro_input!(input as DeriveInput); + + if matches!(input.data, Data::Union(_)) { + todo!("return a better error message"); + } + + let doc = format_docstrings(&input.attrs); + let mut generics = input.generics.clone(); + add_generic_bounds(&mut generics); + + let format_type_name = generate_type_name_body(&input); + + let common = DeriveCommon { + doc, + generics, + format_type_name, + }; + let output = match &input.data { - Data::Struct(s) => process_struct(&input, s), + Data::Struct(s) => process_struct(&input, s, common), + Data::Enum(e) => process_enum(&input, e, common), _ => todo!("need to figure this out"), }; output.into() } -fn process_struct(input: &DeriveInput, s: &syn::DataStruct) -> proc_macro2::TokenStream { +/// Common pre-processed items. +struct DeriveCommon { + /// Documentation for the top-level derive input. + doc: TokenStream, + /// The generics with a `where T: Reflect` added for all generic parameters. + generics: syn::Generics, + /// The implementation of `format_type_name`. + format_type_name: TokenStream, +} + +/// Add a bound `T: Reflect` for each type parameter in the generic list. +/// +/// This is needed to correctly render type-names. +/// +/// For example +/// ``` +/// struct Foo { +/// bar: Vec, +/// } +/// ``` +/// should have a type name like `Foo`. We get the value in the brackets from the +/// `Reflect` bound as well, so we need to add the following bounds: +/// ```text +/// impl Reflect for Foo +/// where +/// T: Reflect, +/// { +/// ... +/// } +/// ``` +fn add_generic_bounds(generics: &mut syn::Generics) { + // Get the raw type parameters. + let type_params: Vec<_> = generics.type_params().map(|p| p.ident.clone()).collect(); + + let predicates = &mut generics.make_where_clause().predicates; + let path = crate_name(); + + for p in type_params { + predicates.push(parse_quote!(#p: #path::Reflect)); + } +} + +/// Add a bound `T: Reflect` for each type in the field. +/// +/// For example, if a struct definition looks like this: +/// ``` +/// struct Foo { +/// bar: usize, +/// } +/// ``` +/// This will add bounds like this +/// ```text +/// impl Reflect for Foo +/// where +/// usize: Reflect +/// { +/// ... +/// } +/// ``` +fn add_field_bounds<'a, I>(generics: &mut syn::Generics, fields: I) +where + I: IntoIterator, +{ + let path = crate_name(); + for field in fields { + let ty = &field.ty; + generics + .make_where_clause() + .predicates + .push(parse_quote!(#ty: #path::Reflect)); + } +} + +/// Generate the expression for a type name. +/// +/// The main complexity comes from dealing with generics. +/// +/// When there are generics (say `Foo`), we want the following implement +/// to look something like this when `T == u32` and `N == 10`: +/// ``` +/// fn type_name(f: &mut dyn std::fmt::Write) -> std::fmt::Result { +/// f.write_str("Foo"); +/// f.write_str("<"); +/// // This would come from the `Reflection::type_name` instead. +/// type_name_u32(f); +/// f.write_str(">"); +/// } +/// +/// fn type_name_u32(f: &mut dyn std::fmt::Write) -> std::fmt::Result { +/// f.write_str("u32") +/// } +/// ``` +/// When there are no generics - we can print the type name directly. +fn generate_type_name_body(input: &DeriveInput) -> TokenStream { + let path = crate_name(); + let name = input.ident.to_string(); + let arguments: Vec<_> = input + .generics + .params + .iter() + .filter_map(|p| match p { + syn::GenericParam::Type(p) => { + let ident = &p.ident; + Some(quote! { + ::std::write!( + f, + "{}", + #path::Reflection::new::<#ident>().type_name(), + )?; + }) + } + syn::GenericParam::Const(p) => { + let ident = &p.ident; + Some(quote! { + ::std::write!(f, "{}", #ident)?; + }) + } + syn::GenericParam::Lifetime(_) => None, + }) + .collect(); + + // If there are no generics, we can dump the typename directly. + if arguments.is_empty() { + return quote! { + f.write_str(#name) + }; + } else { + let writes = arguments.iter().enumerate().map(|(index, argument)| { + if index == 0 { + quote! { + #argument + } + } else { + quote! { + f.write_str(", ")?; + #argument + } + } + }); + + quote! { + f.write_str(#name)?; + f.write_str("<")?; + #(#writes)* + f.write_str(">") + } + } +} + +fn build_fields(fields: &syn::Fields, generics: &mut syn::Generics) -> TokenStream { + let path = crate_name(); + + match fields { + Fields::Named(fields) => { + add_field_bounds(generics, &fields.named); + let list = named_fields(&fields.named); + quote!(#path::Fields::Named(vec![#(#list),*])) + } + Fields::Unnamed(fields) => { + add_field_bounds(generics, &fields.unnamed); + let list = unnamed_fields(&fields.unnamed); + quote!(#path::Fields::Unnamed(vec![#(#list),*])) + } + Fields::Unit => quote!(#path::Fields::Unit), + } +} + +fn named_fields<'a, I>(fields: I) -> impl Iterator +where + I: IntoIterator, +{ + let path = crate_name(); + fields.into_iter().map(move |f| { + let ty = &f.ty; + let ident = f + .ident + .as_ref() + .expect("named fields should have identifiers") + .to_string(); + let doc = format_docstrings(&f.attrs); + quote_spanned! { ty.span()=> #path::NamedField::new::<#ty>(#ident, #doc) } + }) +} + +fn unnamed_fields<'a, I>(fields: I) -> impl Iterator +where + I: IntoIterator, +{ + let path = crate_name(); + fields.into_iter().map(move |f| { + let ty = &f.ty; + let doc = format_docstrings(&f.attrs); + quote_spanned! { ty.span()=> #path::UnnamedField::new::<#ty>(#doc) } + }) +} + +/// Generate the `Reflect` implementation. +fn process_struct(input: &DeriveInput, s: &syn::DataStruct, common: DeriveCommon) -> TokenStream { + let DeriveCommon { + doc, + mut generics, + format_type_name, + } = common; + let type_name = &input.ident; - let type_name_str = type_name.to_string(); let path = crate_name(); - let doc = format_docstrings(&input.attrs); - match &s.fields { - Fields::Named(named) => { - // To handle generics and automatically apply the `Reflect` bound to inner types, - // we extract each field type `T` and add the bound `T: Reflect` to the type's - // where clause. - let mut generics = input.generics.clone(); - for field in &named.named { - let ty = &field.ty; - generics - .make_where_clause() - .predicates - .push(parse_quote!(#ty: #path::Reflect)); + let fields = build_fields(&s.fields, &mut generics); + let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); + + quote! { + impl #impl_generics #path::Reflect for #type_name #ty_generics #where_clause { + fn ty() -> #path::Type { + #path::Type::aggregate( + #fields, + #doc, + ) } - let fields = named.named.iter().map(|f| { - let ty = &f.ty; + fn format_type_name(f: &mut dyn ::std::fmt::Write) -> ::std::fmt::Result { + #format_type_name + } + } + } +} - let ident = f.ident.as_ref().unwrap().to_string(); - let doc = format_docstrings(&f.attrs); +//-------// +// Enums // +//-------// - quote_spanned! { ty.span()=> #path::NamedField::new::<#ty>(#ident, #doc) } - }); +fn process_enum(input: &DeriveInput, e: &syn::DataEnum, common: DeriveCommon) -> TokenStream { + let DeriveCommon { + doc, + mut generics, + format_type_name, + } = common; - let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); + // TODO: For now, we just assume that identifiers are taken as-is. + let type_name = &input.ident; + let path = crate_name(); - quote! { - impl #impl_generics #path::Reflect for #type_name #ty_generics #where_clause { - fn reflect() -> #path::Type { - #path::Type::aggregate( - #path::Fields::Named(vec![#(#fields),*]), - #doc, - ) - } + let variants: Vec<_> = e + .variants + .iter() + .map(|v| { + let doc = format_docstrings(&v.attrs); + let ident = v.ident.to_string(); - fn type_name(f: &mut dyn ::std::fmt::Write) -> ::std::fmt::Result { - f.write_str(#type_name_str) - } - } + let fields = build_fields(&v.fields, &mut generics); + + quote!(#path::Variant::new(#ident, #fields, #doc)) + }) + .collect(); + + let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); + + quote! { + impl #impl_generics #path::Reflect for #type_name #ty_generics #where_clause { + fn ty() -> #path::Type { + #path::Type::enum_( + #path::EnumRepr::External, + [#(#variants),*], + #doc, + ) + } + + fn format_type_name(f: &mut dyn ::std::fmt::Write) -> ::std::fmt::Result { + #format_type_name } } - _ => todo!("more todos"), } } -fn format_docstrings(attributes: &[syn::Attribute]) -> proc_macro2::TokenStream { +//-------------// +// Doc Strings // +//-------------// + +fn format_docstrings(attributes: &[syn::Attribute]) -> TokenStream { match extract_docs(attributes) { None => quote! { ::std::option::Option::None }, Some(docs) => quote! { ::std::option::Option::Some(#docs.into()) }, diff --git a/diskann-benchmark-runner/src/lib.rs b/diskann-benchmark-runner/src/lib.rs index 40220674d7..c6eb596aa8 100644 --- a/diskann-benchmark-runner/src/lib.rs +++ b/diskann-benchmark-runner/src/lib.rs @@ -29,7 +29,7 @@ pub use checker::Checker; pub use features::Features; pub use input::Input; pub use output::Output; -pub use reflect::Reflect; +pub use reflect::{Reflect, Reflection}; pub use registry::{Registry, RegistryError}; pub use result::Checkpoint; diff --git a/diskann-benchmark-runner/src/reflect.rs b/diskann-benchmark-runner/src/reflect.rs index ec2cc85c55..fa41a2bb08 100644 --- a/diskann-benchmark-runner/src/reflect.rs +++ b/diskann-benchmark-runner/src/reflect.rs @@ -16,15 +16,8 @@ use crate::utils::fmt::Quote; const INDENT: usize = 2; pub trait Reflect: 'static { - fn reflect() -> Type; - fn type_name(f: &mut dyn Write) -> fmt::Result; -} - -pub fn reflect() -> Reflection -where - T: Reflect, -{ - Reflection::new::() + fn ty() -> Type; + fn format_type_name(f: &mut dyn Write) -> fmt::Result; } #[derive(Clone, Copy)] @@ -42,8 +35,8 @@ impl Reflection { } } - pub fn reflect(&self) -> Type { - self.reflection.reflect() + pub fn ty(&self) -> Type { + self.reflection.ty() } pub fn type_name(&self) -> TypeName { @@ -72,20 +65,20 @@ pub type Doc = Cow<'static, str>; pub struct TypeName(Reflection); impl TypeName { - fn format_into(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.reflection.type_name(f) + fn format_type_name(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.reflection.format_type_name(f) } } impl std::fmt::Debug for TypeName { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.format_into(f) + self.format_type_name(f) } } impl std::fmt::Display for TypeName { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.format_into(f) + self.format_type_name(f) } } @@ -147,6 +140,22 @@ impl Type { Self::Sequence(_) => true, } } + + fn as_aggregate(&self) -> Option<&Aggregate> { + if let Self::Aggregate(aggregate) = self { + Some(aggregate) + } else { + None + } + } + + fn as_enum(&self) -> Option<&Enum> { + if let Self::Enum(enum_) = self { + Some(enum_) + } else { + None + } + } } impl From for Type { @@ -211,6 +220,10 @@ impl Aggregate { self.doc.as_deref() } + fn fields(&self) -> &Fields { + &self.fields + } + fn has_body(&self) -> bool { self.fields.has_body() } @@ -239,6 +252,22 @@ impl Fields { pub fn unnamed(itr: impl IntoIterator) -> Self { Self::Unnamed(itr.into_iter().collect()) } + + fn as_named(&self) -> Option<&[NamedField]> { + if let Self::Named(fields) = self { + Some(fields) + } else { + None + } + } + + fn as_unnamed(&self) -> Option<&[UnnamedField]> { + if let Self::Unnamed(fields) = self { + Some(fields) + } else { + None + } + } } #[derive(Debug)] @@ -255,7 +284,7 @@ impl NamedField { { Self { name, - field: reflect::(), + field: Reflection::new::(), doc, } } @@ -277,7 +306,7 @@ impl UnnamedField { T: Reflect, { Self { - field: reflect::(), + field: Reflection::new::(), doc, } } @@ -285,6 +314,10 @@ impl UnnamedField { fn doc(&self) -> Option<&str> { self.doc.as_deref() } + + fn field(&self) -> Reflection { + self.field + } } //------// @@ -315,6 +348,10 @@ impl Enum { self.doc.as_deref() } + fn variants(&self) -> &[Variant] { + &self.variants + } + fn has_body(&self) -> bool { !self.variants.is_empty() } @@ -335,7 +372,7 @@ pub enum EnumRepr { /// "value": 10, /// "members": [ /// 1, - /// "world", + /// "world" /// ] /// } /// ``` @@ -350,7 +387,7 @@ pub enum EnumRepr { /// "value": 10, /// "members": [ /// 1, - /// "world", + /// "world" /// ] /// } /// } @@ -373,6 +410,10 @@ impl Variant { Self { name, fields, doc } } + fn name(&self) -> &'static str { + self.name + } + fn doc(&self) -> Option<&str> { self.doc.as_deref() } @@ -394,7 +435,7 @@ impl Sequence { T: Reflect, { Self { - element: reflect::(), + element: Reflection::new::(), doc, } } @@ -417,7 +458,7 @@ struct Tagged { impl Tagged { fn new(reflection: Reflection) -> Self { Self { - ty: reflection.reflect(), + ty: reflection.ty(), reflection, } } @@ -449,7 +490,7 @@ impl<'a> Renderer<'a> { } fn at_bottom(&self) -> bool { - self.depth == self.max_depth + self.depth >= self.max_depth } fn line(&mut self, display: D) -> fmt::Result @@ -566,7 +607,7 @@ impl<'a> Renderer<'a> { fn render_enum(&mut self, enum_: &Enum) -> fmt::Result { match enum_.repr { - EnumRepr::External => self.line("Representation: string")?, + EnumRepr::External => self.line("Representation: externally tagged")?, EnumRepr::Internal { tag } => { self.line(format_args!("Discriminant field: {}", Quote(tag)))? } @@ -696,18 +737,31 @@ impl<'a> Renderer<'a> { // Bootstrap // /////////////// +impl Reflect for std::marker::PhantomData +where + T: Reflect, +{ + fn ty() -> Type { + Type::primitive(None) + } + + fn format_type_name(f: &mut dyn Write) -> fmt::Result { + write!(f, "PhantomData<{}>", Reflection::new::().type_name()) + } +} + macro_rules! primitive { ($T:ty, $doc:literal, $type_name:literal) => { impl Reflect for $T { - fn reflect() -> Type { + fn ty() -> Type { Type::primitive(Some($doc.into())) } - fn type_name(f: &mut dyn Write) -> fmt::Result { + fn format_type_name(f: &mut dyn Write) -> fmt::Result { f.write_str($type_name) } } - } + }; } primitive!(usize, "A system dependent unsigned integer", "usize"); @@ -719,15 +773,23 @@ impl Reflect for Vec where T: Reflect, { - fn reflect() -> Type { + fn ty() -> Type { Type::sequence::(Some("An ordered collection of elements".into())) } - fn type_name(f: &mut dyn Write) -> fmt::Result { - write!(f, "Vec<{}>", reflect::().type_name()) + fn format_type_name(f: &mut dyn Write) -> fmt::Result { + write!(f, "Vec<{}>", Reflection::new::().type_name()) } } +#[derive(Reflect)] +pub struct UnitWithConst {} + +#[derive(Reflect)] +pub struct GenericBoundAdded { + uses_t: Vec, +} + /// This is a test! /// /// Hello world! @@ -740,12 +802,22 @@ pub struct Test { b: usize, } +#[derive(Reflect)] +pub struct TestUnnamed( + /// Can I document this? + usize, + T, +); + /// This is a nother test! #[derive(Reflect)] pub struct Test2 { /// This field affects this value. a: usize, + /// This field doesn't have any names. + unnamed: TestUnnamed, + other: Test, /// How are we going to compute distances? @@ -762,68 +834,56 @@ struct Wrapper { } /// An enum with no payloads. -#[derive(Debug, Clone, Copy)] +#[derive(Debug, Clone, Copy, Reflect)] pub enum Metric { SquaredL2, InnerProduct, Cosine, } -impl Reflect for Metric { - fn reflect() -> Type { - Type::enum_( - EnumRepr::External, - [ - Variant::new("squared-l2", Fields::Unit, Some("Squared Euclidean".into())), - Variant::new("inner-product", Fields::Unit, Some("Inner Product".into())), - Variant::new("cosine", Fields::Unit, Some("Cosine Similarity".into())), - ], - Some("The similarity measure to use".into()), - ) - } - - fn type_name(f: &mut dyn Write) -> fmt::Result { - f.write_str("Metric") - } -} - /// An enum with no payloads. -#[derive(Debug)] +#[derive(Debug, Reflect)] pub enum AdjacentEnum { SquaredL2, + /// Let me see if this works InnerProduct(u32), - Cosine { test: String }, -} -impl Reflect for AdjacentEnum { - fn reflect() -> Type { - Type::enum_( - EnumRepr::Adjacent { - tag: "enum-type", - content: "content", - }, - [ - Variant::new("squared-l2", Fields::Unit, None), - Variant::new( - "inner-product", - Fields::unnamed([UnnamedField::new::(Some("testing".into()))]), - Some("Inner Product with some payload".into()), - ), - Variant::new( - "cosine", - Fields::named([NamedField::new::("test", None)]), - Some("Cosine Similarity".into()), - ), - ], - Some("The similarity measure to use".into()), - ) - } - - fn type_name(f: &mut dyn Write) -> fmt::Result { - f.write_str("Metric") - } + /// Compute the cosine similarity + Cosine { + /// Thos actually doesn't do anything. + test: String, + }, } +// impl Reflect for AdjacentEnum { +// fn ty() -> Type { +// Type::enum_( +// EnumRepr::Adjacent { +// tag: "enum-type", +// content: "content", +// }, +// [ +// Variant::new("squared-l2", Fields::Unit, None), +// Variant::new( +// "inner-product", +// Fields::unnamed([UnnamedField::new::(Some("testing".into()))]), +// Some("Inner Product with some payload".into()), +// ), +// Variant::new( +// "cosine", +// Fields::named([NamedField::new::("test", None)]), +// Some("Cosine Similarity".into()), +// ), +// ], +// Some("The similarity measure to use".into()), +// ) +// } +// +// fn format_type_name(f: &mut dyn Write) -> fmt::Result { +// f.write_str("Metric") +// } +// } + ////////////// // Internal // ////////////// @@ -832,8 +892,8 @@ pub(crate) mod internal { use std::marker::PhantomData; pub(crate) trait Reflect { - fn reflect(&self) -> super::Type; - fn type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result; + fn ty(&self) -> super::Type; + fn format_type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result; fn type_id(&self) -> std::any::TypeId; } @@ -851,12 +911,12 @@ pub(crate) mod internal { where T: super::Reflect, { - fn reflect(&self) -> super::Type { - ::reflect() + fn ty(&self) -> super::Type { + ::ty() } - fn type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result { - ::type_name(f) + fn format_type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result { + ::format_type_name(f) } fn type_id(&self) -> std::any::TypeId { @@ -864,3 +924,285 @@ pub(crate) mod internal { } } } + +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + use std::assert_matches; + + #[test] + fn test_unit() { + /// A unit struct. + #[derive(Reflect)] + struct Unit; + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(r.ty().doc().unwrap(), "A unit struct."); + assert_eq!(r.type_name().to_string(), "Unit"); + assert!(!ty.has_body()); + } + + #[test] + fn test_unit_const_generic() { + /// A unit struct. + #[derive(Reflect)] + struct Unit; + + let r = Reflection::new::>(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "A unit struct."); + assert_eq!(r.type_name().to_string(), "Unit<10>"); + assert!(!ty.has_body()); + } + + #[test] + fn test_unit_const_generic_2() { + #[derive(Reflect)] + struct Unit; + + let r = Reflection::new::>(); + let ty = r.ty(); + + assert!(ty.doc().is_none()); + assert_eq!(r.type_name().to_string(), "Unit<10, 20>"); + assert!(!ty.has_body()); + } + + #[test] + fn test_empty_tuple_like() { + /// An empty tuple-like struct. + #[derive(Reflect)] + struct Empty(); + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "An empty tuple-like struct."); + assert_eq!(r.type_name().to_string(), "Empty"); + assert!(!ty.has_body()); + } + + #[test] + fn test_empty_struct_like() { + /// An empty struct. + #[derive(Reflect)] + struct Empty {} + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "An empty struct."); + assert_eq!(r.type_name().to_string(), "Empty"); + assert!(!ty.has_body()); + } + + #[test] + fn test_struct() { + /// A struct with two fields. + #[expect(unused)] + #[derive(Reflect)] + struct Woo { + /// Foo + foo: usize, + /// Bar + bar: usize, + }; + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "A struct with two fields."); + assert_eq!(r.type_name().to_string(), "Woo"); + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_named().unwrap(); + assert_eq!(f.len(), 2); + + assert_eq!(f[0].name, "foo"); + assert_eq!(f[0].doc().unwrap(), "Foo"); + assert_eq!(f[0].field.type_name().to_string(), "usize"); + + assert_eq!(f[1].name, "bar"); + assert_eq!(f[1].doc().unwrap(), "Bar"); + assert_eq!(f[1].field.type_name().to_string(), "usize"); + } + + #[test] + fn test_tuple1() { + /// A tuple with one field. + #[expect(unused)] + #[derive(Reflect)] + struct Tuple1( + /// Field 0. + usize, + ); + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "A tuple with one field."); + assert_eq!(r.type_name().to_string(), "Tuple1"); + + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); + assert_eq!(f.len(), 1); + assert_eq!(f[0].doc().unwrap(), "Field 0."); + assert_eq!(f[0].field().type_name().to_string(), "usize"); + } + + #[test] + fn test_tuple2() { + #[expect(unused)] + #[derive(Reflect)] + struct Tuple2( + T, + /// Field 1. + Vec, + ); + + let r = Reflection::new::>(); + let ty = r.ty(); + assert!(ty.doc().is_none()); + assert_eq!(r.type_name().to_string(), "Tuple2"); + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); + assert_eq!(f.len(), 2); + + assert!(f[0].doc().is_none()); + assert_eq!(f[0].field().type_name().to_string(), "usize"); + + assert_eq!(f[1].doc().unwrap(), "Field 1."); + assert_eq!(f[1].field().type_name().to_string(), "Vec"); + } + + #[test] + fn test_empty_enum() { + /// An empty enum. + #[derive(Reflect)] + enum Empty {} + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "An empty enum."); + assert_eq!(r.type_name().to_string(), "Empty"); + assert!(!ty.has_body()); + } + + #[test] + fn test_enum_with_generics() { + /// An enum with generics. + #[expect(unused)] + #[derive(Reflect)] + enum Either { + A(Vec), + /// It's a bee! + B( + /// Buzz buzz + B + ), + } + + let r = Reflection::new::>(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "An enum with generics."); + assert_eq!(r.type_name().to_string(), "Either"); + assert!(ty.has_body()); + + let variants = ty.as_enum().unwrap().variants(); + assert_eq!(variants.len(), 2); + + // Variant 0 + assert!(variants[0].doc().is_none()); + assert_eq!(variants[0].name(), "A"); + assert!(variants[0].fields.has_body()); + + let f = variants[0].fields.as_unnamed().unwrap(); + assert_eq!(f.len(), 1); + assert!(f[0].doc().is_none()); + assert_eq!(f[0].field().type_name().to_string(), "Vec"); + + // Variant 1 + assert_eq!(variants[1].doc().unwrap(), "It's a bee!"); + assert_eq!(variants[1].name(), "B"); + assert!(variants[1].fields.has_body()); + + let f = variants[1].fields.as_unnamed().unwrap(); + assert_eq!(f.len(), 1); + assert_eq!(f[0].doc().unwrap(), "Buzz buzz"); + assert_eq!(f[0].field().type_name().to_string(), "string"); + } + + #[test] + fn test_enum_variants() { + /// All the enums. + #[expect(unused)] + #[derive(Reflect)] + enum All { + /// A unit variant. + Unit, + /// A tuple-like variant. + Tuple( + /// Field 0. + usize, + String, + ), + /// Struct-like. + Struct { + foo: usize, + /// All the bars! + bar: u32, + }, + } + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "All the enums."); + assert_eq!(r.type_name().to_string(), "All"); + assert!(ty.has_body()); + + let variants = ty.as_enum().unwrap().variants(); + assert_eq!(variants.len(), 3); + + // Variant 0 + assert_eq!(variants[0].doc().unwrap(), "A unit variant."); + assert_eq!(variants[0].name(), "Unit"); + assert!(!variants[0].fields.has_body()); + assert_matches!(variants[0].fields, Fields::Unit); + + // Variant 1 + assert_eq!(variants[1].doc().unwrap(), "A tuple-like variant."); + assert_eq!(variants[1].name(), "Tuple"); + assert!(variants[1].fields.has_body()); + + let f = variants[1].fields.as_unnamed().unwrap(); + assert_eq!(f.len(), 2); + + assert_eq!(f[0].doc().unwrap(), "Field 0."); + assert_eq!(f[0].field.type_name().to_string(), "usize"); + + assert!(f[1].doc().is_none()); + assert_eq!(f[1].field.type_name().to_string(), "string"); + + // Variant 2 + assert_eq!(variants[2].name(), "Struct"); + assert_eq!(variants[2].doc().unwrap(), "Struct-like."); + let f = variants[2].fields.as_named().unwrap(); + assert_eq!(f.len(), 2); + + assert_eq!(f[0].name, "foo"); + assert!(f[0].doc().is_none()); + assert_eq!(f[0].field.type_name().to_string(), "usize"); + + assert_eq!(f[1].name, "bar"); + assert_eq!(f[1].doc().unwrap(), "All the bars!"); + assert_eq!(f[1].field.type_name().to_string(), "u32"); + } +} diff --git a/diskann-benchmark-runner/temp/main.rs b/diskann-benchmark-runner/temp/main.rs index 8a66f3ddb2..90d24fa9c6 100644 --- a/diskann-benchmark-runner/temp/main.rs +++ b/diskann-benchmark-runner/temp/main.rs @@ -5,14 +5,14 @@ //! Development CLI for exercising the benchmark runner with its test registry. -use diskann_benchmark_runner::reflect; +use diskann_benchmark_runner::{reflect, Reflection}; fn main() -> anyhow::Result<()> { - println!("{}", reflect::reflect::().render()); + println!("{}", Reflection::new::().render()); - println!("{}", reflect::reflect::().render()); + println!("{}", Reflection::new::().render()); - println!("{}", reflect::reflect::().render()); + println!("{}", Reflection::new::().render()); Ok(()) } From a91e2c9bf4b6951512ae9c99e276f0bb93bb7928 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Fri, 25 Sep 2026 18:33:12 -0700 Subject: [PATCH 06/21] Progress! --- Cargo.lock | 1 + diskann-benchmark-runner-derive/Cargo.toml | 1 + diskann-benchmark-runner-derive/SERDE_TODO.md | 57 ++++ diskann-benchmark-runner-derive/src/lib.rs | 149 ++++++--- diskann-benchmark-runner-derive/src/serde.rs | 294 ++++++++++++++++++ diskann-benchmark-runner/src/reflect.rs | 6 +- 6 files changed, 470 insertions(+), 38 deletions(-) create mode 100644 diskann-benchmark-runner-derive/SERDE_TODO.md create mode 100644 diskann-benchmark-runner-derive/src/serde.rs diff --git a/Cargo.lock b/Cargo.lock index 229e36d466..0dd8a1e6fc 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -530,6 +530,7 @@ dependencies = [ name = "diskann-benchmark-runner-derive" version = "0.59.0" dependencies = [ + "heck", "proc-macro2", "quote", "syn 2.0.117", diff --git a/diskann-benchmark-runner-derive/Cargo.toml b/diskann-benchmark-runner-derive/Cargo.toml index 9672dbebcb..4bedaf395c 100644 --- a/diskann-benchmark-runner-derive/Cargo.toml +++ b/diskann-benchmark-runner-derive/Cargo.toml @@ -13,6 +13,7 @@ proc-macro = true syn = { version = "2", features = ["full"] } quote = "1" proc-macro2 = "1" +heck = "0.5.0" [lints] workspace = true diff --git a/diskann-benchmark-runner-derive/SERDE_TODO.md b/diskann-benchmark-runner-derive/SERDE_TODO.md new file mode 100644 index 0000000000..58ad1db011 --- /dev/null +++ b/diskann-benchmark-runner-derive/SERDE_TODO.md @@ -0,0 +1,57 @@ +# Serde Support: Resume Notes + +The initial Serde attribute parser and code generation are in place. It currently handles +field and variant names, container-level `rename_all`, and external, internal, and adjacent +enum representations. + +## Correctness + +- [ ] Use separate Serde-compatible case conversion for fields and variants. + - Serde applies different rules in each context. + - In particular, Serde converts an enum variant such as `XMLHttpRequest` to + `x_m_l_http_request` under `snake_case`, while `heck` produces + `xml_http_request`. + - For fields, `snake_case` is identity and `kebab-case` replaces underscores. + - Implementing the supported transformations directly should allow removal of the `heck` + dependency. +- [ ] Reject tuple variants on internally tagged enums with an error at the variant. +- [ ] Confirm the intended treatment of internally tagged newtype variants. Serde accepts + some forms syntactically, but compatibility depends on the wrapped value's serialized + shape. +- [ ] Check raw identifiers such as `r#type`; reflected names must match Serde's wire names + rather than include the raw-identifier prefix. + +## Tests + +- [ ] Update the enum smoke test to expect renamed variants such as `"unit"` rather than + `"Unit"`. +- [ ] Assert the generated enum representation, including both `tag` and `content` for an + adjacently tagged enum. +- [ ] Add naming tests that compare reflection metadata with `serde_json`, covering: + - explicit field and variant `rename`; + - struct field `rename_all`; + - enum variant `rename_all`; + - variant-level `rename_all` for struct-variant fields; + - acronym-heavy variants such as `XMLHttpRequest`; + - explicit `rename` taking precedence over `rename_all`. +- [ ] Add representation tests for external, internal, and adjacent tagging. +- [ ] Add compile-fail tests for duplicate attributes, `content` without `tag`, enum-only + attributes on structs, internally tagged tuple variants, unsupported rename rules, and + unsupported Serde attributes. + +## Cleanup and validation + +- [ ] Fix the `generate_type_name_body` doctest by returning the final + `f.write_str(">")` result instead of discarding it with a semicolon. +- [ ] Run `cargo fmt --all`. +- [ ] Run the targeted derive and runner tests. +- [ ] Run Clippy with warnings denied once the implementation and tests settle. + +## Deferred Serde features + +- [ ] Decide how `default` and `alias` should appear in reflection metadata before accepting + them. +- [ ] Continue rejecting asymmetric serialization/deserialization names unless the metadata + model represents both. +- [ ] Continue rejecting `untagged`, `flatten`, `remote`, `with`, `deserialize_with`, + `try_from`, and skipped input fields until each has an explicit metadata design. diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index 7c2d789817..21bb7edf91 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -7,6 +7,8 @@ use proc_macro2::TokenStream; use quote::{quote, quote_spanned}; use syn::{Data, DeriveInput, Fields, parse_macro_input, parse_quote, spanned::Spanned}; +mod serde; + fn crate_name() -> syn::Path { syn::parse_quote!(::diskann_benchmark_runner::reflect) } @@ -28,12 +30,32 @@ fn crate_name() -> syn::Path { /// threads: usize, /// } /// ``` -#[proc_macro_derive(Reflect, attributes(reflect))] +/// +/// # Serde Compatibility +/// +/// This is meant to be used in conjunction with `serde` for documenting benchmark inputs +/// and as such, it respects several common `serde` attributes: +/// +/// * rename_all = "lowercase" | "snake_case" | "kebab-case" +/// * rename = "..." +/// * tag = "..." +/// * tag = "...", content = "..." +#[proc_macro_derive(Reflect, attributes(reflect, serde))] pub fn derive_reflect(input: proc_macro::TokenStream) -> proc_macro::TokenStream { let input = parse_macro_input!(input as DeriveInput); + expand(&input) + .unwrap_or_else(syn::Error::into_compile_error) + .into() +} + +fn expand(input: &DeriveInput) -> syn::Result { + // Error as early as possible. if matches!(input.data, Data::Union(_)) { - todo!("return a better error message"); + return Err(syn::Error::new_spanned( + input, + "Reflect cannot be derived for unions", + )); } let doc = format_docstrings(&input.attrs); @@ -41,19 +63,20 @@ pub fn derive_reflect(input: proc_macro::TokenStream) -> proc_macro::TokenStream add_generic_bounds(&mut generics); let format_type_name = generate_type_name_body(&input); + let container = serde::Container::parse(&input.attrs)?; let common = DeriveCommon { doc, generics, format_type_name, + container, }; - let output = match &input.data { + match &input.data { Data::Struct(s) => process_struct(&input, s, common), Data::Enum(e) => process_enum(&input, e, common), - _ => todo!("need to figure this out"), - }; - output.into() + Data::Union(_) => unreachable!("this has already been checked"), + } } /// Common pre-processed items. @@ -64,6 +87,8 @@ struct DeriveCommon { generics: syn::Generics, /// The implementation of `format_type_name`. format_type_name: TokenStream, + /// Serde container-level attributes. + container: serde::Container, } /// Add a bound `T: Reflect` for each type parameter in the generic list. @@ -205,39 +230,53 @@ fn generate_type_name_body(input: &DeriveInput) -> TokenStream { } } -fn build_fields(fields: &syn::Fields, generics: &mut syn::Generics) -> TokenStream { +fn build_fields( + fields: &syn::Fields, + generics: &mut syn::Generics, + rename: &dyn Fn(syn::LitStr) -> syn::LitStr, +) -> syn::Result { let path = crate_name(); match fields { Fields::Named(fields) => { add_field_bounds(generics, &fields.named); - let list = named_fields(&fields.named); - quote!(#path::Fields::Named(vec![#(#list),*])) + let list = named_fields(&fields.named, rename)?; + Ok(quote!(#path::Fields::Named(vec![#(#list),*]))) } Fields::Unnamed(fields) => { add_field_bounds(generics, &fields.unnamed); let list = unnamed_fields(&fields.unnamed); - quote!(#path::Fields::Unnamed(vec![#(#list),*])) + Ok(quote!(#path::Fields::Unnamed(vec![#(#list),*]))) } - Fields::Unit => quote!(#path::Fields::Unit), + Fields::Unit => Ok(quote!(#path::Fields::Unit)), } } -fn named_fields<'a, I>(fields: I) -> impl Iterator +fn named_fields<'a, I>( + fields: I, + rename: &dyn Fn(syn::LitStr) -> syn::LitStr, +) -> syn::Result> where I: IntoIterator, { let path = crate_name(); - fields.into_iter().map(move |f| { - let ty = &f.ty; - let ident = f - .ident - .as_ref() - .expect("named fields should have identifiers") - .to_string(); - let doc = format_docstrings(&f.attrs); - quote_spanned! { ty.span()=> #path::NamedField::new::<#ty>(#ident, #doc) } - }) + fields + .into_iter() + .map(move |f| { + let ty = &f.ty; + let ident = f + .ident + .as_ref() + .expect("named fields should have identifiers"); + + let name = syn::LitStr::new(&ident.to_string(), ident.span()); + + let doc = format_docstrings(&f.attrs); + let field = serde::Field::parse(&f.attrs)?; + let name = field.rename_field_or(name, rename); + Ok(quote_spanned! { ty.span()=> #path::NamedField::new::<#ty>(#name, #doc) }) + }) + .collect() } fn unnamed_fields<'a, I>(fields: I) -> impl Iterator @@ -253,20 +292,33 @@ where } /// Generate the `Reflect` implementation. -fn process_struct(input: &DeriveInput, s: &syn::DataStruct, common: DeriveCommon) -> TokenStream { +fn process_struct( + input: &DeriveInput, + s: &syn::DataStruct, + common: DeriveCommon, +) -> syn::Result { let DeriveCommon { doc, mut generics, format_type_name, + container, } = common; + // Validate that the attributes we parsed are compatible with a `struct` definition. + let container: serde::Struct = container.try_as_struct()?; + let type_name = &input.ident; let path = crate_name(); - let fields = build_fields(&s.fields, &mut generics); + let fields = build_fields( + &s.fields, + &mut generics, + &serde::RenameAll::visitor(container.rename_all), + )?; + let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); - quote! { + let ts = quote! { impl #impl_generics #path::Reflect for #type_name #ty_generics #where_clause { fn ty() -> #path::Type { #path::Type::aggregate( @@ -279,44 +331,67 @@ fn process_struct(input: &DeriveInput, s: &syn::DataStruct, common: DeriveCommon #format_type_name } } - } + }; + + Ok(ts) } //-------// // Enums // //-------// -fn process_enum(input: &DeriveInput, e: &syn::DataEnum, common: DeriveCommon) -> TokenStream { +fn process_enum( + input: &DeriveInput, + e: &syn::DataEnum, + common: DeriveCommon, +) -> syn::Result { let DeriveCommon { doc, mut generics, format_type_name, + container, } = common; + // Validate that the attributes we parsed are compatible with an `enum` definition. + let container: serde::Enum = container.as_enum(); + // TODO: For now, we just assume that identifiers are taken as-is. let type_name = &input.ident; let path = crate_name(); - let variants: Vec<_> = e + let renamer = serde::RenameAll::visitor(container.rename_all); + let variants = e .variants .iter() - .map(|v| { + .map(|v| -> syn::Result { let doc = format_docstrings(&v.attrs); - let ident = v.ident.to_string(); + let name = syn::LitStr::new(&v.ident.to_string(), v.ident.span()); + let attrs = serde::Variant::parse(&v.attrs)?; - let fields = build_fields(&v.fields, &mut generics); + let fields = build_fields(&v.fields, &mut generics, &attrs.renamer())?; - quote!(#path::Variant::new(#ident, #fields, #doc)) + // Rename the variant as needed. + let name = attrs.rename_variant_or(name, &renamer); + Ok(quote!(#path::Variant::new(#name, #fields, #doc))) }) - .collect(); + .collect::>>()?; + + // Build the enum representation. + let enum_repr = match container.enum_repr { + serde::EnumRepr::External => quote!(#path::EnumRepr::External), + serde::EnumRepr::Internal { tag } => quote!(#path::EnumRepr::Internal { tag: #tag }), + serde::EnumRepr::Adjacent { tag, content } => { + quote!(#path::EnumRepr::Adjacent { tag: #tag, content: #content }) + } + }; let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); - quote! { + let ts = quote! { impl #impl_generics #path::Reflect for #type_name #ty_generics #where_clause { fn ty() -> #path::Type { #path::Type::enum_( - #path::EnumRepr::External, + #enum_repr, [#(#variants),*], #doc, ) @@ -326,7 +401,9 @@ fn process_enum(input: &DeriveInput, e: &syn::DataEnum, common: DeriveCommon) -> #format_type_name } } - } + }; + + Ok(ts) } //-------------// diff --git a/diskann-benchmark-runner-derive/src/serde.rs b/diskann-benchmark-runner-derive/src/serde.rs new file mode 100644 index 0000000000..cd034e1301 --- /dev/null +++ b/diskann-benchmark-runner-derive/src/serde.rs @@ -0,0 +1,294 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +//! These are modeled after the attributes documented in . + +#[must_use] +fn is_serde_attr(attr: &syn::Attribute) -> bool { + attr.path().is_ident("serde") +} + +fn identity(x: T) -> T { + x +} + +fn set_unique(opt: &mut Option, value: syn::LitStr, attr: &str) -> syn::Result<()> { + if opt.is_some() { + Err(syn::Error::new_spanned( + value, + format!("serde attribute `{}` found multiple times", attr), + )) + } else { + *opt = Some(value); + Ok(()) + } +} + +fn set_rename_all( + opt: &mut Option, + rename_all: RenameAll, + lit: syn::LitStr, +) -> syn::Result<()> { + if opt.is_some() { + Err(syn::Error::new_spanned( + lit, + "serde attribute `rename_all` found multiple times", + )) + } else { + *opt = Some(rename_all); + Ok(()) + } +} + +pub(crate) struct Struct { + pub(crate) rename_all: Option, +} + +pub(crate) struct Enum { + pub(crate) rename_all: Option, + pub(crate) enum_repr: EnumRepr, +} + +pub(crate) struct Container { + pub(crate) rename_all: Option, + pub(crate) enum_repr: EnumRepr, +} + +impl Container { + pub(crate) fn parse(attrs: &[syn::Attribute]) -> syn::Result { + let mut rename_all = Option::None; + let mut tag = Option::None; + let mut content = Option::None; + + for attr in attrs.iter().filter(|a| is_serde_attr(*a)) { + attr.parse_nested_meta(|meta| { + // serde(rename_all = "...") + if meta.path.is_ident("rename_all") { + let value: syn::LitStr = meta.value()?.parse()?; + let parsed = RenameAll::parse(&value)?; + set_rename_all(&mut rename_all, parsed, value)?; + return Ok(()); + } + + // serde(tag = "...") + if meta.path.is_ident("tag") { + let value: syn::LitStr = meta.value()?.parse()?; + set_unique(&mut tag, value, "tag")?; + return Ok(()); + } + + // serde(content = "...") + if meta.path.is_ident("content") { + let value: syn::LitStr = meta.value()?.parse()?; + set_unique(&mut content, value, "content")?; + return Ok(()); + } + + Err(meta.error("unsupported Serde attribute for Reflect")) + })?; + } + + Ok(Self { + rename_all, + enum_repr: EnumRepr::from_parsed(tag, content)?, + }) + } + + pub(crate) fn try_as_struct(self) -> syn::Result { + self.enum_repr.assert_struct_compatible()?; + Ok(Struct { + rename_all: self.rename_all, + }) + } + + pub(crate) fn as_enum(self) -> Enum { + Enum { + rename_all: self.rename_all, + enum_repr: self.enum_repr, + } + } +} + +pub(crate) enum EnumRepr { + External, + Internal { + tag: syn::LitStr, + }, + Adjacent { + tag: syn::LitStr, + content: syn::LitStr, + }, +} + +impl EnumRepr { + pub(crate) fn from_parsed( + tag: Option, + content: Option, + ) -> syn::Result { + match (tag, content) { + (None, None) => Ok(EnumRepr::External), + (Some(tag), None) => Ok(EnumRepr::Internal { tag }), + (Some(tag), Some(content)) => Ok(EnumRepr::Adjacent { tag, content }), + (None, Some(content)) => Err(syn::Error::new_spanned( + content, + "serde attribute `content` provided without a `tag`", + )), + } + } + + pub(crate) fn assert_struct_compatible(&self) -> syn::Result<()> { + match self { + Self::External => Ok(()), + Self::Internal { tag } => Err(syn::Error::new_spanned( + tag, + "serde attribute `tag` provided on a non-enum", + )), + Self::Adjacent { tag, .. } => Err(syn::Error::new_spanned( + tag, + "serde attributes `tag` and `content` provided on a non-enum", + )), + } + } +} + +#[derive(Debug, Clone, Copy)] +pub(crate) enum RenameAll { + Lower, + Snake, + Kebab, +} + +impl RenameAll { + fn supported() -> &'static str { + "\"lowercase\", \"snake_case\", or \"kebab-case\"" + } + + fn parse(lit: &syn::LitStr) -> syn::Result { + let s = lit.value(); + + match &*s { + "lowercase" => Ok(Self::Lower), + "snake_case" => Ok(Self::Snake), + "kebab-case" => Ok(Self::Kebab), + _ => Err(syn::Error::new_spanned( + lit, + format!( + "unsupported serde `rename_all` rule \"{}\" - expected one of {}", + s, + Self::supported() + ), + )), + } + } + + fn apply(&self, lit: syn::LitStr) -> syn::LitStr { + let s = lit.value(); + let s = match self { + Self::Lower => s.to_lowercase(), + Self::Snake => heck::AsSnakeCase(s).to_string(), + Self::Kebab => heck::AsKebabCase(s).to_string(), + }; + + syn::LitStr::new(&s, lit.span()) + } + + pub(crate) fn visitor(me: Option) -> impl Fn(syn::LitStr) -> syn::LitStr { + move |v| { + if let Some(rename_all) = me { + rename_all.apply(v) + } else { + v + } + } + } +} + +#[derive(Default)] +pub(crate) struct Variant { + rename: Option, + rename_all: Option, +} + +impl Variant { + pub(crate) fn parse(attrs: &[syn::Attribute]) -> syn::Result { + let mut me = Self::default(); + + for attr in attrs.iter().filter(|a| is_serde_attr(*a)) { + attr.parse_nested_meta(|meta| { + // serde(rename_all = "...") + if meta.path.is_ident("rename_all") { + let value: syn::LitStr = meta.value()?.parse()?; + let rename_all = RenameAll::parse(&value)?; + set_rename_all(&mut me.rename_all, rename_all, value)?; + return Ok(()); + } + + // serde(rename = "...") + if meta.path.is_ident("rename") { + let value: syn::LitStr = meta.value()?.parse()?; + set_unique(&mut me.rename, value, "rename")?; + return Ok(()); + } + + Err(meta.error("unsupported Serde attribute for Reflect")) + })?; + } + + Ok(me) + } + + /// Replace the name of this variant if directed by `serde(rename = "...")`. + /// + /// If the rename attribute exists, it takes precedence. Otherwise, the fallback is used. + pub(crate) fn rename_variant_or( + self, + value: syn::LitStr, + or_else: &dyn Fn(syn::LitStr) -> syn::LitStr, + ) -> syn::LitStr { + self.rename.map_or_else(|| or_else(value), identity) + } + + pub(crate) fn renamer(&self) -> impl Fn(syn::LitStr) -> syn::LitStr { + RenameAll::visitor(self.rename_all) + } +} + +#[derive(Default, Clone)] +pub(crate) struct Field { + rename: Option, +} + +impl Field { + pub(crate) fn parse(attrs: &[syn::Attribute]) -> syn::Result { + let mut me = Self::default(); + + for attr in attrs.iter().filter(|a| is_serde_attr(*a)) { + attr.parse_nested_meta(|meta| { + // serde(rename = "...") + if meta.path.is_ident("rename") { + let value: syn::LitStr = meta.value()?.parse()?; + set_unique(&mut me.rename, value, "rename")?; + return Ok(()); + } + + Err(meta.error("unsupported Serde attribute for Reflect")) + })?; + } + + Ok(me) + } + + /// Apply the renaming rules defined in `self`. + /// + /// If no renaming rules are present, instead invoke `or_else`. + pub(crate) fn rename_field_or( + self, + value: syn::LitStr, + or_else: &dyn Fn(syn::LitStr) -> syn::LitStr, + ) -> syn::LitStr { + let Self { rename } = self; + rename.map_or_else(|| or_else(value), identity) + } +} diff --git a/diskann-benchmark-runner/src/reflect.rs b/diskann-benchmark-runner/src/reflect.rs index fa41a2bb08..3d4bc2708c 100644 --- a/diskann-benchmark-runner/src/reflect.rs +++ b/diskann-benchmark-runner/src/reflect.rs @@ -843,6 +843,7 @@ pub enum Metric { /// An enum with no payloads. #[derive(Debug, Reflect)] +#[serde(rename_all = "kebab-case")] pub enum AdjacentEnum { SquaredL2, /// Let me see if this works @@ -1014,7 +1015,7 @@ mod tests { foo: usize, /// Bar bar: usize, - }; + } let r = Reflection::new::(); let ty = r.ty(); @@ -1106,7 +1107,7 @@ mod tests { /// It's a bee! B( /// Buzz buzz - B + B, ), } @@ -1145,6 +1146,7 @@ mod tests { /// All the enums. #[expect(unused)] #[derive(Reflect)] + #[serde(tag = "tag", content = "content")] enum All { /// A unit variant. Unit, From 4598bc16f19a963432ae8d4a05e1ac6471df3080 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Sat, 26 Sep 2026 08:53:30 -0700 Subject: [PATCH 07/21] Let's see if we can wire it up ... --- diskann-benchmark-runner-derive/src/lib.rs | 31 ++-- diskann-benchmark-runner-derive/src/serde.rs | 158 +++++++++++-------- 2 files changed, 102 insertions(+), 87 deletions(-) diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index 21bb7edf91..72ac73c54b 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -233,14 +233,14 @@ fn generate_type_name_body(input: &DeriveInput) -> TokenStream { fn build_fields( fields: &syn::Fields, generics: &mut syn::Generics, - rename: &dyn Fn(syn::LitStr) -> syn::LitStr, + rename_all: serde::RenameAll, ) -> syn::Result { let path = crate_name(); match fields { Fields::Named(fields) => { add_field_bounds(generics, &fields.named); - let list = named_fields(&fields.named, rename)?; + let list = named_fields(&fields.named, rename_all)?; Ok(quote!(#path::Fields::Named(vec![#(#list),*]))) } Fields::Unnamed(fields) => { @@ -252,10 +252,7 @@ fn build_fields( } } -fn named_fields<'a, I>( - fields: I, - rename: &dyn Fn(syn::LitStr) -> syn::LitStr, -) -> syn::Result> +fn named_fields<'a, I>(fields: I, rename_all: serde::RenameAll) -> syn::Result> where I: IntoIterator, { @@ -273,7 +270,7 @@ where let doc = format_docstrings(&f.attrs); let field = serde::Field::parse(&f.attrs)?; - let name = field.rename_field_or(name, rename); + let name = field.rename_field_or(name, rename_all); Ok(quote_spanned! { ty.span()=> #path::NamedField::new::<#ty>(#name, #doc) }) }) .collect() @@ -305,16 +302,12 @@ fn process_struct( } = common; // Validate that the attributes we parsed are compatible with a `struct` definition. - let container: serde::Struct = container.try_as_struct()?; + let serde::Struct { rename_all } = container.try_as_struct()?; let type_name = &input.ident; let path = crate_name(); - let fields = build_fields( - &s.fields, - &mut generics, - &serde::RenameAll::visitor(container.rename_all), - )?; + let fields = build_fields(&s.fields, &mut generics, rename_all)?; let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); @@ -353,13 +346,15 @@ fn process_enum( } = common; // Validate that the attributes we parsed are compatible with an `enum` definition. - let container: serde::Enum = container.as_enum(); + let serde::Enum { + rename_all, + enum_repr, + } = container.as_enum(); // TODO: For now, we just assume that identifiers are taken as-is. let type_name = &input.ident; let path = crate_name(); - let renamer = serde::RenameAll::visitor(container.rename_all); let variants = e .variants .iter() @@ -368,16 +363,16 @@ fn process_enum( let name = syn::LitStr::new(&v.ident.to_string(), v.ident.span()); let attrs = serde::Variant::parse(&v.attrs)?; - let fields = build_fields(&v.fields, &mut generics, &attrs.renamer())?; + let fields = build_fields(&v.fields, &mut generics, attrs.field_rename_all())?; // Rename the variant as needed. - let name = attrs.rename_variant_or(name, &renamer); + let name = attrs.rename_variant_or(name, rename_all); Ok(quote!(#path::Variant::new(#name, #fields, #doc))) }) .collect::>>()?; // Build the enum representation. - let enum_repr = match container.enum_repr { + let enum_repr = match enum_repr { serde::EnumRepr::External => quote!(#path::EnumRepr::External), serde::EnumRepr::Internal { tag } => quote!(#path::EnumRepr::Internal { tag: #tag }), serde::EnumRepr::Adjacent { tag, content } => { diff --git a/diskann-benchmark-runner-derive/src/serde.rs b/diskann-benchmark-runner-derive/src/serde.rs index cd034e1301..4a6539d973 100644 --- a/diskann-benchmark-runner-derive/src/serde.rs +++ b/diskann-benchmark-runner-derive/src/serde.rs @@ -26,39 +26,23 @@ fn set_unique(opt: &mut Option, value: syn::LitStr, attr: &str) -> } } -fn set_rename_all( - opt: &mut Option, - rename_all: RenameAll, - lit: syn::LitStr, -) -> syn::Result<()> { - if opt.is_some() { - Err(syn::Error::new_spanned( - lit, - "serde attribute `rename_all` found multiple times", - )) - } else { - *opt = Some(rename_all); - Ok(()) - } -} - pub(crate) struct Struct { - pub(crate) rename_all: Option, + pub(crate) rename_all: RenameAll, } pub(crate) struct Enum { - pub(crate) rename_all: Option, + pub(crate) rename_all: RenameAll, pub(crate) enum_repr: EnumRepr, } pub(crate) struct Container { - pub(crate) rename_all: Option, + pub(crate) rename_all: RenameAll, pub(crate) enum_repr: EnumRepr, } impl Container { pub(crate) fn parse(attrs: &[syn::Attribute]) -> syn::Result { - let mut rename_all = Option::None; + let mut rename_all = RenameAll::None; let mut tag = Option::None; let mut content = Option::None; @@ -67,8 +51,7 @@ impl Container { // serde(rename_all = "...") if meta.path.is_ident("rename_all") { let value: syn::LitStr = meta.value()?.parse()?; - let parsed = RenameAll::parse(&value)?; - set_rename_all(&mut rename_all, parsed, value)?; + rename_all.parse_in(value)?; return Ok(()); } @@ -153,8 +136,10 @@ impl EnumRepr { } } -#[derive(Debug, Clone, Copy)] +#[derive(Default, Debug, Clone, Copy, PartialEq)] pub(crate) enum RenameAll { + #[default] + None, Lower, Snake, Kebab, @@ -165,42 +150,85 @@ impl RenameAll { "\"lowercase\", \"snake_case\", or \"kebab-case\"" } - fn parse(lit: &syn::LitStr) -> syn::Result { - let s = lit.value(); - - match &*s { - "lowercase" => Ok(Self::Lower), - "snake_case" => Ok(Self::Snake), - "kebab-case" => Ok(Self::Kebab), - _ => Err(syn::Error::new_spanned( - lit, - format!( - "unsupported serde `rename_all` rule \"{}\" - expected one of {}", + fn parse(s: &str) -> Option { + match s { + "lowercase" => Some(Self::Lower), + "snake_case" => Some(Self::Snake), + "kebab-case" => Some(Self::Kebab), + _ => None, + } + } + + fn parse_in(&mut self, s: syn::LitStr) -> syn::Result { + if *self != Self::None { + Err(syn::Error::new_spanned( + s, + "serde attribute `rename_all` found multiple times", + )) + } else { + let value = s.value(); + match Self::parse(&value) { + Some(me) => Ok(me), + None => Err(syn::Error::new_spanned( s, - Self::supported() - ), - )), + format!( + "unsupported serde `rename_all` rule \"{}\" - expected one of {}", + value, + Self::supported() + ), + )), + } } } - fn apply(&self, lit: syn::LitStr) -> syn::LitStr { - let s = lit.value(); - let s = match self { - Self::Lower => s.to_lowercase(), - Self::Snake => heck::AsSnakeCase(s).to_string(), - Self::Kebab => heck::AsKebabCase(s).to_string(), - }; + /// These methods are taken from the `serde_derive` internals as they need to match. + /// + /// See: + fn apply_to_variant_str(&self, variant: &str) -> String { + match self { + Self::None => variant.to_owned(), + Self::Lower => variant.to_ascii_lowercase(), + Self::Snake => { + let mut snake = String::new(); + for (i, ch) in variant.char_indices() { + if i > 0 && ch.is_uppercase() { + snake.push('_'); + } + snake.push(ch.to_ascii_lowercase()); + } + snake + } + Self::Kebab => (Self::Snake) + .apply_to_variant_str(variant) + .replace('_', "-"), + } + } - syn::LitStr::new(&s, lit.span()) + pub(crate) fn apply_to_variant(&self, variant: syn::LitStr) -> syn::LitStr { + if *self == Self::None { + variant + } else { + syn::LitStr::new(&self.apply_to_variant_str(&variant.value()), variant.span()) + } } - pub(crate) fn visitor(me: Option) -> impl Fn(syn::LitStr) -> syn::LitStr { - move |v| { - if let Some(rename_all) = me { - rename_all.apply(v) - } else { - v - } + /// These methods are taken from the `serde_derive` internals as they need to match. + /// + /// Since Rust field are generally in lower snake case, there's less work to do. + /// + /// See: + fn apply_to_field_str(&self, field: &str) -> String { + match self { + Self::None | Self::Lower | Self::Snake => field.to_owned(), + Self::Kebab => field.replace('_', "-"), + } + } + + pub(crate) fn apply_to_field(&self, variant: syn::LitStr) -> syn::LitStr { + if *self == Self::None { + variant + } else { + syn::LitStr::new(&self.apply_to_field_str(&variant.value()), variant.span()) } } } @@ -208,7 +236,7 @@ impl RenameAll { #[derive(Default)] pub(crate) struct Variant { rename: Option, - rename_all: Option, + rename_all: RenameAll, } impl Variant { @@ -220,8 +248,7 @@ impl Variant { // serde(rename_all = "...") if meta.path.is_ident("rename_all") { let value: syn::LitStr = meta.value()?.parse()?; - let rename_all = RenameAll::parse(&value)?; - set_rename_all(&mut me.rename_all, rename_all, value)?; + me.rename_all.parse_in(value)?; return Ok(()); } @@ -242,16 +269,13 @@ impl Variant { /// Replace the name of this variant if directed by `serde(rename = "...")`. /// /// If the rename attribute exists, it takes precedence. Otherwise, the fallback is used. - pub(crate) fn rename_variant_or( - self, - value: syn::LitStr, - or_else: &dyn Fn(syn::LitStr) -> syn::LitStr, - ) -> syn::LitStr { - self.rename.map_or_else(|| or_else(value), identity) + pub(crate) fn rename_variant_or(self, variant: syn::LitStr, or_else: RenameAll) -> syn::LitStr { + self.rename + .map_or_else(|| or_else.apply_to_variant(variant), identity) } - pub(crate) fn renamer(&self) -> impl Fn(syn::LitStr) -> syn::LitStr { - RenameAll::visitor(self.rename_all) + pub(crate) fn field_rename_all(&self) -> RenameAll { + self.rename_all } } @@ -283,12 +307,8 @@ impl Field { /// Apply the renaming rules defined in `self`. /// /// If no renaming rules are present, instead invoke `or_else`. - pub(crate) fn rename_field_or( - self, - value: syn::LitStr, - or_else: &dyn Fn(syn::LitStr) -> syn::LitStr, - ) -> syn::LitStr { + pub(crate) fn rename_field_or(self, field: syn::LitStr, or_else: RenameAll) -> syn::LitStr { let Self { rename } = self; - rename.map_or_else(|| or_else(value), identity) + rename.map_or_else(|| or_else.apply_to_field(field), identity) } } From 0deff777f1a7f17690fce8da027c7bf57b5d626b Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Sat, 26 Sep 2026 09:09:45 -0700 Subject: [PATCH 08/21] Refactor --- diskann-benchmark-runner-derive/src/lib.rs | 18 +- diskann-benchmark-runner/src/reflect.rs | 1210 ----------------- diskann-benchmark-runner/src/reflect/mod.rs | 1207 ++++++++++++++++ .../src/reflect/render.rs | 302 ++++ diskann-benchmark-runner/src/reflect/tree.rs | 385 ++++++ 5 files changed, 1903 insertions(+), 1219 deletions(-) delete mode 100644 diskann-benchmark-runner/src/reflect.rs create mode 100644 diskann-benchmark-runner/src/reflect/mod.rs create mode 100644 diskann-benchmark-runner/src/reflect/render.rs create mode 100644 diskann-benchmark-runner/src/reflect/tree.rs diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index 72ac73c54b..e7906a6cf4 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -241,14 +241,14 @@ fn build_fields( Fields::Named(fields) => { add_field_bounds(generics, &fields.named); let list = named_fields(&fields.named, rename_all)?; - Ok(quote!(#path::Fields::Named(vec![#(#list),*]))) + Ok(quote!(#path::tree::Fields::Named(vec![#(#list),*]))) } Fields::Unnamed(fields) => { add_field_bounds(generics, &fields.unnamed); let list = unnamed_fields(&fields.unnamed); - Ok(quote!(#path::Fields::Unnamed(vec![#(#list),*]))) + Ok(quote!(#path::tree::Fields::Unnamed(vec![#(#list),*]))) } - Fields::Unit => Ok(quote!(#path::Fields::Unit)), + Fields::Unit => Ok(quote!(#path::tree::Fields::Unit)), } } @@ -271,7 +271,7 @@ where let doc = format_docstrings(&f.attrs); let field = serde::Field::parse(&f.attrs)?; let name = field.rename_field_or(name, rename_all); - Ok(quote_spanned! { ty.span()=> #path::NamedField::new::<#ty>(#name, #doc) }) + Ok(quote_spanned! { ty.span()=> #path::tree::NamedField::new::<#ty>(#name, #doc) }) }) .collect() } @@ -284,7 +284,7 @@ where fields.into_iter().map(move |f| { let ty = &f.ty; let doc = format_docstrings(&f.attrs); - quote_spanned! { ty.span()=> #path::UnnamedField::new::<#ty>(#doc) } + quote_spanned! { ty.span()=> #path::tree::UnnamedField::new::<#ty>(#doc) } }) } @@ -367,16 +367,16 @@ fn process_enum( // Rename the variant as needed. let name = attrs.rename_variant_or(name, rename_all); - Ok(quote!(#path::Variant::new(#name, #fields, #doc))) + Ok(quote!(#path::tree::Variant::new(#name, #fields, #doc))) }) .collect::>>()?; // Build the enum representation. let enum_repr = match enum_repr { - serde::EnumRepr::External => quote!(#path::EnumRepr::External), - serde::EnumRepr::Internal { tag } => quote!(#path::EnumRepr::Internal { tag: #tag }), + serde::EnumRepr::External => quote!(#path::tree::EnumRepr::External), + serde::EnumRepr::Internal { tag } => quote!(#path::tree::EnumRepr::Internal { tag: #tag }), serde::EnumRepr::Adjacent { tag, content } => { - quote!(#path::EnumRepr::Adjacent { tag: #tag, content: #content }) + quote!(#path::tree::EnumRepr::Adjacent { tag: #tag, content: #content }) } }; diff --git a/diskann-benchmark-runner/src/reflect.rs b/diskann-benchmark-runner/src/reflect.rs deleted file mode 100644 index 3d4bc2708c..0000000000 --- a/diskann-benchmark-runner/src/reflect.rs +++ /dev/null @@ -1,1210 +0,0 @@ -/* - * Copyright (c) Microsoft Corporation. - * Licensed under the MIT license. - */ - -use std::{ - any::TypeId, - borrow::Cow, - fmt::{self, Write}, -}; - -pub use diskann_benchmark_runner_derive::Reflect; - -use crate::utils::fmt::Quote; - -const INDENT: usize = 2; - -pub trait Reflect: 'static { - fn ty() -> Type; - fn format_type_name(f: &mut dyn Write) -> fmt::Result; -} - -#[derive(Clone, Copy)] -pub struct Reflection { - reflection: &'static dyn internal::Reflect, -} - -impl Reflection { - pub const fn new() -> Self - where - T: Reflect, - { - Self { - reflection: &internal::Wrapper::::INSTANCE, - } - } - - pub fn ty(&self) -> Type { - self.reflection.ty() - } - - pub fn type_name(&self) -> TypeName { - TypeName(*self) - } - - pub fn type_id(&self) -> TypeId { - self.reflection.type_id() - } - - pub fn render(&self) -> Render { - Render(*self) - } -} - -impl fmt::Debug for Reflection { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - f.debug_struct("Reflection") - .field("type_name", &self.type_name()) - .finish_non_exhaustive() - } -} - -pub type Doc = Cow<'static, str>; - -pub struct TypeName(Reflection); - -impl TypeName { - fn format_type_name(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.reflection.format_type_name(f) - } -} - -impl std::fmt::Debug for TypeName { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.format_type_name(f) - } -} - -impl std::fmt::Display for TypeName { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.format_type_name(f) - } -} - -pub struct Render(Reflection); - -impl std::fmt::Display for Render { - fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - let mut r = Renderer::new(f, 3); - r.render_subject(self.0) - } -} - -#[derive(Debug)] -pub enum Type { - Primitive(Primitive), - Aggregate(Aggregate), - Enum(Enum), - Sequence(Sequence), -} - -impl Type { - pub fn primitive(doc: Option) -> Self { - Self::from(Primitive::new(doc)) - } - - pub fn aggregate(fields: Fields, doc: Option) -> Self { - Self::from(Aggregate::new(fields, doc)) - } - - pub fn enum_( - repr: EnumRepr, - variants: impl IntoIterator, - doc: Option, - ) -> Self { - Self::from(Enum::new(repr, variants, doc)) - } - - pub fn sequence(doc: Option) -> Self - where - T: Reflect, - { - Self::from(Sequence::new::(doc)) - } - - pub fn doc(&self) -> Option<&str> { - match self { - Self::Primitive(p) => p.doc(), - Self::Aggregate(a) => a.doc(), - Self::Enum(e) => e.doc(), - Self::Sequence(s) => s.doc(), - } - } - - fn has_body(&self) -> bool { - match self { - Self::Primitive(_) => false, - Self::Aggregate(a) => a.has_body(), - Self::Enum(e) => e.has_body(), - Self::Sequence(_) => true, - } - } - - fn as_aggregate(&self) -> Option<&Aggregate> { - if let Self::Aggregate(aggregate) = self { - Some(aggregate) - } else { - None - } - } - - fn as_enum(&self) -> Option<&Enum> { - if let Self::Enum(enum_) = self { - Some(enum_) - } else { - None - } - } -} - -impl From for Type { - fn from(primitive: Primitive) -> Self { - Self::Primitive(primitive) - } -} - -impl From for Type { - fn from(aggergate: Aggregate) -> Self { - Self::Aggregate(aggergate) - } -} - -impl From for Type { - fn from(e: Enum) -> Self { - Self::Enum(e) - } -} - -impl From for Type { - fn from(s: Sequence) -> Self { - Self::Sequence(s) - } -} - -//-----------// -// Primitive // -//-----------// - -#[derive(Debug)] -pub struct Primitive { - doc: Option, -} - -impl Primitive { - pub fn new(doc: Option) -> Self { - Self { doc } - } - - fn doc(&self) -> Option<&str> { - self.doc.as_deref() - } -} - -//-----------// -// Aggregate // -//-----------// - -#[derive(Debug)] -pub struct Aggregate { - fields: Fields, - doc: Option, -} - -impl Aggregate { - pub fn new(fields: Fields, doc: Option) -> Self { - Self { fields, doc } - } - - fn doc(&self) -> Option<&str> { - self.doc.as_deref() - } - - fn fields(&self) -> &Fields { - &self.fields - } - - fn has_body(&self) -> bool { - self.fields.has_body() - } -} - -#[derive(Debug)] -pub enum Fields { - Named(Vec), - Unnamed(Vec), - Unit, -} - -impl Fields { - fn has_body(&self) -> bool { - match self { - Self::Named(fields) => !fields.is_empty(), - Self::Unnamed(fields) => !fields.is_empty(), - Self::Unit => false, - } - } - - pub fn named(itr: impl IntoIterator) -> Self { - Self::Named(itr.into_iter().collect()) - } - - pub fn unnamed(itr: impl IntoIterator) -> Self { - Self::Unnamed(itr.into_iter().collect()) - } - - fn as_named(&self) -> Option<&[NamedField]> { - if let Self::Named(fields) = self { - Some(fields) - } else { - None - } - } - - fn as_unnamed(&self) -> Option<&[UnnamedField]> { - if let Self::Unnamed(fields) = self { - Some(fields) - } else { - None - } - } -} - -#[derive(Debug)] -pub struct NamedField { - name: &'static str, - field: Reflection, - doc: Option, -} - -impl NamedField { - pub fn new(name: &'static str, doc: Option) -> Self - where - T: Reflect, - { - Self { - name, - field: Reflection::new::(), - doc, - } - } - - fn doc(&self) -> Option<&str> { - self.doc.as_deref() - } -} - -#[derive(Debug)] -pub struct UnnamedField { - field: Reflection, - doc: Option, -} - -impl UnnamedField { - pub fn new(doc: Option) -> Self - where - T: Reflect, - { - Self { - field: Reflection::new::(), - doc, - } - } - - fn doc(&self) -> Option<&str> { - self.doc.as_deref() - } - - fn field(&self) -> Reflection { - self.field - } -} - -//------// -// Enum // -//------// - -#[derive(Debug)] -pub struct Enum { - repr: EnumRepr, - variants: Vec, - doc: Option, -} - -impl Enum { - pub fn new( - repr: EnumRepr, - variants: impl IntoIterator, - doc: Option, - ) -> Self { - Self { - repr, - variants: variants.into_iter().collect(), - doc, - } - } - - pub fn doc(&self) -> Option<&str> { - self.doc.as_deref() - } - - fn variants(&self) -> &[Variant] { - &self.variants - } - - fn has_body(&self) -> bool { - !self.variants.is_empty() - } -} - -#[derive(Debug)] -pub enum EnumRepr { - /// Enums are tagged as the key in a collection. - External, - - /// Enums are tagged by an internal field. - /// - /// This contains the name of the field that is used as the tag. - /// - /// ```json - /// { - /// "tag": "some-enum-tag", - /// "value": 10, - /// "members": [ - /// 1, - /// "world" - /// ] - /// } - /// ``` - Internal { tag: &'static str }, - - /// Enums have a separate tag and content payload. - /// - /// ```json - /// { - /// "tag": "some-enum-tag", - /// "content": { - /// "value": 10, - /// "members": [ - /// 1, - /// "world" - /// ] - /// } - /// } - /// ``` - Adjacent { - tag: &'static str, - content: &'static str, - }, -} - -#[derive(Debug)] -pub struct Variant { - name: &'static str, - fields: Fields, - doc: Option, -} - -impl Variant { - pub fn new(name: &'static str, fields: Fields, doc: Option) -> Self { - Self { name, fields, doc } - } - - fn name(&self) -> &'static str { - self.name - } - - fn doc(&self) -> Option<&str> { - self.doc.as_deref() - } -} - -//----------// -// Sequence // -//----------// - -#[derive(Debug)] -pub struct Sequence { - element: Reflection, - doc: Option, -} - -impl Sequence { - pub fn new(doc: Option) -> Self - where - T: Reflect, - { - Self { - element: Reflection::new::(), - doc, - } - } - - pub fn doc(&self) -> Option<&str> { - self.doc.as_deref() - } -} - -////////////// -// Renderer // -////////////// - -#[derive(Debug)] -struct Tagged { - ty: Type, - reflection: Reflection, -} - -impl Tagged { - fn new(reflection: Reflection) -> Self { - Self { - ty: reflection.ty(), - reflection, - } - } - - fn ty(&self) -> &Type { - &self.ty - } - - fn reflection(&self) -> Reflection { - self.reflection - } -} - -struct Renderer<'a> { - output: &'a mut dyn Write, - indent: usize, - depth: usize, - max_depth: usize, -} - -impl<'a> Renderer<'a> { - fn new(output: &'a mut dyn Write, max_depth: usize) -> Self { - Self { - output, - indent: 0, - depth: 0, - max_depth, - } - } - - fn at_bottom(&self) -> bool { - self.depth >= self.max_depth - } - - fn line(&mut self, display: D) -> fmt::Result - where - D: std::fmt::Display, - { - let indent = INDENT * self.indent; - write!(self.output, "{: >indent$}{}\n", "", display) - } - - fn blank(&mut self) -> fmt::Result { - self.output.write_char('\n') - } - - fn maybe_indent(&mut self, indent: bool, f: F) -> Result - where - F: FnOnce(&mut Self) -> Result, - { - if indent { - self.indent += 1; - } - let result = f(self); - if indent { - self.indent -= 1; - } - result - } - - fn indent(&mut self, f: F) -> Result - where - F: FnOnce(&mut Self) -> Result, - { - self.maybe_indent(true, f) - } - - fn next_with( - &mut self, - pre: impl FnOnce(&mut Self) -> fmt::Result, - body: impl FnOnce(&mut Self) -> fmt::Result, - post: impl FnOnce(&mut Self) -> fmt::Result, - ) -> fmt::Result { - if self.at_bottom() { - Ok(()) - } else { - pre(self)?; - self.depth += 1; - let result = self.indent(body); - self.depth -= 1; - post(self)?; - result - } - } - - fn next(&mut self, f: F) -> fmt::Result - where - F: FnOnce(&mut Self) -> fmt::Result, - { - self.next_with(|_| Ok(()), f, |_| Ok(())) - } - - fn render_doc(&mut self, s: Option<&str>) -> Result { - if let Some(s) = s { - let mut rendered = false; - for ln in s.lines() { - if ln.is_empty() { - self.blank()?; - } else { - self.line(ln)?; - } - - rendered = true; - } - Ok(rendered) - } else { - Ok(false) - } - } - - /// Return `true` if a nested type will be rendered. - fn will_render(&self, ty: &Type) -> bool { - !self.at_bottom() && ty.has_body() - } - - //-------// - // Types // - //-------// - - fn render_subject(&mut self, reflection: Reflection) -> fmt::Result { - self.line(reflection.type_name())?; - self.indent(|r| { - let tagged = Tagged::new(reflection); - let wrote_doc = r.render_doc(tagged.ty().doc())?; - - if wrote_doc && r.will_render(tagged.ty()) { - r.blank()?; - } - - r.render_body(&tagged) - }) - } - - fn render_body(&mut self, tagged: &Tagged) -> fmt::Result { - match tagged.ty() { - Type::Primitive(_) => Ok(()), - Type::Aggregate(aggregate) => self.render_aggregate(&aggregate), - Type::Enum(enum_) => self.render_enum(enum_), - Type::Sequence(sequence) => self.render_sequence(&sequence), - } - } - - fn render_aggregate(&mut self, aggregate: &Aggregate) -> fmt::Result { - self.render_fields(&aggregate.fields) - } - - fn render_enum(&mut self, enum_: &Enum) -> fmt::Result { - match enum_.repr { - EnumRepr::External => self.line("Representation: externally tagged")?, - EnumRepr::Internal { tag } => { - self.line(format_args!("Discriminant field: {}", Quote(tag)))? - } - EnumRepr::Adjacent { tag, content } => { - self.line(format_args!("Discriminant field: {}", Quote(tag)))?; - self.line(format_args!("Content field: {}", Quote(content)))?; - } - } - - self.blank()?; - self.line("Options:")?; - self.indent(|r| { - let mut first = true; - for variant in enum_.variants.iter() { - if !first { - r.blank()?; - } - - r.render_variant(variant)?; - first = false; - } - - Ok(()) - }) - } - - fn render_sequence(&mut self, sequence: &Sequence) -> fmt::Result { - self.next_with( - |r| r.line(format_args!("Elements: {}", sequence.element.type_name())), - |r| r.render_body(&Tagged::new(sequence.element)), - |_| Ok(()), - ) - } - - //--------// - // Fields // - //--------// - - fn render_fields(&mut self, fields: &Fields) -> fmt::Result { - match fields { - Fields::Named(named) => { - let mut first = true; - for field in named.iter() { - if !first { - self.blank()?; - } - self.render_named_field(field)?; - first = false; - } - } - Fields::Unnamed(unnamed) => { - for (i, field) in unnamed.iter().enumerate() { - if i != 0 { - self.blank()?; - } - self.render_unnamed_field(i, field)?; - } - } - Fields::Unit => {} - } - - Ok(()) - } - - fn render_named_field(&mut self, field: &NamedField) -> fmt::Result { - let tagged = Tagged::new(field.field); - let will_render_body = self.will_render(tagged.ty()); - - self.line(format_args!( - "{}: {}", - Quote(field.name), - tagged.reflection().type_name() - ))?; - - let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; - if rendered_doc && will_render_body { - self.blank()?; - } - - if will_render_body { - self.next(|r| r.render_body(&tagged))?; - } - Ok(()) - } - - fn render_unnamed_field(&mut self, index: usize, field: &UnnamedField) -> fmt::Result { - let tagged = Tagged::new(field.field); - let will_render_body = self.will_render(tagged.ty()); - - self.line(format_args!( - "{}: {}", - index, - tagged.reflection().type_name() - ))?; - - let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; - if rendered_doc && will_render_body { - self.blank()?; - } - - if will_render_body { - self.next(|r| r.render_body(&tagged))?; - } - Ok(()) - } - - //---------// - // Variant // - //---------// - - fn render_variant(&mut self, variant: &Variant) -> fmt::Result { - self.line(Quote(variant.name))?; - - self.indent(|r| { - let rendered_doc = r.render_doc(variant.doc())?; - - if rendered_doc && variant.fields.has_body() { - r.blank()?; - } - - r.render_fields(&variant.fields) - }) - } -} - -/////////////// -// Bootstrap // -/////////////// - -impl Reflect for std::marker::PhantomData -where - T: Reflect, -{ - fn ty() -> Type { - Type::primitive(None) - } - - fn format_type_name(f: &mut dyn Write) -> fmt::Result { - write!(f, "PhantomData<{}>", Reflection::new::().type_name()) - } -} - -macro_rules! primitive { - ($T:ty, $doc:literal, $type_name:literal) => { - impl Reflect for $T { - fn ty() -> Type { - Type::primitive(Some($doc.into())) - } - - fn format_type_name(f: &mut dyn Write) -> fmt::Result { - f.write_str($type_name) - } - } - }; -} - -primitive!(usize, "A system dependent unsigned integer", "usize"); -primitive!(u32, "A 32-bit unsigned integer", "u32"); - -primitive!(String, "A string", "string"); - -impl Reflect for Vec -where - T: Reflect, -{ - fn ty() -> Type { - Type::sequence::(Some("An ordered collection of elements".into())) - } - - fn format_type_name(f: &mut dyn Write) -> fmt::Result { - write!(f, "Vec<{}>", Reflection::new::().type_name()) - } -} - -#[derive(Reflect)] -pub struct UnitWithConst {} - -#[derive(Reflect)] -pub struct GenericBoundAdded { - uses_t: Vec, -} - -/// This is a test! -/// -/// Hello world! -#[derive(Reflect)] -pub struct Test { - /// This field affects this value. - a: usize, - - /// This field does something else. - b: usize, -} - -#[derive(Reflect)] -pub struct TestUnnamed( - /// Can I document this? - usize, - T, -); - -/// This is a nother test! -#[derive(Reflect)] -pub struct Test2 { - /// This field affects this value. - a: usize, - - /// This field doesn't have any names. - unnamed: TestUnnamed, - - other: Test, - - /// How are we going to compute distances? - metric: AdjacentEnum, - - /// These control a bunch of parameters. - seq: Vec, -} - -#[derive(Reflect)] -struct Wrapper { - /// Inner - a: T, -} - -/// An enum with no payloads. -#[derive(Debug, Clone, Copy, Reflect)] -pub enum Metric { - SquaredL2, - InnerProduct, - Cosine, -} - -/// An enum with no payloads. -#[derive(Debug, Reflect)] -#[serde(rename_all = "kebab-case")] -pub enum AdjacentEnum { - SquaredL2, - /// Let me see if this works - InnerProduct(u32), - - /// Compute the cosine similarity - Cosine { - /// Thos actually doesn't do anything. - test: String, - }, -} - -// impl Reflect for AdjacentEnum { -// fn ty() -> Type { -// Type::enum_( -// EnumRepr::Adjacent { -// tag: "enum-type", -// content: "content", -// }, -// [ -// Variant::new("squared-l2", Fields::Unit, None), -// Variant::new( -// "inner-product", -// Fields::unnamed([UnnamedField::new::(Some("testing".into()))]), -// Some("Inner Product with some payload".into()), -// ), -// Variant::new( -// "cosine", -// Fields::named([NamedField::new::("test", None)]), -// Some("Cosine Similarity".into()), -// ), -// ], -// Some("The similarity measure to use".into()), -// ) -// } -// -// fn format_type_name(f: &mut dyn Write) -> fmt::Result { -// f.write_str("Metric") -// } -// } - -////////////// -// Internal // -////////////// - -pub(crate) mod internal { - use std::marker::PhantomData; - - pub(crate) trait Reflect { - fn ty(&self) -> super::Type; - fn format_type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result; - fn type_id(&self) -> std::any::TypeId; - } - - pub(crate) struct Wrapper(PhantomData); - - impl Wrapper { - pub(crate) const INSTANCE: Self = Self::new(); - - pub(crate) const fn new() -> Self { - Self(PhantomData) - } - } - - impl Reflect for Wrapper - where - T: super::Reflect, - { - fn ty(&self) -> super::Type { - ::ty() - } - - fn format_type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result { - ::format_type_name(f) - } - - fn type_id(&self) -> std::any::TypeId { - std::any::TypeId::of::() - } - } -} - -/////////// -// Tests // -/////////// - -#[cfg(test)] -mod tests { - use super::*; - - use std::assert_matches; - - #[test] - fn test_unit() { - /// A unit struct. - #[derive(Reflect)] - struct Unit; - - let r = Reflection::new::(); - let ty = r.ty(); - - assert_eq!(r.ty().doc().unwrap(), "A unit struct."); - assert_eq!(r.type_name().to_string(), "Unit"); - assert!(!ty.has_body()); - } - - #[test] - fn test_unit_const_generic() { - /// A unit struct. - #[derive(Reflect)] - struct Unit; - - let r = Reflection::new::>(); - let ty = r.ty(); - - assert_eq!(ty.doc().unwrap(), "A unit struct."); - assert_eq!(r.type_name().to_string(), "Unit<10>"); - assert!(!ty.has_body()); - } - - #[test] - fn test_unit_const_generic_2() { - #[derive(Reflect)] - struct Unit; - - let r = Reflection::new::>(); - let ty = r.ty(); - - assert!(ty.doc().is_none()); - assert_eq!(r.type_name().to_string(), "Unit<10, 20>"); - assert!(!ty.has_body()); - } - - #[test] - fn test_empty_tuple_like() { - /// An empty tuple-like struct. - #[derive(Reflect)] - struct Empty(); - - let r = Reflection::new::(); - let ty = r.ty(); - - assert_eq!(ty.doc().unwrap(), "An empty tuple-like struct."); - assert_eq!(r.type_name().to_string(), "Empty"); - assert!(!ty.has_body()); - } - - #[test] - fn test_empty_struct_like() { - /// An empty struct. - #[derive(Reflect)] - struct Empty {} - - let r = Reflection::new::(); - let ty = r.ty(); - - assert_eq!(ty.doc().unwrap(), "An empty struct."); - assert_eq!(r.type_name().to_string(), "Empty"); - assert!(!ty.has_body()); - } - - #[test] - fn test_struct() { - /// A struct with two fields. - #[expect(unused)] - #[derive(Reflect)] - struct Woo { - /// Foo - foo: usize, - /// Bar - bar: usize, - } - - let r = Reflection::new::(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "A struct with two fields."); - assert_eq!(r.type_name().to_string(), "Woo"); - assert!(ty.has_body()); - - let f = ty.as_aggregate().unwrap().fields().as_named().unwrap(); - assert_eq!(f.len(), 2); - - assert_eq!(f[0].name, "foo"); - assert_eq!(f[0].doc().unwrap(), "Foo"); - assert_eq!(f[0].field.type_name().to_string(), "usize"); - - assert_eq!(f[1].name, "bar"); - assert_eq!(f[1].doc().unwrap(), "Bar"); - assert_eq!(f[1].field.type_name().to_string(), "usize"); - } - - #[test] - fn test_tuple1() { - /// A tuple with one field. - #[expect(unused)] - #[derive(Reflect)] - struct Tuple1( - /// Field 0. - usize, - ); - - let r = Reflection::new::(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "A tuple with one field."); - assert_eq!(r.type_name().to_string(), "Tuple1"); - - assert!(ty.has_body()); - - let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); - assert_eq!(f.len(), 1); - assert_eq!(f[0].doc().unwrap(), "Field 0."); - assert_eq!(f[0].field().type_name().to_string(), "usize"); - } - - #[test] - fn test_tuple2() { - #[expect(unused)] - #[derive(Reflect)] - struct Tuple2( - T, - /// Field 1. - Vec, - ); - - let r = Reflection::new::>(); - let ty = r.ty(); - assert!(ty.doc().is_none()); - assert_eq!(r.type_name().to_string(), "Tuple2"); - assert!(ty.has_body()); - - let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); - assert_eq!(f.len(), 2); - - assert!(f[0].doc().is_none()); - assert_eq!(f[0].field().type_name().to_string(), "usize"); - - assert_eq!(f[1].doc().unwrap(), "Field 1."); - assert_eq!(f[1].field().type_name().to_string(), "Vec"); - } - - #[test] - fn test_empty_enum() { - /// An empty enum. - #[derive(Reflect)] - enum Empty {} - - let r = Reflection::new::(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "An empty enum."); - assert_eq!(r.type_name().to_string(), "Empty"); - assert!(!ty.has_body()); - } - - #[test] - fn test_enum_with_generics() { - /// An enum with generics. - #[expect(unused)] - #[derive(Reflect)] - enum Either { - A(Vec), - /// It's a bee! - B( - /// Buzz buzz - B, - ), - } - - let r = Reflection::new::>(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "An enum with generics."); - assert_eq!(r.type_name().to_string(), "Either"); - assert!(ty.has_body()); - - let variants = ty.as_enum().unwrap().variants(); - assert_eq!(variants.len(), 2); - - // Variant 0 - assert!(variants[0].doc().is_none()); - assert_eq!(variants[0].name(), "A"); - assert!(variants[0].fields.has_body()); - - let f = variants[0].fields.as_unnamed().unwrap(); - assert_eq!(f.len(), 1); - assert!(f[0].doc().is_none()); - assert_eq!(f[0].field().type_name().to_string(), "Vec"); - - // Variant 1 - assert_eq!(variants[1].doc().unwrap(), "It's a bee!"); - assert_eq!(variants[1].name(), "B"); - assert!(variants[1].fields.has_body()); - - let f = variants[1].fields.as_unnamed().unwrap(); - assert_eq!(f.len(), 1); - assert_eq!(f[0].doc().unwrap(), "Buzz buzz"); - assert_eq!(f[0].field().type_name().to_string(), "string"); - } - - #[test] - fn test_enum_variants() { - /// All the enums. - #[expect(unused)] - #[derive(Reflect)] - #[serde(tag = "tag", content = "content")] - enum All { - /// A unit variant. - Unit, - /// A tuple-like variant. - Tuple( - /// Field 0. - usize, - String, - ), - /// Struct-like. - Struct { - foo: usize, - /// All the bars! - bar: u32, - }, - } - - let r = Reflection::new::(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "All the enums."); - assert_eq!(r.type_name().to_string(), "All"); - assert!(ty.has_body()); - - let variants = ty.as_enum().unwrap().variants(); - assert_eq!(variants.len(), 3); - - // Variant 0 - assert_eq!(variants[0].doc().unwrap(), "A unit variant."); - assert_eq!(variants[0].name(), "Unit"); - assert!(!variants[0].fields.has_body()); - assert_matches!(variants[0].fields, Fields::Unit); - - // Variant 1 - assert_eq!(variants[1].doc().unwrap(), "A tuple-like variant."); - assert_eq!(variants[1].name(), "Tuple"); - assert!(variants[1].fields.has_body()); - - let f = variants[1].fields.as_unnamed().unwrap(); - assert_eq!(f.len(), 2); - - assert_eq!(f[0].doc().unwrap(), "Field 0."); - assert_eq!(f[0].field.type_name().to_string(), "usize"); - - assert!(f[1].doc().is_none()); - assert_eq!(f[1].field.type_name().to_string(), "string"); - - // Variant 2 - assert_eq!(variants[2].name(), "Struct"); - assert_eq!(variants[2].doc().unwrap(), "Struct-like."); - let f = variants[2].fields.as_named().unwrap(); - assert_eq!(f.len(), 2); - - assert_eq!(f[0].name, "foo"); - assert!(f[0].doc().is_none()); - assert_eq!(f[0].field.type_name().to_string(), "usize"); - - assert_eq!(f[1].name, "bar"); - assert_eq!(f[1].doc().unwrap(), "All the bars!"); - assert_eq!(f[1].field.type_name().to_string(), "u32"); - } -} diff --git a/diskann-benchmark-runner/src/reflect/mod.rs b/diskann-benchmark-runner/src/reflect/mod.rs new file mode 100644 index 0000000000..fde153736f --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/mod.rs @@ -0,0 +1,1207 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::{ + any::TypeId, + fmt::{self, Write}, +}; + +pub use diskann_benchmark_runner_derive::Reflect; + +mod render; +pub mod tree; +pub use tree::Type; + +pub trait Reflect: 'static { + fn ty() -> Type; + fn format_type_name(f: &mut dyn Write) -> fmt::Result; +} + +#[derive(Clone, Copy)] +pub struct Reflection { + reflection: &'static dyn internal::Reflect, +} + +impl Reflection { + pub const fn new() -> Self + where + T: Reflect, + { + Self { + reflection: &internal::Wrapper::::INSTANCE, + } + } + + pub fn ty(&self) -> Type { + self.reflection.ty() + } + + pub fn type_name(&self) -> TypeName { + TypeName(*self) + } + + pub fn type_id(&self) -> TypeId { + self.reflection.type_id() + } + + pub fn render(&self) -> Render { + Render(*self) + } +} + +impl fmt::Debug for Reflection { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_struct("Reflection") + .field("type_name", &self.type_name()) + .finish_non_exhaustive() + } +} + +pub struct TypeName(Reflection); + +impl TypeName { + fn format_type_name(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.0.reflection.format_type_name(f) + } +} + +impl std::fmt::Debug for TypeName { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.format_type_name(f) + } +} + +impl std::fmt::Display for TypeName { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + self.format_type_name(f) + } +} + +pub struct Render(Reflection); + +impl std::fmt::Display for Render { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + let mut r = render::Renderer::new(f, 3); + r.render_subject(self.0) + } +} + +// #[derive(Debug)] +// pub enum Type { +// Primitive(Primitive), +// Aggregate(Aggregate), +// Enum(Enum), +// Sequence(Sequence), +// } +// +// impl Type { +// pub fn primitive(doc: Option) -> Self { +// Self::from(Primitive::new(doc)) +// } +// +// pub fn aggregate(fields: Fields, doc: Option) -> Self { +// Self::from(Aggregate::new(fields, doc)) +// } +// +// pub fn enum_( +// repr: EnumRepr, +// variants: impl IntoIterator, +// doc: Option, +// ) -> Self { +// Self::from(Enum::new(repr, variants, doc)) +// } +// +// pub fn sequence(doc: Option) -> Self +// where +// T: Reflect, +// { +// Self::from(Sequence::new::(doc)) +// } +// +// pub fn doc(&self) -> Option<&str> { +// match self { +// Self::Primitive(p) => p.doc(), +// Self::Aggregate(a) => a.doc(), +// Self::Enum(e) => e.doc(), +// Self::Sequence(s) => s.doc(), +// } +// } +// +// fn has_body(&self) -> bool { +// match self { +// Self::Primitive(_) => false, +// Self::Aggregate(a) => a.has_body(), +// Self::Enum(e) => e.has_body(), +// Self::Sequence(_) => true, +// } +// } +// +// fn as_aggregate(&self) -> Option<&Aggregate> { +// if let Self::Aggregate(aggregate) = self { +// Some(aggregate) +// } else { +// None +// } +// } +// +// fn as_enum(&self) -> Option<&Enum> { +// if let Self::Enum(enum_) = self { +// Some(enum_) +// } else { +// None +// } +// } +// } +// +// impl From for Type { +// fn from(primitive: Primitive) -> Self { +// Self::Primitive(primitive) +// } +// } +// +// impl From for Type { +// fn from(aggergate: Aggregate) -> Self { +// Self::Aggregate(aggergate) +// } +// } +// +// impl From for Type { +// fn from(e: Enum) -> Self { +// Self::Enum(e) +// } +// } +// +// impl From for Type { +// fn from(s: Sequence) -> Self { +// Self::Sequence(s) +// } +// } +// +// //-----------// +// // Primitive // +// //-----------// +// +// #[derive(Debug)] +// pub struct Primitive { +// doc: Option, +// } +// +// impl Primitive { +// pub fn new(doc: Option) -> Self { +// Self { doc } +// } +// +// fn doc(&self) -> Option<&str> { +// self.doc.as_deref() +// } +// } +// +// //-----------// +// // Aggregate // +// //-----------// +// +// #[derive(Debug)] +// pub struct Aggregate { +// fields: Fields, +// doc: Option, +// } +// +// impl Aggregate { +// pub fn new(fields: Fields, doc: Option) -> Self { +// Self { fields, doc } +// } +// +// fn doc(&self) -> Option<&str> { +// self.doc.as_deref() +// } +// +// fn fields(&self) -> &Fields { +// &self.fields +// } +// +// fn has_body(&self) -> bool { +// self.fields.has_body() +// } +// } +// +// #[derive(Debug)] +// pub enum Fields { +// Named(Vec), +// Unnamed(Vec), +// Unit, +// } +// +// impl Fields { +// fn has_body(&self) -> bool { +// match self { +// Self::Named(fields) => !fields.is_empty(), +// Self::Unnamed(fields) => !fields.is_empty(), +// Self::Unit => false, +// } +// } +// +// pub fn named(itr: impl IntoIterator) -> Self { +// Self::Named(itr.into_iter().collect()) +// } +// +// pub fn unnamed(itr: impl IntoIterator) -> Self { +// Self::Unnamed(itr.into_iter().collect()) +// } +// +// fn as_named(&self) -> Option<&[NamedField]> { +// if let Self::Named(fields) = self { +// Some(fields) +// } else { +// None +// } +// } +// +// fn as_unnamed(&self) -> Option<&[UnnamedField]> { +// if let Self::Unnamed(fields) = self { +// Some(fields) +// } else { +// None +// } +// } +// } +// +// #[derive(Debug)] +// pub struct NamedField { +// name: &'static str, +// field: Reflection, +// doc: Option, +// } +// +// impl NamedField { +// pub fn new(name: &'static str, doc: Option) -> Self +// where +// T: Reflect, +// { +// Self { +// name, +// field: Reflection::new::(), +// doc, +// } +// } +// +// fn doc(&self) -> Option<&str> { +// self.doc.as_deref() +// } +// } +// +// #[derive(Debug)] +// pub struct UnnamedField { +// field: Reflection, +// doc: Option, +// } +// +// impl UnnamedField { +// pub fn new(doc: Option) -> Self +// where +// T: Reflect, +// { +// Self { +// field: Reflection::new::(), +// doc, +// } +// } +// +// fn doc(&self) -> Option<&str> { +// self.doc.as_deref() +// } +// +// fn field(&self) -> Reflection { +// self.field +// } +// } +// +// //------// +// // Enum // +// //------// +// +// #[derive(Debug)] +// pub struct Enum { +// repr: EnumRepr, +// variants: Vec, +// doc: Option, +// } +// +// impl Enum { +// pub fn new( +// repr: EnumRepr, +// variants: impl IntoIterator, +// doc: Option, +// ) -> Self { +// Self { +// repr, +// variants: variants.into_iter().collect(), +// doc, +// } +// } +// +// pub fn doc(&self) -> Option<&str> { +// self.doc.as_deref() +// } +// +// fn variants(&self) -> &[Variant] { +// &self.variants +// } +// +// fn has_body(&self) -> bool { +// !self.variants.is_empty() +// } +// } +// +// #[derive(Debug)] +// pub enum EnumRepr { +// /// Enums are tagged as the key in a collection. +// External, +// +// /// Enums are tagged by an internal field. +// /// +// /// This contains the name of the field that is used as the tag. +// /// +// /// ```json +// /// { +// /// "tag": "some-enum-tag", +// /// "value": 10, +// /// "members": [ +// /// 1, +// /// "world" +// /// ] +// /// } +// /// ``` +// Internal { tag: &'static str }, +// +// /// Enums have a separate tag and content payload. +// /// +// /// ```json +// /// { +// /// "tag": "some-enum-tag", +// /// "content": { +// /// "value": 10, +// /// "members": [ +// /// 1, +// /// "world" +// /// ] +// /// } +// /// } +// /// ``` +// Adjacent { +// tag: &'static str, +// content: &'static str, +// }, +// } +// +// #[derive(Debug)] +// pub struct Variant { +// name: &'static str, +// fields: Fields, +// doc: Option, +// } +// +// impl Variant { +// pub fn new(name: &'static str, fields: Fields, doc: Option) -> Self { +// Self { name, fields, doc } +// } +// +// fn name(&self) -> &'static str { +// self.name +// } +// +// fn doc(&self) -> Option<&str> { +// self.doc.as_deref() +// } +// } +// +// //----------// +// // Sequence // +// //----------// +// +// #[derive(Debug)] +// pub struct Sequence { +// element: Reflection, +// doc: Option, +// } +// +// impl Sequence { +// pub fn new(doc: Option) -> Self +// where +// T: Reflect, +// { +// Self { +// element: Reflection::new::(), +// doc, +// } +// } +// +// pub fn doc(&self) -> Option<&str> { +// self.doc.as_deref() +// } +// } +// +// ////////////// +// // Renderer // +// ////////////// +// +// #[derive(Debug)] +// struct Tagged { +// ty: Type, +// reflection: Reflection, +// } +// +// impl Tagged { +// fn new(reflection: Reflection) -> Self { +// Self { +// ty: reflection.ty(), +// reflection, +// } +// } +// +// fn ty(&self) -> &Type { +// &self.ty +// } +// +// fn reflection(&self) -> Reflection { +// self.reflection +// } +// } +// +// struct Renderer<'a> { +// output: &'a mut dyn Write, +// indent: usize, +// depth: usize, +// max_depth: usize, +// } +// +// impl<'a> Renderer<'a> { +// fn new(output: &'a mut dyn Write, max_depth: usize) -> Self { +// Self { +// output, +// indent: 0, +// depth: 0, +// max_depth, +// } +// } +// +// fn at_bottom(&self) -> bool { +// self.depth >= self.max_depth +// } +// +// fn line(&mut self, display: D) -> fmt::Result +// where +// D: std::fmt::Display, +// { +// let indent = INDENT * self.indent; +// write!(self.output, "{: >indent$}{}\n", "", display) +// } +// +// fn blank(&mut self) -> fmt::Result { +// self.output.write_char('\n') +// } +// +// fn maybe_indent(&mut self, indent: bool, f: F) -> Result +// where +// F: FnOnce(&mut Self) -> Result, +// { +// if indent { +// self.indent += 1; +// } +// let result = f(self); +// if indent { +// self.indent -= 1; +// } +// result +// } +// +// fn indent(&mut self, f: F) -> Result +// where +// F: FnOnce(&mut Self) -> Result, +// { +// self.maybe_indent(true, f) +// } +// +// fn next_with( +// &mut self, +// pre: impl FnOnce(&mut Self) -> fmt::Result, +// body: impl FnOnce(&mut Self) -> fmt::Result, +// post: impl FnOnce(&mut Self) -> fmt::Result, +// ) -> fmt::Result { +// if self.at_bottom() { +// Ok(()) +// } else { +// pre(self)?; +// self.depth += 1; +// let result = self.indent(body); +// self.depth -= 1; +// post(self)?; +// result +// } +// } +// +// fn next(&mut self, f: F) -> fmt::Result +// where +// F: FnOnce(&mut Self) -> fmt::Result, +// { +// self.next_with(|_| Ok(()), f, |_| Ok(())) +// } +// +// fn render_doc(&mut self, s: Option<&str>) -> Result { +// if let Some(s) = s { +// let mut rendered = false; +// for ln in s.lines() { +// if ln.is_empty() { +// self.blank()?; +// } else { +// self.line(ln)?; +// } +// +// rendered = true; +// } +// Ok(rendered) +// } else { +// Ok(false) +// } +// } +// +// /// Return `true` if a nested type will be rendered. +// fn will_render(&self, ty: &Type) -> bool { +// !self.at_bottom() && ty.has_body() +// } +// +// //-------// +// // Types // +// //-------// +// +// fn render_subject(&mut self, reflection: Reflection) -> fmt::Result { +// self.line(reflection.type_name())?; +// self.indent(|r| { +// let tagged = Tagged::new(reflection); +// let wrote_doc = r.render_doc(tagged.ty().doc())?; +// +// if wrote_doc && r.will_render(tagged.ty()) { +// r.blank()?; +// } +// +// r.render_body(&tagged) +// }) +// } +// +// fn render_body(&mut self, tagged: &Tagged) -> fmt::Result { +// match tagged.ty() { +// Type::Primitive(_) => Ok(()), +// Type::Aggregate(aggregate) => self.render_aggregate(&aggregate), +// Type::Enum(enum_) => self.render_enum(enum_), +// Type::Sequence(sequence) => self.render_sequence(&sequence), +// } +// } +// +// fn render_aggregate(&mut self, aggregate: &Aggregate) -> fmt::Result { +// self.render_fields(&aggregate.fields) +// } +// +// fn render_enum(&mut self, enum_: &Enum) -> fmt::Result { +// match enum_.repr { +// EnumRepr::External => self.line("Representation: externally tagged")?, +// EnumRepr::Internal { tag } => { +// self.line(format_args!("Discriminant field: {}", Quote(tag)))? +// } +// EnumRepr::Adjacent { tag, content } => { +// self.line(format_args!("Discriminant field: {}", Quote(tag)))?; +// self.line(format_args!("Content field: {}", Quote(content)))?; +// } +// } +// +// self.blank()?; +// self.line("Options:")?; +// self.indent(|r| { +// let mut first = true; +// for variant in enum_.variants.iter() { +// if !first { +// r.blank()?; +// } +// +// r.render_variant(variant)?; +// first = false; +// } +// +// Ok(()) +// }) +// } +// +// fn render_sequence(&mut self, sequence: &Sequence) -> fmt::Result { +// self.next_with( +// |r| r.line(format_args!("Elements: {}", sequence.element.type_name())), +// |r| r.render_body(&Tagged::new(sequence.element)), +// |_| Ok(()), +// ) +// } +// +// //--------// +// // Fields // +// //--------// +// +// fn render_fields(&mut self, fields: &Fields) -> fmt::Result { +// match fields { +// Fields::Named(named) => { +// let mut first = true; +// for field in named.iter() { +// if !first { +// self.blank()?; +// } +// self.render_named_field(field)?; +// first = false; +// } +// } +// Fields::Unnamed(unnamed) => { +// for (i, field) in unnamed.iter().enumerate() { +// if i != 0 { +// self.blank()?; +// } +// self.render_unnamed_field(i, field)?; +// } +// } +// Fields::Unit => {} +// } +// +// Ok(()) +// } +// +// fn render_named_field(&mut self, field: &NamedField) -> fmt::Result { +// let tagged = Tagged::new(field.field); +// let will_render_body = self.will_render(tagged.ty()); +// +// self.line(format_args!( +// "{}: {}", +// Quote(field.name), +// tagged.reflection().type_name() +// ))?; +// +// let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; +// if rendered_doc && will_render_body { +// self.blank()?; +// } +// +// if will_render_body { +// self.next(|r| r.render_body(&tagged))?; +// } +// Ok(()) +// } +// +// fn render_unnamed_field(&mut self, index: usize, field: &UnnamedField) -> fmt::Result { +// let tagged = Tagged::new(field.field); +// let will_render_body = self.will_render(tagged.ty()); +// +// self.line(format_args!( +// "{}: {}", +// index, +// tagged.reflection().type_name() +// ))?; +// +// let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; +// if rendered_doc && will_render_body { +// self.blank()?; +// } +// +// if will_render_body { +// self.next(|r| r.render_body(&tagged))?; +// } +// Ok(()) +// } +// +// //---------// +// // Variant // +// //---------// +// +// fn render_variant(&mut self, variant: &Variant) -> fmt::Result { +// self.line(Quote(variant.name))?; +// +// self.indent(|r| { +// let rendered_doc = r.render_doc(variant.doc())?; +// +// if rendered_doc && variant.fields.has_body() { +// r.blank()?; +// } +// +// r.render_fields(&variant.fields) +// }) +// } +// } + +/////////////// +// Bootstrap // +/////////////// + +impl Reflect for std::marker::PhantomData +where + T: Reflect, +{ + fn ty() -> Type { + Type::primitive(None) + } + + fn format_type_name(f: &mut dyn Write) -> fmt::Result { + write!(f, "PhantomData<{}>", Reflection::new::().type_name()) + } +} + +macro_rules! primitive { + ($T:ty, $doc:literal, $type_name:literal) => { + impl Reflect for $T { + fn ty() -> Type { + Type::primitive(Some($doc.into())) + } + + fn format_type_name(f: &mut dyn Write) -> fmt::Result { + f.write_str($type_name) + } + } + }; +} + +primitive!(usize, "A system dependent unsigned integer", "usize"); +primitive!(u32, "A 32-bit unsigned integer", "u32"); + +primitive!(String, "A string", "string"); + +impl Reflect for Vec +where + T: Reflect, +{ + fn ty() -> Type { + Type::sequence::(Some("An ordered collection of elements".into())) + } + + fn format_type_name(f: &mut dyn Write) -> fmt::Result { + write!(f, "Vec<{}>", Reflection::new::().type_name()) + } +} + +#[derive(Reflect)] +pub struct UnitWithConst {} + +#[derive(Reflect)] +pub struct GenericBoundAdded { + uses_t: Vec, +} + +/// This is a test! +/// +/// Hello world! +#[derive(Reflect)] +pub struct Test { + /// This field affects this value. + a: usize, + + /// This field does something else. + b: usize, +} + +#[derive(Reflect)] +pub struct TestUnnamed( + /// Can I document this? + usize, + T, +); + +/// This is a nother test! +#[derive(Reflect)] +pub struct Test2 { + /// This field affects this value. + a: usize, + + /// This field doesn't have any names. + unnamed: TestUnnamed, + + other: Test, + + /// How are we going to compute distances? + metric: AdjacentEnum, + + /// These control a bunch of parameters. + seq: Vec, +} + +#[derive(Reflect)] +struct Wrapper { + /// Inner + a: T, +} + +/// An enum with no payloads. +#[derive(Debug, Clone, Copy, Reflect)] +pub enum Metric { + SquaredL2, + InnerProduct, + Cosine, +} + +/// An enum with no payloads. +#[derive(Debug, Reflect)] +#[serde(rename_all = "kebab-case")] +pub enum AdjacentEnum { + SquaredL2, + /// Let me see if this works + InnerProduct(u32), + + /// Compute the cosine similarity + Cosine { + /// Thos actually doesn't do anything. + test: String, + }, +} + +// impl Reflect for AdjacentEnum { +// fn ty() -> Type { +// Type::enum_( +// EnumRepr::Adjacent { +// tag: "enum-type", +// content: "content", +// }, +// [ +// Variant::new("squared-l2", Fields::Unit, None), +// Variant::new( +// "inner-product", +// Fields::unnamed([UnnamedField::new::(Some("testing".into()))]), +// Some("Inner Product with some payload".into()), +// ), +// Variant::new( +// "cosine", +// Fields::named([NamedField::new::("test", None)]), +// Some("Cosine Similarity".into()), +// ), +// ], +// Some("The similarity measure to use".into()), +// ) +// } +// +// fn format_type_name(f: &mut dyn Write) -> fmt::Result { +// f.write_str("Metric") +// } +// } + +////////////// +// Internal // +////////////// + +pub(crate) mod internal { + use std::marker::PhantomData; + + pub(crate) trait Reflect { + fn ty(&self) -> super::Type; + fn format_type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result; + fn type_id(&self) -> std::any::TypeId; + } + + pub(crate) struct Wrapper(PhantomData); + + impl Wrapper { + pub(crate) const INSTANCE: Self = Self::new(); + + pub(crate) const fn new() -> Self { + Self(PhantomData) + } + } + + impl Reflect for Wrapper + where + T: super::Reflect, + { + fn ty(&self) -> super::Type { + ::ty() + } + + fn format_type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result { + ::format_type_name(f) + } + + fn type_id(&self) -> std::any::TypeId { + std::any::TypeId::of::() + } + } +} + +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + use std::assert_matches; + + #[test] + fn test_unit() { + /// A unit struct. + #[derive(Reflect)] + struct Unit; + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(r.ty().doc().unwrap(), "A unit struct."); + assert_eq!(r.type_name().to_string(), "Unit"); + assert!(!ty.has_body()); + } + + #[test] + fn test_unit_const_generic() { + /// A unit struct. + #[derive(Reflect)] + struct Unit; + + let r = Reflection::new::>(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "A unit struct."); + assert_eq!(r.type_name().to_string(), "Unit<10>"); + assert!(!ty.has_body()); + } + + #[test] + fn test_unit_const_generic_2() { + #[derive(Reflect)] + struct Unit; + + let r = Reflection::new::>(); + let ty = r.ty(); + + assert!(ty.doc().is_none()); + assert_eq!(r.type_name().to_string(), "Unit<10, 20>"); + assert!(!ty.has_body()); + } + + #[test] + fn test_empty_tuple_like() { + /// An empty tuple-like struct. + #[derive(Reflect)] + struct Empty(); + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "An empty tuple-like struct."); + assert_eq!(r.type_name().to_string(), "Empty"); + assert!(!ty.has_body()); + } + + #[test] + fn test_empty_struct_like() { + /// An empty struct. + #[derive(Reflect)] + struct Empty {} + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "An empty struct."); + assert_eq!(r.type_name().to_string(), "Empty"); + assert!(!ty.has_body()); + } + + #[test] + fn test_struct() { + /// A struct with two fields. + #[expect(unused)] + #[derive(Reflect)] + struct Woo { + /// Foo + foo: usize, + /// Bar + bar: usize, + } + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "A struct with two fields."); + assert_eq!(r.type_name().to_string(), "Woo"); + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_named().unwrap(); + assert_eq!(f.len(), 2); + + assert_eq!(f[0].name, "foo"); + assert_eq!(f[0].doc().unwrap(), "Foo"); + assert_eq!(f[0].field.type_name().to_string(), "usize"); + + assert_eq!(f[1].name, "bar"); + assert_eq!(f[1].doc().unwrap(), "Bar"); + assert_eq!(f[1].field.type_name().to_string(), "usize"); + } + + #[test] + fn test_tuple1() { + /// A tuple with one field. + #[expect(unused)] + #[derive(Reflect)] + struct Tuple1( + /// Field 0. + usize, + ); + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "A tuple with one field."); + assert_eq!(r.type_name().to_string(), "Tuple1"); + + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); + assert_eq!(f.len(), 1); + assert_eq!(f[0].doc().unwrap(), "Field 0."); + assert_eq!(f[0].field().type_name().to_string(), "usize"); + } + + #[test] + fn test_tuple2() { + #[expect(unused)] + #[derive(Reflect)] + struct Tuple2( + T, + /// Field 1. + Vec, + ); + + let r = Reflection::new::>(); + let ty = r.ty(); + assert!(ty.doc().is_none()); + assert_eq!(r.type_name().to_string(), "Tuple2"); + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); + assert_eq!(f.len(), 2); + + assert!(f[0].doc().is_none()); + assert_eq!(f[0].field().type_name().to_string(), "usize"); + + assert_eq!(f[1].doc().unwrap(), "Field 1."); + assert_eq!(f[1].field().type_name().to_string(), "Vec"); + } + + #[test] + fn test_empty_enum() { + /// An empty enum. + #[derive(Reflect)] + enum Empty {} + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "An empty enum."); + assert_eq!(r.type_name().to_string(), "Empty"); + assert!(!ty.has_body()); + } + + #[test] + fn test_enum_with_generics() { + /// An enum with generics. + #[expect(unused)] + #[derive(Reflect)] + enum Either { + A(Vec), + /// It's a bee! + B( + /// Buzz buzz + B, + ), + } + + let r = Reflection::new::>(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "An enum with generics."); + assert_eq!(r.type_name().to_string(), "Either"); + assert!(ty.has_body()); + + let variants = ty.as_enum().unwrap().variants(); + assert_eq!(variants.len(), 2); + + // Variant 0 + assert!(variants[0].doc().is_none()); + assert_eq!(variants[0].name(), "A"); + assert!(variants[0].fields.has_body()); + + let f = variants[0].fields.as_unnamed().unwrap(); + assert_eq!(f.len(), 1); + assert!(f[0].doc().is_none()); + assert_eq!(f[0].field().type_name().to_string(), "Vec"); + + // Variant 1 + assert_eq!(variants[1].doc().unwrap(), "It's a bee!"); + assert_eq!(variants[1].name(), "B"); + assert!(variants[1].fields.has_body()); + + let f = variants[1].fields.as_unnamed().unwrap(); + assert_eq!(f.len(), 1); + assert_eq!(f[0].doc().unwrap(), "Buzz buzz"); + assert_eq!(f[0].field().type_name().to_string(), "string"); + } + + #[test] + fn test_enum_variants() { + /// All the enums. + #[expect(unused)] + #[derive(Reflect)] + #[serde(tag = "tag", content = "content")] + enum All { + /// A unit variant. + Unit, + /// A tuple-like variant. + Tuple( + /// Field 0. + usize, + String, + ), + /// Struct-like. + Struct { + foo: usize, + /// All the bars! + bar: u32, + }, + } + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "All the enums."); + assert_eq!(r.type_name().to_string(), "All"); + assert!(ty.has_body()); + + let variants = ty.as_enum().unwrap().variants(); + assert_eq!(variants.len(), 3); + + // Variant 0 + assert_eq!(variants[0].doc().unwrap(), "A unit variant."); + assert_eq!(variants[0].name(), "Unit"); + assert!(!variants[0].fields.has_body()); + assert_matches!(variants[0].fields, Fields::Unit); + + // Variant 1 + assert_eq!(variants[1].doc().unwrap(), "A tuple-like variant."); + assert_eq!(variants[1].name(), "Tuple"); + assert!(variants[1].fields.has_body()); + + let f = variants[1].fields.as_unnamed().unwrap(); + assert_eq!(f.len(), 2); + + assert_eq!(f[0].doc().unwrap(), "Field 0."); + assert_eq!(f[0].field.type_name().to_string(), "usize"); + + assert!(f[1].doc().is_none()); + assert_eq!(f[1].field.type_name().to_string(), "string"); + + // Variant 2 + assert_eq!(variants[2].name(), "Struct"); + assert_eq!(variants[2].doc().unwrap(), "Struct-like."); + let f = variants[2].fields.as_named().unwrap(); + assert_eq!(f.len(), 2); + + assert_eq!(f[0].name, "foo"); + assert!(f[0].doc().is_none()); + assert_eq!(f[0].field.type_name().to_string(), "usize"); + + assert_eq!(f[1].name, "bar"); + assert_eq!(f[1].doc().unwrap(), "All the bars!"); + assert_eq!(f[1].field.type_name().to_string(), "u32"); + } +} diff --git a/diskann-benchmark-runner/src/reflect/render.rs b/diskann-benchmark-runner/src/reflect/render.rs new file mode 100644 index 0000000000..003c2aee22 --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/render.rs @@ -0,0 +1,302 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::fmt::{self, Write}; + +use crate::utils::fmt::Quote; + +use super::{ + Reflection, + tree::{ + Type, Aggregate, Enum, EnumRepr, Sequence, Fields, NamedField, UnnamedField, Variant, + } +}; + +const INDENT: usize = 2; + +#[derive(Debug)] +struct Tagged { + ty: Type, + reflection: Reflection, +} + +impl Tagged { + fn new(reflection: Reflection) -> Self { + Self { + ty: reflection.ty(), + reflection, + } + } + + fn ty(&self) -> &Type { + &self.ty + } + + fn reflection(&self) -> Reflection { + self.reflection + } +} + +pub(super) struct Renderer<'a> { + output: &'a mut dyn Write, + indent: usize, + depth: usize, + max_depth: usize, +} + +impl<'a> Renderer<'a> { + pub(super) fn new(output: &'a mut dyn Write, max_depth: usize) -> Self { + Self { + output, + indent: 0, + depth: 0, + max_depth, + } + } + + fn at_bottom(&self) -> bool { + self.depth >= self.max_depth + } + + fn line(&mut self, display: D) -> fmt::Result + where + D: std::fmt::Display, + { + let indent = INDENT * self.indent; + write!(self.output, "{: >indent$}{}\n", "", display) + } + + fn blank(&mut self) -> fmt::Result { + self.output.write_char('\n') + } + + fn maybe_indent(&mut self, indent: bool, f: F) -> Result + where + F: FnOnce(&mut Self) -> Result, + { + if indent { + self.indent += 1; + } + let result = f(self); + if indent { + self.indent -= 1; + } + result + } + + fn indent(&mut self, f: F) -> Result + where + F: FnOnce(&mut Self) -> Result, + { + self.maybe_indent(true, f) + } + + fn next_with( + &mut self, + pre: impl FnOnce(&mut Self) -> fmt::Result, + body: impl FnOnce(&mut Self) -> fmt::Result, + post: impl FnOnce(&mut Self) -> fmt::Result, + ) -> fmt::Result { + if self.at_bottom() { + Ok(()) + } else { + pre(self)?; + self.depth += 1; + let result = self.indent(body); + self.depth -= 1; + post(self)?; + result + } + } + + fn next(&mut self, f: F) -> fmt::Result + where + F: FnOnce(&mut Self) -> fmt::Result, + { + self.next_with(|_| Ok(()), f, |_| Ok(())) + } + + fn render_doc(&mut self, s: Option<&str>) -> Result { + if let Some(s) = s { + let mut rendered = false; + for ln in s.lines() { + if ln.is_empty() { + self.blank()?; + } else { + self.line(ln)?; + } + + rendered = true; + } + Ok(rendered) + } else { + Ok(false) + } + } + + /// Return `true` if a nested type will be rendered. + fn will_render(&self, ty: &Type) -> bool { + !self.at_bottom() && ty.has_body() + } + + //-------// + // Types // + //-------// + + pub(super) fn render_subject(&mut self, reflection: Reflection) -> fmt::Result { + self.line(reflection.type_name())?; + self.indent(|r| { + let tagged = Tagged::new(reflection); + let wrote_doc = r.render_doc(tagged.ty().doc())?; + + if wrote_doc && r.will_render(tagged.ty()) { + r.blank()?; + } + + r.render_body(&tagged) + }) + } + + fn render_body(&mut self, tagged: &Tagged) -> fmt::Result { + match tagged.ty() { + Type::Primitive(_) => Ok(()), + Type::Aggregate(aggregate) => self.render_aggregate(&aggregate), + Type::Enum(enum_) => self.render_enum(enum_), + Type::Sequence(sequence) => self.render_sequence(&sequence), + } + } + + fn render_aggregate(&mut self, aggregate: &Aggregate) -> fmt::Result { + self.render_fields(aggregate.fields()) + } + + fn render_enum(&mut self, enum_: &Enum) -> fmt::Result { + match enum_.repr() { + EnumRepr::External => self.line("Representation: externally tagged")?, + EnumRepr::Internal { tag } => { + self.line(format_args!("Discriminant field: {}", Quote(tag)))? + } + EnumRepr::Adjacent { tag, content } => { + self.line(format_args!("Discriminant field: {}", Quote(tag)))?; + self.line(format_args!("Content field: {}", Quote(content)))?; + } + } + + self.blank()?; + self.line("Options:")?; + self.indent(|r| { + let mut first = true; + for variant in enum_.variants().iter() { + if !first { + r.blank()?; + } + + r.render_variant(variant)?; + first = false; + } + + Ok(()) + }) + } + + fn render_sequence(&mut self, sequence: &Sequence) -> fmt::Result { + self.next_with( + |r| r.line(format_args!("Elements: {}", sequence.element().type_name())), + |r| r.render_body(&Tagged::new(sequence.element())), + |_| Ok(()), + ) + } + + //--------// + // Fields // + //--------// + + fn render_fields(&mut self, fields: &Fields) -> fmt::Result { + match fields { + Fields::Named(named) => { + let mut first = true; + for field in named.iter() { + if !first { + self.blank()?; + } + self.render_named_field(field)?; + first = false; + } + } + Fields::Unnamed(unnamed) => { + for (i, field) in unnamed.iter().enumerate() { + if i != 0 { + self.blank()?; + } + self.render_unnamed_field(i, field)?; + } + } + Fields::Unit => {} + } + + Ok(()) + } + + fn render_named_field(&mut self, field: &NamedField) -> fmt::Result { + let tagged = Tagged::new(field.field()); + let will_render_body = self.will_render(tagged.ty()); + + self.line(format_args!( + "{}: {}", + Quote(field.name()), + tagged.reflection().type_name() + ))?; + + let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; + if rendered_doc && will_render_body { + self.blank()?; + } + + if will_render_body { + self.next(|r| r.render_body(&tagged))?; + } + Ok(()) + } + + fn render_unnamed_field(&mut self, index: usize, field: &UnnamedField) -> fmt::Result { + let tagged = Tagged::new(field.field()); + let will_render_body = self.will_render(tagged.ty()); + + self.line(format_args!( + "{}: {}", + index, + tagged.reflection().type_name() + ))?; + + let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; + if rendered_doc && will_render_body { + self.blank()?; + } + + if will_render_body { + self.next(|r| r.render_body(&tagged))?; + } + Ok(()) + } + + //---------// + // Variant // + //---------// + + fn render_variant(&mut self, variant: &Variant) -> fmt::Result { + self.line(Quote(variant.name()))?; + + self.indent(|r| { + let rendered_doc = r.render_doc(variant.doc())?; + + if rendered_doc && variant.fields().has_body() { + r.blank()?; + } + + r.render_fields(&variant.fields()) + }) + } +} + diff --git a/diskann-benchmark-runner/src/reflect/tree.rs b/diskann-benchmark-runner/src/reflect/tree.rs new file mode 100644 index 0000000000..dfb3bcfce7 --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/tree.rs @@ -0,0 +1,385 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use super::{Reflect, Reflection}; + +pub type Doc = std::borrow::Cow<'static, str>; + +#[derive(Debug)] +pub enum Type { + Primitive(Primitive), + Aggregate(Aggregate), + Enum(Enum), + Sequence(Sequence), +} + +impl Type { + pub fn primitive(doc: Option) -> Self { + Self::from(Primitive::new(doc)) + } + + pub fn aggregate(fields: Fields, doc: Option) -> Self { + Self::from(Aggregate::new(fields, doc)) + } + + pub fn enum_( + repr: EnumRepr, + variants: impl IntoIterator, + doc: Option, + ) -> Self { + Self::from(Enum::new(repr, variants, doc)) + } + + pub fn sequence(doc: Option) -> Self + where + T: Reflect, + { + Self::from(Sequence::new::(doc)) + } + + pub fn doc(&self) -> Option<&str> { + match self { + Self::Primitive(p) => p.doc(), + Self::Aggregate(a) => a.doc(), + Self::Enum(e) => e.doc(), + Self::Sequence(s) => s.doc(), + } + } + + pub(super) fn has_body(&self) -> bool { + match self { + Self::Primitive(_) => false, + Self::Aggregate(a) => a.has_body(), + Self::Enum(e) => e.has_body(), + Self::Sequence(_) => true, + } + } + + fn as_aggregate(&self) -> Option<&Aggregate> { + if let Self::Aggregate(aggregate) = self { + Some(aggregate) + } else { + None + } + } + + fn as_enum(&self) -> Option<&Enum> { + if let Self::Enum(enum_) = self { + Some(enum_) + } else { + None + } + } +} + +impl From for Type { + fn from(primitive: Primitive) -> Self { + Self::Primitive(primitive) + } +} + +impl From for Type { + fn from(aggergate: Aggregate) -> Self { + Self::Aggregate(aggergate) + } +} + +impl From for Type { + fn from(e: Enum) -> Self { + Self::Enum(e) + } +} + +impl From for Type { + fn from(s: Sequence) -> Self { + Self::Sequence(s) + } +} + +//-----------// +// Primitive // +//-----------// + +#[derive(Debug)] +pub struct Primitive { + doc: Option, +} + +impl Primitive { + pub fn new(doc: Option) -> Self { + Self { doc } + } + + fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } +} + +//-----------// +// Aggregate // +//-----------// + +#[derive(Debug)] +pub struct Aggregate { + fields: Fields, + doc: Option, +} + +impl Aggregate { + pub fn new(fields: Fields, doc: Option) -> Self { + Self { fields, doc } + } + + pub(super) fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } + + pub(super) fn fields(&self) -> &Fields { + &self.fields + } + + pub(super) fn has_body(&self) -> bool { + self.fields.has_body() + } +} + +#[derive(Debug)] +pub enum Fields { + Named(Vec), + Unnamed(Vec), + Unit, +} + +impl Fields { + pub(super) fn has_body(&self) -> bool { + match self { + Self::Named(fields) => !fields.is_empty(), + Self::Unnamed(fields) => !fields.is_empty(), + Self::Unit => false, + } + } + + pub fn named(itr: impl IntoIterator) -> Self { + Self::Named(itr.into_iter().collect()) + } + + pub fn unnamed(itr: impl IntoIterator) -> Self { + Self::Unnamed(itr.into_iter().collect()) + } + + fn as_named(&self) -> Option<&[NamedField]> { + if let Self::Named(fields) = self { + Some(fields) + } else { + None + } + } + + fn as_unnamed(&self) -> Option<&[UnnamedField]> { + if let Self::Unnamed(fields) = self { + Some(fields) + } else { + None + } + } +} + +#[derive(Debug)] +pub struct NamedField { + name: &'static str, + field: Reflection, + doc: Option, +} + +impl NamedField { + pub fn new(name: &'static str, doc: Option) -> Self + where + T: Reflect, + { + Self { + name, + field: Reflection::new::(), + doc, + } + } + + pub(super) fn name(&self) -> &str { + self.name + } + + pub(super) fn field(&self) -> Reflection { + self.field + } + + pub(super) fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } +} + +#[derive(Debug)] +pub struct UnnamedField { + field: Reflection, + doc: Option, +} + +impl UnnamedField { + pub fn new(doc: Option) -> Self + where + T: Reflect, + { + Self { + field: Reflection::new::(), + doc, + } + } + + pub(super) fn field(&self) -> Reflection { + self.field + } + + pub(super) fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } + +} + +//------// +// Enum // +//------// + +#[derive(Debug)] +pub struct Enum { + repr: EnumRepr, + variants: Vec, + doc: Option, +} + +impl Enum { + pub fn new( + repr: EnumRepr, + variants: impl IntoIterator, + doc: Option, + ) -> Self { + Self { + repr, + variants: variants.into_iter().collect(), + doc, + } + } + + pub(super) fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } + + pub(super) fn repr(&self) -> &EnumRepr { + &self.repr + } + + pub(super) fn variants(&self) -> &[Variant] { + &self.variants + } + + pub(super) fn has_body(&self) -> bool { + !self.variants.is_empty() + } +} + +#[derive(Debug)] +pub enum EnumRepr { + /// Enums are tagged as the key in a collection. + External, + + /// Enums are tagged by an internal field. + /// + /// This contains the name of the field that is used as the tag. + /// + /// ```json + /// { + /// "tag": "some-enum-tag", + /// "value": 10, + /// "members": [ + /// 1, + /// "world" + /// ] + /// } + /// ``` + Internal { tag: &'static str }, + + /// Enums have a separate tag and content payload. + /// + /// ```json + /// { + /// "tag": "some-enum-tag", + /// "content": { + /// "value": 10, + /// "members": [ + /// 1, + /// "world" + /// ] + /// } + /// } + /// ``` + Adjacent { + tag: &'static str, + content: &'static str, + }, +} + +#[derive(Debug)] +pub struct Variant { + name: &'static str, + fields: Fields, + doc: Option, +} + +impl Variant { + pub fn new(name: &'static str, fields: Fields, doc: Option) -> Self { + Self { name, fields, doc } + } + + pub(super) fn name(&self) -> &'static str { + self.name + } + + pub(super) fn fields(&self) -> &Fields { + &self.fields + } + + pub(super) fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } +} + +//----------// +// Sequence // +//----------// + +#[derive(Debug)] +pub struct Sequence { + element: Reflection, + doc: Option, +} + +impl Sequence { + pub fn new(doc: Option) -> Self + where + T: Reflect, + { + Self { + element: Reflection::new::(), + doc, + } + } + + pub(super) fn element(&self) -> Reflection { + self.element + } + + pub(super) fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } +} + + From 84097c33235a623c30ac2041129711829ca7f6ef Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Sat, 26 Sep 2026 09:10:22 -0700 Subject: [PATCH 09/21] Clean-up. --- diskann-benchmark-runner/src/reflect/mod.rs | 642 -------------------- 1 file changed, 642 deletions(-) diff --git a/diskann-benchmark-runner/src/reflect/mod.rs b/diskann-benchmark-runner/src/reflect/mod.rs index fde153736f..075948d5b4 100644 --- a/diskann-benchmark-runner/src/reflect/mod.rs +++ b/diskann-benchmark-runner/src/reflect/mod.rs @@ -88,648 +88,6 @@ impl std::fmt::Display for Render { } } -// #[derive(Debug)] -// pub enum Type { -// Primitive(Primitive), -// Aggregate(Aggregate), -// Enum(Enum), -// Sequence(Sequence), -// } -// -// impl Type { -// pub fn primitive(doc: Option) -> Self { -// Self::from(Primitive::new(doc)) -// } -// -// pub fn aggregate(fields: Fields, doc: Option) -> Self { -// Self::from(Aggregate::new(fields, doc)) -// } -// -// pub fn enum_( -// repr: EnumRepr, -// variants: impl IntoIterator, -// doc: Option, -// ) -> Self { -// Self::from(Enum::new(repr, variants, doc)) -// } -// -// pub fn sequence(doc: Option) -> Self -// where -// T: Reflect, -// { -// Self::from(Sequence::new::(doc)) -// } -// -// pub fn doc(&self) -> Option<&str> { -// match self { -// Self::Primitive(p) => p.doc(), -// Self::Aggregate(a) => a.doc(), -// Self::Enum(e) => e.doc(), -// Self::Sequence(s) => s.doc(), -// } -// } -// -// fn has_body(&self) -> bool { -// match self { -// Self::Primitive(_) => false, -// Self::Aggregate(a) => a.has_body(), -// Self::Enum(e) => e.has_body(), -// Self::Sequence(_) => true, -// } -// } -// -// fn as_aggregate(&self) -> Option<&Aggregate> { -// if let Self::Aggregate(aggregate) = self { -// Some(aggregate) -// } else { -// None -// } -// } -// -// fn as_enum(&self) -> Option<&Enum> { -// if let Self::Enum(enum_) = self { -// Some(enum_) -// } else { -// None -// } -// } -// } -// -// impl From for Type { -// fn from(primitive: Primitive) -> Self { -// Self::Primitive(primitive) -// } -// } -// -// impl From for Type { -// fn from(aggergate: Aggregate) -> Self { -// Self::Aggregate(aggergate) -// } -// } -// -// impl From for Type { -// fn from(e: Enum) -> Self { -// Self::Enum(e) -// } -// } -// -// impl From for Type { -// fn from(s: Sequence) -> Self { -// Self::Sequence(s) -// } -// } -// -// //-----------// -// // Primitive // -// //-----------// -// -// #[derive(Debug)] -// pub struct Primitive { -// doc: Option, -// } -// -// impl Primitive { -// pub fn new(doc: Option) -> Self { -// Self { doc } -// } -// -// fn doc(&self) -> Option<&str> { -// self.doc.as_deref() -// } -// } -// -// //-----------// -// // Aggregate // -// //-----------// -// -// #[derive(Debug)] -// pub struct Aggregate { -// fields: Fields, -// doc: Option, -// } -// -// impl Aggregate { -// pub fn new(fields: Fields, doc: Option) -> Self { -// Self { fields, doc } -// } -// -// fn doc(&self) -> Option<&str> { -// self.doc.as_deref() -// } -// -// fn fields(&self) -> &Fields { -// &self.fields -// } -// -// fn has_body(&self) -> bool { -// self.fields.has_body() -// } -// } -// -// #[derive(Debug)] -// pub enum Fields { -// Named(Vec), -// Unnamed(Vec), -// Unit, -// } -// -// impl Fields { -// fn has_body(&self) -> bool { -// match self { -// Self::Named(fields) => !fields.is_empty(), -// Self::Unnamed(fields) => !fields.is_empty(), -// Self::Unit => false, -// } -// } -// -// pub fn named(itr: impl IntoIterator) -> Self { -// Self::Named(itr.into_iter().collect()) -// } -// -// pub fn unnamed(itr: impl IntoIterator) -> Self { -// Self::Unnamed(itr.into_iter().collect()) -// } -// -// fn as_named(&self) -> Option<&[NamedField]> { -// if let Self::Named(fields) = self { -// Some(fields) -// } else { -// None -// } -// } -// -// fn as_unnamed(&self) -> Option<&[UnnamedField]> { -// if let Self::Unnamed(fields) = self { -// Some(fields) -// } else { -// None -// } -// } -// } -// -// #[derive(Debug)] -// pub struct NamedField { -// name: &'static str, -// field: Reflection, -// doc: Option, -// } -// -// impl NamedField { -// pub fn new(name: &'static str, doc: Option) -> Self -// where -// T: Reflect, -// { -// Self { -// name, -// field: Reflection::new::(), -// doc, -// } -// } -// -// fn doc(&self) -> Option<&str> { -// self.doc.as_deref() -// } -// } -// -// #[derive(Debug)] -// pub struct UnnamedField { -// field: Reflection, -// doc: Option, -// } -// -// impl UnnamedField { -// pub fn new(doc: Option) -> Self -// where -// T: Reflect, -// { -// Self { -// field: Reflection::new::(), -// doc, -// } -// } -// -// fn doc(&self) -> Option<&str> { -// self.doc.as_deref() -// } -// -// fn field(&self) -> Reflection { -// self.field -// } -// } -// -// //------// -// // Enum // -// //------// -// -// #[derive(Debug)] -// pub struct Enum { -// repr: EnumRepr, -// variants: Vec, -// doc: Option, -// } -// -// impl Enum { -// pub fn new( -// repr: EnumRepr, -// variants: impl IntoIterator, -// doc: Option, -// ) -> Self { -// Self { -// repr, -// variants: variants.into_iter().collect(), -// doc, -// } -// } -// -// pub fn doc(&self) -> Option<&str> { -// self.doc.as_deref() -// } -// -// fn variants(&self) -> &[Variant] { -// &self.variants -// } -// -// fn has_body(&self) -> bool { -// !self.variants.is_empty() -// } -// } -// -// #[derive(Debug)] -// pub enum EnumRepr { -// /// Enums are tagged as the key in a collection. -// External, -// -// /// Enums are tagged by an internal field. -// /// -// /// This contains the name of the field that is used as the tag. -// /// -// /// ```json -// /// { -// /// "tag": "some-enum-tag", -// /// "value": 10, -// /// "members": [ -// /// 1, -// /// "world" -// /// ] -// /// } -// /// ``` -// Internal { tag: &'static str }, -// -// /// Enums have a separate tag and content payload. -// /// -// /// ```json -// /// { -// /// "tag": "some-enum-tag", -// /// "content": { -// /// "value": 10, -// /// "members": [ -// /// 1, -// /// "world" -// /// ] -// /// } -// /// } -// /// ``` -// Adjacent { -// tag: &'static str, -// content: &'static str, -// }, -// } -// -// #[derive(Debug)] -// pub struct Variant { -// name: &'static str, -// fields: Fields, -// doc: Option, -// } -// -// impl Variant { -// pub fn new(name: &'static str, fields: Fields, doc: Option) -> Self { -// Self { name, fields, doc } -// } -// -// fn name(&self) -> &'static str { -// self.name -// } -// -// fn doc(&self) -> Option<&str> { -// self.doc.as_deref() -// } -// } -// -// //----------// -// // Sequence // -// //----------// -// -// #[derive(Debug)] -// pub struct Sequence { -// element: Reflection, -// doc: Option, -// } -// -// impl Sequence { -// pub fn new(doc: Option) -> Self -// where -// T: Reflect, -// { -// Self { -// element: Reflection::new::(), -// doc, -// } -// } -// -// pub fn doc(&self) -> Option<&str> { -// self.doc.as_deref() -// } -// } -// -// ////////////// -// // Renderer // -// ////////////// -// -// #[derive(Debug)] -// struct Tagged { -// ty: Type, -// reflection: Reflection, -// } -// -// impl Tagged { -// fn new(reflection: Reflection) -> Self { -// Self { -// ty: reflection.ty(), -// reflection, -// } -// } -// -// fn ty(&self) -> &Type { -// &self.ty -// } -// -// fn reflection(&self) -> Reflection { -// self.reflection -// } -// } -// -// struct Renderer<'a> { -// output: &'a mut dyn Write, -// indent: usize, -// depth: usize, -// max_depth: usize, -// } -// -// impl<'a> Renderer<'a> { -// fn new(output: &'a mut dyn Write, max_depth: usize) -> Self { -// Self { -// output, -// indent: 0, -// depth: 0, -// max_depth, -// } -// } -// -// fn at_bottom(&self) -> bool { -// self.depth >= self.max_depth -// } -// -// fn line(&mut self, display: D) -> fmt::Result -// where -// D: std::fmt::Display, -// { -// let indent = INDENT * self.indent; -// write!(self.output, "{: >indent$}{}\n", "", display) -// } -// -// fn blank(&mut self) -> fmt::Result { -// self.output.write_char('\n') -// } -// -// fn maybe_indent(&mut self, indent: bool, f: F) -> Result -// where -// F: FnOnce(&mut Self) -> Result, -// { -// if indent { -// self.indent += 1; -// } -// let result = f(self); -// if indent { -// self.indent -= 1; -// } -// result -// } -// -// fn indent(&mut self, f: F) -> Result -// where -// F: FnOnce(&mut Self) -> Result, -// { -// self.maybe_indent(true, f) -// } -// -// fn next_with( -// &mut self, -// pre: impl FnOnce(&mut Self) -> fmt::Result, -// body: impl FnOnce(&mut Self) -> fmt::Result, -// post: impl FnOnce(&mut Self) -> fmt::Result, -// ) -> fmt::Result { -// if self.at_bottom() { -// Ok(()) -// } else { -// pre(self)?; -// self.depth += 1; -// let result = self.indent(body); -// self.depth -= 1; -// post(self)?; -// result -// } -// } -// -// fn next(&mut self, f: F) -> fmt::Result -// where -// F: FnOnce(&mut Self) -> fmt::Result, -// { -// self.next_with(|_| Ok(()), f, |_| Ok(())) -// } -// -// fn render_doc(&mut self, s: Option<&str>) -> Result { -// if let Some(s) = s { -// let mut rendered = false; -// for ln in s.lines() { -// if ln.is_empty() { -// self.blank()?; -// } else { -// self.line(ln)?; -// } -// -// rendered = true; -// } -// Ok(rendered) -// } else { -// Ok(false) -// } -// } -// -// /// Return `true` if a nested type will be rendered. -// fn will_render(&self, ty: &Type) -> bool { -// !self.at_bottom() && ty.has_body() -// } -// -// //-------// -// // Types // -// //-------// -// -// fn render_subject(&mut self, reflection: Reflection) -> fmt::Result { -// self.line(reflection.type_name())?; -// self.indent(|r| { -// let tagged = Tagged::new(reflection); -// let wrote_doc = r.render_doc(tagged.ty().doc())?; -// -// if wrote_doc && r.will_render(tagged.ty()) { -// r.blank()?; -// } -// -// r.render_body(&tagged) -// }) -// } -// -// fn render_body(&mut self, tagged: &Tagged) -> fmt::Result { -// match tagged.ty() { -// Type::Primitive(_) => Ok(()), -// Type::Aggregate(aggregate) => self.render_aggregate(&aggregate), -// Type::Enum(enum_) => self.render_enum(enum_), -// Type::Sequence(sequence) => self.render_sequence(&sequence), -// } -// } -// -// fn render_aggregate(&mut self, aggregate: &Aggregate) -> fmt::Result { -// self.render_fields(&aggregate.fields) -// } -// -// fn render_enum(&mut self, enum_: &Enum) -> fmt::Result { -// match enum_.repr { -// EnumRepr::External => self.line("Representation: externally tagged")?, -// EnumRepr::Internal { tag } => { -// self.line(format_args!("Discriminant field: {}", Quote(tag)))? -// } -// EnumRepr::Adjacent { tag, content } => { -// self.line(format_args!("Discriminant field: {}", Quote(tag)))?; -// self.line(format_args!("Content field: {}", Quote(content)))?; -// } -// } -// -// self.blank()?; -// self.line("Options:")?; -// self.indent(|r| { -// let mut first = true; -// for variant in enum_.variants.iter() { -// if !first { -// r.blank()?; -// } -// -// r.render_variant(variant)?; -// first = false; -// } -// -// Ok(()) -// }) -// } -// -// fn render_sequence(&mut self, sequence: &Sequence) -> fmt::Result { -// self.next_with( -// |r| r.line(format_args!("Elements: {}", sequence.element.type_name())), -// |r| r.render_body(&Tagged::new(sequence.element)), -// |_| Ok(()), -// ) -// } -// -// //--------// -// // Fields // -// //--------// -// -// fn render_fields(&mut self, fields: &Fields) -> fmt::Result { -// match fields { -// Fields::Named(named) => { -// let mut first = true; -// for field in named.iter() { -// if !first { -// self.blank()?; -// } -// self.render_named_field(field)?; -// first = false; -// } -// } -// Fields::Unnamed(unnamed) => { -// for (i, field) in unnamed.iter().enumerate() { -// if i != 0 { -// self.blank()?; -// } -// self.render_unnamed_field(i, field)?; -// } -// } -// Fields::Unit => {} -// } -// -// Ok(()) -// } -// -// fn render_named_field(&mut self, field: &NamedField) -> fmt::Result { -// let tagged = Tagged::new(field.field); -// let will_render_body = self.will_render(tagged.ty()); -// -// self.line(format_args!( -// "{}: {}", -// Quote(field.name), -// tagged.reflection().type_name() -// ))?; -// -// let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; -// if rendered_doc && will_render_body { -// self.blank()?; -// } -// -// if will_render_body { -// self.next(|r| r.render_body(&tagged))?; -// } -// Ok(()) -// } -// -// fn render_unnamed_field(&mut self, index: usize, field: &UnnamedField) -> fmt::Result { -// let tagged = Tagged::new(field.field); -// let will_render_body = self.will_render(tagged.ty()); -// -// self.line(format_args!( -// "{}: {}", -// index, -// tagged.reflection().type_name() -// ))?; -// -// let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; -// if rendered_doc && will_render_body { -// self.blank()?; -// } -// -// if will_render_body { -// self.next(|r| r.render_body(&tagged))?; -// } -// Ok(()) -// } -// -// //---------// -// // Variant // -// //---------// -// -// fn render_variant(&mut self, variant: &Variant) -> fmt::Result { -// self.line(Quote(variant.name))?; -// -// self.indent(|r| { -// let rendered_doc = r.render_doc(variant.doc())?; -// -// if rendered_doc && variant.fields.has_body() { -// r.blank()?; -// } -// -// r.render_fields(&variant.fields) -// }) -// } -// } - /////////////// // Bootstrap // /////////////// From fae0fe9870cd4d8f55c675e604df69fb8b94bd14 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Sat, 26 Sep 2026 12:57:39 -0700 Subject: [PATCH 10/21] Slowly expanding. --- Cargo.lock | 2 +- Cargo.toml | 1 - diskann-benchmark-runner-derive/Cargo.toml | 1 - diskann-benchmark-runner-derive/SERDE_TODO.md | 2 +- .../src/{serde.rs => attributes.rs} | 248 +++++++-- diskann-benchmark-runner-derive/src/lib.rs | 92 +++- diskann-benchmark-runner/Cargo.toml | 1 + diskann-benchmark-runner/src/app.rs | 4 + diskann-benchmark-runner/src/input.rs | 18 +- diskann-benchmark-runner/src/reflect/mod.rs | 487 ++++-------------- .../src/reflect/render.rs | 18 +- diskann-benchmark-runner/src/reflect/test.rs | 373 ++++++++++++++ diskann-benchmark-runner/src/reflect/tree.rs | 15 +- diskann-benchmark-runner/src/registry.rs | 55 +- diskann-benchmark-runner/src/test/dim.rs | 6 +- diskann-benchmark-runner/src/test/gated.rs | 5 +- diskann-benchmark-runner/src/test/typed.rs | 10 +- .../src/utils/datatype.rs | 4 +- diskann-benchmark-runner/temp/main.rs | 6 +- 19 files changed, 877 insertions(+), 471 deletions(-) rename diskann-benchmark-runner-derive/src/{serde.rs => attributes.rs} (54%) create mode 100644 diskann-benchmark-runner/src/reflect/test.rs diff --git a/Cargo.lock b/Cargo.lock index 0dd8a1e6fc..39f0338718 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -519,6 +519,7 @@ dependencies = [ "clap", "diskann-benchmark-runner-derive", "half", + "hashbrown 0.16.1", "indicatif", "serde", "serde_json", @@ -530,7 +531,6 @@ dependencies = [ name = "diskann-benchmark-runner-derive" version = "0.59.0" dependencies = [ - "heck", "proc-macro2", "quote", "syn 2.0.117", diff --git a/Cargo.toml b/Cargo.toml index b884bce878..5e8562cd6c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -71,7 +71,6 @@ diskann-benchmark-runner = { path = "diskann-benchmark-runner", version = "0.59. diskann-benchmark-core = { path = "diskann-benchmark-core", version = "0.59.0" } diskann-tools = { path = "diskann-tools", version = "0.59.0" } diskann-bftree = { path = "diskann-bftree", version = "0.59.0" } -diskann-record = { path = "diskann-record", version = "0.59.0" } # External dependencies (shared versions) anyhow = "1.0.98" diff --git a/diskann-benchmark-runner-derive/Cargo.toml b/diskann-benchmark-runner-derive/Cargo.toml index 4bedaf395c..9672dbebcb 100644 --- a/diskann-benchmark-runner-derive/Cargo.toml +++ b/diskann-benchmark-runner-derive/Cargo.toml @@ -13,7 +13,6 @@ proc-macro = true syn = { version = "2", features = ["full"] } quote = "1" proc-macro2 = "1" -heck = "0.5.0" [lints] workspace = true diff --git a/diskann-benchmark-runner-derive/SERDE_TODO.md b/diskann-benchmark-runner-derive/SERDE_TODO.md index 58ad1db011..b8c19c43bc 100644 --- a/diskann-benchmark-runner-derive/SERDE_TODO.md +++ b/diskann-benchmark-runner-derive/SERDE_TODO.md @@ -6,7 +6,7 @@ enum representations. ## Correctness -- [ ] Use separate Serde-compatible case conversion for fields and variants. +- [X] Use separate Serde-compatible case conversion for fields and variants. - Serde applies different rules in each context. - In particular, Serde converts an enum variant such as `XMLHttpRequest` to `x_m_l_http_request` under `snake_case`, while `heck` produces diff --git a/diskann-benchmark-runner-derive/src/serde.rs b/diskann-benchmark-runner-derive/src/attributes.rs similarity index 54% rename from diskann-benchmark-runner-derive/src/serde.rs rename to diskann-benchmark-runner-derive/src/attributes.rs index 4a6539d973..b82d7b9ba1 100644 --- a/diskann-benchmark-runner-derive/src/serde.rs +++ b/diskann-benchmark-runner-derive/src/attributes.rs @@ -10,6 +10,11 @@ fn is_serde_attr(attr: &syn::Attribute) -> bool { attr.path().is_ident("serde") } +#[must_use] +fn is_reflect_attr(attr: &syn::Attribute) -> bool { + attr.path().is_ident("reflect") +} + fn identity(x: T) -> T { x } @@ -26,18 +31,33 @@ fn set_unique(opt: &mut Option, value: syn::LitStr, attr: &str) -> } } +/// Attributes applicable to struct definitions. pub(crate) struct Struct { + /// A univeral rename rule for all fields. + /// + /// [`Field`] specific renames take precedence. pub(crate) rename_all: RenameAll, } +/// Attributes applicable to enum definitions. pub(crate) struct Enum { + /// A universal rename rule for all variants. + /// + /// [`Variant`] specific renames take precedence. pub(crate) rename_all: RenameAll, + + /// Enum's representation. pub(crate) enum_repr: EnumRepr, } +/// Attributes on the top-level [`syn::DeriveInput`]. +/// +/// Uses should go through [`Container::as_enum`] or [`Container::try_as_struct`] to ensure +/// the attributes are appropriate for the actual type. pub(crate) struct Container { - pub(crate) rename_all: RenameAll, - pub(crate) enum_repr: EnumRepr, + rename_all: RenameAll, + enum_repr: EnumRepr, + type_name: TypeName, } impl Container { @@ -45,13 +65,15 @@ impl Container { let mut rename_all = RenameAll::None; let mut tag = Option::None; let mut content = Option::None; + let mut type_name = TypeName::None; + // Parse serde attributes. for attr in attrs.iter().filter(|a| is_serde_attr(*a)) { attr.parse_nested_meta(|meta| { // serde(rename_all = "...") if meta.path.is_ident("rename_all") { let value: syn::LitStr = meta.value()?.parse()?; - rename_all.parse_in(value)?; + rename_all.parse_once(value)?; return Ok(()); } @@ -73,12 +95,35 @@ impl Container { })?; } + // Parse reflect attributes + for attr in attrs.iter().filter(|a| is_reflect_attr(*a)) { + attr.parse_nested_meta(|meta| { + // reflect(prefix = "...") + if meta.path.is_ident("prefix") { + let value: syn::LitStr = meta.value()?.parse()?; + type_name.set_unique(TypeNameKind::Prefix, value)?; + return Ok(()); + } + + // reflect(type_name = "...") + if meta.path.is_ident("type_name") { + let value: syn::LitStr = meta.value()?.parse()?; + type_name.set_unique(TypeNameKind::Rename, value)?; + return Ok(()); + } + + Err(meta.error("unsupported attribute for Reflect")) + })?; + } + Ok(Self { rename_all, enum_repr: EnumRepr::from_parsed(tag, content)?, + type_name, }) } + /// Verify the parsed attributes are compatible with an aggregate definition. pub(crate) fn try_as_struct(self) -> syn::Result { self.enum_repr.assert_struct_compatible()?; Ok(Struct { @@ -86,14 +131,21 @@ impl Container { }) } + /// Verify the parsed attributes are compatible with an enum definition. pub(crate) fn as_enum(self) -> Enum { Enum { rename_all: self.rename_all, enum_repr: self.enum_repr, } } + + /// Extract the type-name attributes. + pub(crate) fn type_name(&self) -> TypeName { + self.type_name.clone() + } } +/// The `serde` representation for an enum. pub(crate) enum EnumRepr { External, Internal { @@ -106,6 +158,9 @@ pub(crate) enum EnumRepr { } impl EnumRepr { + /// Verify that the parsed `tag` and `content` fields are coherent. + /// + /// This checks that `content` cannot be applied without a `tag`. pub(crate) fn from_parsed( tag: Option, content: Option, @@ -121,6 +176,10 @@ impl EnumRepr { } } + /// Verify that no `enum` tag attributes are present. + /// + /// These do not apply to structs, so we give a compile error with a diagnostic if they + /// are observed. pub(crate) fn assert_struct_compatible(&self) -> syn::Result<()> { match self { Self::External => Ok(()), @@ -136,6 +195,7 @@ impl EnumRepr { } } +/// Supported subset of `serde(rename_all = "...")` #[derive(Default, Debug, Clone, Copy, PartialEq)] pub(crate) enum RenameAll { #[default] @@ -159,7 +219,8 @@ impl RenameAll { } } - fn parse_in(&mut self, s: syn::LitStr) -> syn::Result { + /// Attempt to parse `s`, returning an error if `self` is already parsed. + fn parse_once(&mut self, s: syn::LitStr) -> syn::Result<()> { if *self != Self::None { Err(syn::Error::new_spanned( s, @@ -168,7 +229,10 @@ impl RenameAll { } else { let value = s.value(); match Self::parse(&value) { - Some(me) => Ok(me), + Some(me) => { + *self = me; + Ok(()) + } None => Err(syn::Error::new_spanned( s, format!( @@ -204,6 +268,7 @@ impl RenameAll { } } + /// Apply the rename rule to `variant`. pub(crate) fn apply_to_variant(&self, variant: syn::LitStr) -> syn::LitStr { if *self == Self::None { variant @@ -224,19 +289,76 @@ impl RenameAll { } } - pub(crate) fn apply_to_field(&self, variant: syn::LitStr) -> syn::LitStr { + /// Apply the rename rule to `field`. + pub(crate) fn apply_to_field(&self, field: syn::LitStr) -> syn::LitStr { if *self == Self::None { - variant + field } else { - syn::LitStr::new(&self.apply_to_field_str(&variant.value()), variant.span()) + syn::LitStr::new(&self.apply_to_field_str(&field.value()), field.span()) } } } +/// A one-type variant or field renamer. #[derive(Default)] -pub(crate) struct Variant { +pub(crate) struct RenameOnce { rename: Option, - rename_all: RenameAll, +} + +impl RenameOnce { + /// Apply the configured rename to `variant`. If no rename is configured, apply `or_else`. + pub(crate) fn apply_to_variant(self, variant: syn::LitStr, or_else: RenameAll) -> syn::LitStr { + self.rename + .map_or_else(|| or_else.apply_to_variant(variant), identity) + } + + /// Apply the configured rename to `field`. If no rename is configured, apply `or_else`. + pub(crate) fn apply_to_field(self, field: syn::LitStr, or_else: RenameAll) -> syn::LitStr { + self.rename + .map_or_else(|| or_else.apply_to_field(field), identity) + } +} + +/// Strategy for generting type-names. +#[derive(Default, Clone)] +pub(crate) enum TypeName { + #[default] + None, + Prefix(syn::LitStr), + Rename(syn::LitStr), +} + +enum TypeNameKind { + Prefix, + Rename, +} + +impl TypeName { + fn set_unique(&mut self, kind: TypeNameKind, value: syn::LitStr) -> syn::Result<()> { + if !matches!(self, Self::None) { + Err(syn::Error::new_spanned( + value, + "reflect attribute `prefix` found multiple times", + )) + } else { + match kind { + TypeNameKind::Prefix => *self = Self::Prefix(value), + TypeNameKind::Rename => *self = Self::Rename(value), + } + Ok(()) + } + } +} + +//---------// +// Variant // +//---------// + +/// Variant level attributes. +#[derive(Default)] +pub(crate) struct Variant { + pub(crate) rename_variant: RenameOnce, + pub(crate) rename_variant_fields: RenameAll, } impl Variant { @@ -248,14 +370,14 @@ impl Variant { // serde(rename_all = "...") if meta.path.is_ident("rename_all") { let value: syn::LitStr = meta.value()?.parse()?; - me.rename_all.parse_in(value)?; + me.rename_variant_fields.parse_once(value)?; return Ok(()); } // serde(rename = "...") if meta.path.is_ident("rename") { let value: syn::LitStr = meta.value()?.parse()?; - set_unique(&mut me.rename, value, "rename")?; + set_unique(&mut me.rename_variant.rename, value, "rename")?; return Ok(()); } @@ -265,23 +387,16 @@ impl Variant { Ok(me) } - - /// Replace the name of this variant if directed by `serde(rename = "...")`. - /// - /// If the rename attribute exists, it takes precedence. Otherwise, the fallback is used. - pub(crate) fn rename_variant_or(self, variant: syn::LitStr, or_else: RenameAll) -> syn::LitStr { - self.rename - .map_or_else(|| or_else.apply_to_variant(variant), identity) - } - - pub(crate) fn field_rename_all(&self) -> RenameAll { - self.rename_all - } } -#[derive(Default, Clone)] +//-------// +// Field // +//-------// + +/// Field level attributes. +#[derive(Default)] pub(crate) struct Field { - rename: Option, + pub(crate) rename_field: RenameOnce, } impl Field { @@ -293,7 +408,7 @@ impl Field { // serde(rename = "...") if meta.path.is_ident("rename") { let value: syn::LitStr = meta.value()?.parse()?; - set_unique(&mut me.rename, value, "rename")?; + set_unique(&mut me.rename_field.rename, value, "rename")?; return Ok(()); } @@ -303,12 +418,79 @@ impl Field { Ok(me) } +} - /// Apply the renaming rules defined in `self`. - /// - /// If no renaming rules are present, instead invoke `or_else`. - pub(crate) fn rename_field_or(self, field: syn::LitStr, or_else: RenameAll) -> syn::LitStr { - let Self { rename } = self; - rename.map_or_else(|| or_else.apply_to_field(field), identity) +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn test_rename_all_parse() { + assert!(RenameAll::parse("none").is_none()); + assert_eq!(RenameAll::parse("lowercase").unwrap(), RenameAll::Lower); + assert_eq!(RenameAll::parse("snake_case").unwrap(), RenameAll::Snake); + assert_eq!(RenameAll::parse("kebab-case").unwrap(), RenameAll::Kebab); + + assert!(RenameAll::parse("foo").is_none()); + assert!(RenameAll::parse("bar").is_none()); + } + + #[test] + fn test_apply_to_variant() { + assert_eq!( + RenameAll::None.apply_to_variant_str("MiXeDUpper_Case"), + "MiXeDUpper_Case" + ); + + assert_eq!( + RenameAll::Lower.apply_to_variant_str("MiXeDUpper_Case"), + "mixedupper_case" + ); + assert_eq!( + RenameAll::Lower.apply_to_variant_str("all_lower"), + "all_lower" + ); + + assert_eq!( + RenameAll::Snake.apply_to_variant_str("MixedUpperCase"), + "mixed_upper_case" + ); + assert_eq!( + RenameAll::Snake.apply_to_variant_str("X86_64_V4"), + "x86_64__v4" + ); + + assert_eq!( + RenameAll::Kebab.apply_to_variant_str("MixedUpperCase"), + "mixed-upper-case" + ); + assert_eq!( + RenameAll::Kebab.apply_to_variant_str("X86_64_V4"), + "x86-64--v4" + ); + } + + #[test] + fn test_apply_to_field() { + assert_eq!( + RenameAll::None.apply_to_field_str("a_standard_field"), + "a_standard_field" + ); + assert_eq!( + RenameAll::Lower.apply_to_field_str("a_standard_field"), + "a_standard_field" + ); + assert_eq!( + RenameAll::Snake.apply_to_field_str("a_standard_field"), + "a_standard_field" + ); + assert_eq!( + RenameAll::Kebab.apply_to_field_str("a_standard_field"), + "a-standard-field" + ); } } diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index e7906a6cf4..089afa74ac 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -7,7 +7,7 @@ use proc_macro2::TokenStream; use quote::{quote, quote_spanned}; use syn::{Data, DeriveInput, Fields, parse_macro_input, parse_quote, spanned::Spanned}; -mod serde; +mod attributes; fn crate_name() -> syn::Path { syn::parse_quote!(::diskann_benchmark_runner::reflect) @@ -21,7 +21,7 @@ fn crate_name() -> syn::Path { /// # Example /// /// ```ignore -/// use diskann_benchmark_runner::reflect::Reflect; +/// use diskann_benchmark_runner::Reflect; /// /// /// A test aggregate. /// #[derive(Reflect)] @@ -61,9 +61,9 @@ fn expand(input: &DeriveInput) -> syn::Result { let doc = format_docstrings(&input.attrs); let mut generics = input.generics.clone(); add_generic_bounds(&mut generics); + let container = attributes::Container::parse(&input.attrs)?; - let format_type_name = generate_type_name_body(&input); - let container = serde::Container::parse(&input.attrs)?; + let format_type_name = generate_type_name_body(&input, container.type_name())?; let common = DeriveCommon { doc, @@ -88,7 +88,7 @@ struct DeriveCommon { /// The implementation of `format_type_name`. format_type_name: TokenStream, /// Serde container-level attributes. - container: serde::Container, + container: attributes::Container, } /// Add a bound `T: Reflect` for each type parameter in the generic list. @@ -166,7 +166,7 @@ where /// f.write_str("<"); /// // This would come from the `Reflection::type_name` instead. /// type_name_u32(f); -/// f.write_str(">"); +/// f.write_str(">") /// } /// /// fn type_name_u32(f: &mut dyn std::fmt::Write) -> std::fmt::Result { @@ -174,7 +174,19 @@ where /// } /// ``` /// When there are no generics - we can print the type name directly. -fn generate_type_name_body(input: &DeriveInput) -> TokenStream { +/// +/// # Compatibility with type-name attributes +/// +/// There are two type-name attributes supported: +/// +/// * `reflect(prefix = "...")`: Apply the prefix to the final type-name. +/// * `reflect(type_name = "...")`: Use the given type name literal instead. +/// +/// To prevent mayhem, `type_name` may only be used on non-genric types. +fn generate_type_name_body( + input: &DeriveInput, + type_name: attributes::TypeName, +) -> syn::Result { let path = crate_name(); let name = input.ident.to_string(); let arguments: Vec<_> = input @@ -202,11 +214,37 @@ fn generate_type_name_body(input: &DeriveInput) -> TokenStream { }) .collect(); + // Check that the type-name attributes are compatible with the struct. + let prefix: Vec = match type_name { + attributes::TypeName::Rename(rename) => { + if arguments.is_empty() { + let ts = quote! { + f.write_str(#rename) + }; + return Ok(ts); + } else { + return Err(syn::Error::new_spanned( + rename, + "The `type_name` attribute cannot be applied to types with generics", + )); + } + }, + attributes::TypeName::Prefix(prefix) => { + let ts = quote! { + f.write_str(#prefix)?; + }; + vec![ts] + }, + attributes::TypeName::None => Vec::new(), + }; + // If there are no generics, we can dump the typename directly. if arguments.is_empty() { - return quote! { + let ts = quote! { + #(#prefix)* f.write_str(#name) }; + Ok(ts) } else { let writes = arguments.iter().enumerate().map(|(index, argument)| { if index == 0 { @@ -221,19 +259,21 @@ fn generate_type_name_body(input: &DeriveInput) -> TokenStream { } }); - quote! { + let ts = quote! { + #(#prefix)* f.write_str(#name)?; f.write_str("<")?; #(#writes)* f.write_str(">") - } + }; + Ok(ts) } } fn build_fields( fields: &syn::Fields, generics: &mut syn::Generics, - rename_all: serde::RenameAll, + rename_all: attributes::RenameAll, ) -> syn::Result { let path = crate_name(); @@ -252,7 +292,10 @@ fn build_fields( } } -fn named_fields<'a, I>(fields: I, rename_all: serde::RenameAll) -> syn::Result> +fn named_fields<'a, I>( + fields: I, + rename_all: attributes::RenameAll, +) -> syn::Result> where I: IntoIterator, { @@ -269,8 +312,8 @@ where let name = syn::LitStr::new(&ident.to_string(), ident.span()); let doc = format_docstrings(&f.attrs); - let field = serde::Field::parse(&f.attrs)?; - let name = field.rename_field_or(name, rename_all); + let attributes::Field { rename_field } = attributes::Field::parse(&f.attrs)?; + let name = rename_field.apply_to_field(name, rename_all); Ok(quote_spanned! { ty.span()=> #path::tree::NamedField::new::<#ty>(#name, #doc) }) }) .collect() @@ -302,7 +345,7 @@ fn process_struct( } = common; // Validate that the attributes we parsed are compatible with a `struct` definition. - let serde::Struct { rename_all } = container.try_as_struct()?; + let attributes::Struct { rename_all } = container.try_as_struct()?; let type_name = &input.ident; let path = crate_name(); @@ -346,7 +389,7 @@ fn process_enum( } = common; // Validate that the attributes we parsed are compatible with an `enum` definition. - let serde::Enum { + let attributes::Enum { rename_all, enum_repr, } = container.as_enum(); @@ -361,21 +404,26 @@ fn process_enum( .map(|v| -> syn::Result { let doc = format_docstrings(&v.attrs); let name = syn::LitStr::new(&v.ident.to_string(), v.ident.span()); - let attrs = serde::Variant::parse(&v.attrs)?; + let attributes::Variant { + rename_variant, + rename_variant_fields, + } = attributes::Variant::parse(&v.attrs)?; - let fields = build_fields(&v.fields, &mut generics, attrs.field_rename_all())?; + let fields = build_fields(&v.fields, &mut generics, rename_variant_fields)?; // Rename the variant as needed. - let name = attrs.rename_variant_or(name, rename_all); + let name = rename_variant.apply_to_variant(name, rename_all); Ok(quote!(#path::tree::Variant::new(#name, #fields, #doc))) }) .collect::>>()?; // Build the enum representation. let enum_repr = match enum_repr { - serde::EnumRepr::External => quote!(#path::tree::EnumRepr::External), - serde::EnumRepr::Internal { tag } => quote!(#path::tree::EnumRepr::Internal { tag: #tag }), - serde::EnumRepr::Adjacent { tag, content } => { + attributes::EnumRepr::External => quote!(#path::tree::EnumRepr::External), + attributes::EnumRepr::Internal { tag } => { + quote!(#path::tree::EnumRepr::Internal { tag: #tag }) + } + attributes::EnumRepr::Adjacent { tag, content } => { quote!(#path::tree::EnumRepr::Adjacent { tag: #tag, content: #content }) } }; diff --git a/diskann-benchmark-runner/Cargo.toml b/diskann-benchmark-runner/Cargo.toml index 9d2369827b..331cfc6e4b 100644 --- a/diskann-benchmark-runner/Cargo.toml +++ b/diskann-benchmark-runner/Cargo.toml @@ -12,6 +12,7 @@ anyhow = { workspace = true } clap = { workspace = true, features = ["derive"] } diskann-benchmark-runner-derive = { workspace = true } half = { workspace = true } +hashbrown = { workspace = true } indicatif = "0.18.3" serde = { workspace = true, features = ["derive"] } serde_json = { workspace = true } diff --git a/diskann-benchmark-runner/src/app.rs b/diskann-benchmark-runner/src/app.rs index c54baed43a..1714d8959c 100644 --- a/diskann-benchmark-runner/src/app.rs +++ b/diskann-benchmark-runner/src/app.rs @@ -213,6 +213,10 @@ impl App { describe )?; writeln!(output, "{}", serde_json::to_string_pretty(&repr)?)?; + + // // TODO: Make a little nicer. + // writeln!(output, "{}", input.reflection().unwrap().render())?; + return Ok(()); } else { writeln!(output, "No input found for \"{}\"", describe)?; diff --git a/diskann-benchmark-runner/src/input.rs b/diskann-benchmark-runner/src/input.rs index 3868ac229c..8670d8eb6d 100644 --- a/diskann-benchmark-runner/src/input.rs +++ b/diskann-benchmark-runner/src/input.rs @@ -3,7 +3,7 @@ * Licensed under the MIT license. */ -use crate::{internal::visibility::Visibility, Checker}; +use crate::{internal::visibility::Visibility, Checker, Reflect, Reflection}; /// Inputs to [`Benchmarks`](crate::Benchmark). /// @@ -16,7 +16,7 @@ pub trait Input: Sized + std::fmt::Debug + 'static { /// [`Deserialize`](serde::Deserialize) implementation. /// /// Final object validation is performed via [`from_raw`](Self::from_raw). - type Raw: serde::de::DeserializeOwned + serde::Serialize; + type Raw: serde::de::DeserializeOwned + serde::Serialize + Reflect; /// Return the discriminant associated with this type. /// @@ -73,6 +73,11 @@ impl Registered<'_> { self.0.visibility() } + /// Return the [`Reflection`] for the raw input. + pub(crate) fn raw_reflection(&self) -> Option { + self.0.raw_reflection() + } + /// Return a `std::fmt::Display` implementation that pretty-prints the input tag as well /// as any visibility modifiers. pub(crate) fn display(&self) -> Display<'_> { @@ -125,7 +130,7 @@ pub(crate) fn order_inputs(a: &Registered<'_>, b: &Registered<'_>) -> std::cmp:: pub(crate) mod internal { use super::*; - use crate::Features; + use crate::{Features, Reflection}; /// Runtime representation of a deserialized [`Input`]. #[derive(Debug)] @@ -227,6 +232,7 @@ pub(crate) mod internal { fn visibility(&self) -> Visibility<'_>; // reflection + fn raw_reflection(&self) -> Option; fn as_any(&self) -> &dyn std::any::Any; fn type_name(&self) -> &'static str; } @@ -252,6 +258,9 @@ pub(crate) mod internal { fn visibility(&self) -> Visibility<'_> { Visibility::Available } + fn raw_reflection(&self) -> Option { + Some(Reflection::new::()) + } fn as_any(&self) -> &dyn std::any::Any { self } @@ -310,6 +319,9 @@ pub(crate) mod internal { features: &self.features, } } + fn raw_reflection(&self) -> Option { + None + } fn as_any(&self) -> &dyn std::any::Any { self } diff --git a/diskann-benchmark-runner/src/reflect/mod.rs b/diskann-benchmark-runner/src/reflect/mod.rs index 075948d5b4..211972b086 100644 --- a/diskann-benchmark-runner/src/reflect/mod.rs +++ b/diskann-benchmark-runner/src/reflect/mod.rs @@ -5,6 +5,7 @@ use std::{ any::TypeId, + collections::HashSet, fmt::{self, Write}, }; @@ -14,6 +15,9 @@ mod render; pub mod tree; pub use tree::Type; +#[cfg(test)] +mod test; + pub trait Reflect: 'static { fn ty() -> Type; fn format_type_name(f: &mut dyn Write) -> fmt::Result; @@ -49,6 +53,36 @@ impl Reflection { pub fn render(&self) -> Render { Render(*self) } + + // pub(crate) fn visit_all_reachable(&self, f: F) -> Result<(), E> + // where + // F: FnMut(Reflection) -> Result<(), E>, + // { + // let mut id_map = HashSet::new(); + // visit_all_reachable(*self, &mut id_map, f) + // } + + pub(crate) fn visit_unique(&self, mut f: F) -> Result<(), E> + where + F: FnMut(Reflection) -> Result, + { + let mut seen = HashSet::new(); + self.visit_with(|r: Reflection| { + if seen.insert(r.type_id()) { + f(r)?; + Ok(true) + } else { + Ok(false) + } + }) + } + + pub(crate) fn visit_with(&self, f: F) -> Result<(), E> + where + F: FnMut(Reflection) -> Result, + { + visit_with(*self, f) + } } impl fmt::Debug for Reflection { @@ -88,6 +122,49 @@ impl std::fmt::Display for Render { } } +//------------// +// Algorithms // +//------------// + +fn visit_with(mut reflection: Reflection, mut f: F) -> Result<(), E> +where + F: FnMut(Reflection) -> Result, +{ + use tree::Fields; + + let mut stack = Vec::new(); + loop { + // visit: expand this node if the closure returns `true`. + if f(reflection)? { + let mut push = |r: Reflection| stack.push(r); + let mut push_fields = |fields: &Fields| match fields { + Fields::Named(named) => named.iter().for_each(|field| push(field.field())), + Fields::Unnamed(unnamed) => unnamed.iter().for_each(|field| push(field.field())), + Fields::Unit => {} + }; + + // explore + match reflection.ty() { + // Nothing to do for primitives as there is no other object that can be reached. + Type::Primitive(_) => {} + Type::Aggregate(aggregate) => push_fields(aggregate.fields()), + Type::Enum(enum_) => enum_ + .variants() + .iter() + .for_each(|variant| push_fields(variant.fields())), + Type::Sequence(seq) => push(seq.element()), + } + } + + // loop + if let Some(r) = stack.pop() { + reflection = r; + } else { + break Ok(()); + } + } +} + /////////////// // Bootstrap // /////////////// @@ -119,127 +196,57 @@ macro_rules! primitive { }; } +primitive!((), "empty", "()"); primitive!(usize, "A system dependent unsigned integer", "usize"); primitive!(u32, "A 32-bit unsigned integer", "u32"); +primitive!(bool, "A value of \"true\" or \"false\"", "bool"); primitive!(String, "A string", "string"); -impl Reflect for Vec +impl Reflect for Option where T: Reflect, { fn ty() -> Type { - Type::sequence::(Some("An ordered collection of elements".into())) + Type::enum_( + // TODO: Untagged + tree::EnumRepr::External, + [ + tree::Variant::new( + "", + tree::Fields::Unit, + Some("Use `null` to indicate that this value does not exist".into()), + ), + tree::Variant::new( + "", + tree::Fields::unnamed([tree::UnnamedField::new::(None)]), + Some("Presence indicates the value is present".into()), + ), + ], + Some("An optional configuration".into()), + ) } fn format_type_name(f: &mut dyn Write) -> fmt::Result { - write!(f, "Vec<{}>", Reflection::new::().type_name()) + f.write_str("Option<")?; + T::format_type_name(f)?; + f.write_str(">") } } -#[derive(Reflect)] -pub struct UnitWithConst {} - -#[derive(Reflect)] -pub struct GenericBoundAdded { - uses_t: Vec, -} - -/// This is a test! -/// -/// Hello world! -#[derive(Reflect)] -pub struct Test { - /// This field affects this value. - a: usize, - - /// This field does something else. - b: usize, -} - -#[derive(Reflect)] -pub struct TestUnnamed( - /// Can I document this? - usize, - T, -); - -/// This is a nother test! -#[derive(Reflect)] -pub struct Test2 { - /// This field affects this value. - a: usize, - - /// This field doesn't have any names. - unnamed: TestUnnamed, - - other: Test, - - /// How are we going to compute distances? - metric: AdjacentEnum, - - /// These control a bunch of parameters. - seq: Vec, -} - -#[derive(Reflect)] -struct Wrapper { - /// Inner - a: T, -} - -/// An enum with no payloads. -#[derive(Debug, Clone, Copy, Reflect)] -pub enum Metric { - SquaredL2, - InnerProduct, - Cosine, -} +impl Reflect for Vec +where + T: Reflect, +{ + fn ty() -> Type { + Type::sequence::(Some("An ordered collection of elements".into())) + } -/// An enum with no payloads. -#[derive(Debug, Reflect)] -#[serde(rename_all = "kebab-case")] -pub enum AdjacentEnum { - SquaredL2, - /// Let me see if this works - InnerProduct(u32), - - /// Compute the cosine similarity - Cosine { - /// Thos actually doesn't do anything. - test: String, - }, + fn format_type_name(f: &mut dyn Write) -> fmt::Result { + write!(f, "Vec<{}>", Reflection::new::().type_name()) + } } -// impl Reflect for AdjacentEnum { -// fn ty() -> Type { -// Type::enum_( -// EnumRepr::Adjacent { -// tag: "enum-type", -// content: "content", -// }, -// [ -// Variant::new("squared-l2", Fields::Unit, None), -// Variant::new( -// "inner-product", -// Fields::unnamed([UnnamedField::new::(Some("testing".into()))]), -// Some("Inner Product with some payload".into()), -// ), -// Variant::new( -// "cosine", -// Fields::named([NamedField::new::("test", None)]), -// Some("Cosine Similarity".into()), -// ), -// ], -// Some("The similarity measure to use".into()), -// ) -// } -// -// fn format_type_name(f: &mut dyn Write) -> fmt::Result { -// f.write_str("Metric") -// } -// } - ////////////// // Internal // ////////////// @@ -288,278 +295,4 @@ pub(crate) mod internal { #[cfg(test)] mod tests { use super::*; - - use std::assert_matches; - - #[test] - fn test_unit() { - /// A unit struct. - #[derive(Reflect)] - struct Unit; - - let r = Reflection::new::(); - let ty = r.ty(); - - assert_eq!(r.ty().doc().unwrap(), "A unit struct."); - assert_eq!(r.type_name().to_string(), "Unit"); - assert!(!ty.has_body()); - } - - #[test] - fn test_unit_const_generic() { - /// A unit struct. - #[derive(Reflect)] - struct Unit; - - let r = Reflection::new::>(); - let ty = r.ty(); - - assert_eq!(ty.doc().unwrap(), "A unit struct."); - assert_eq!(r.type_name().to_string(), "Unit<10>"); - assert!(!ty.has_body()); - } - - #[test] - fn test_unit_const_generic_2() { - #[derive(Reflect)] - struct Unit; - - let r = Reflection::new::>(); - let ty = r.ty(); - - assert!(ty.doc().is_none()); - assert_eq!(r.type_name().to_string(), "Unit<10, 20>"); - assert!(!ty.has_body()); - } - - #[test] - fn test_empty_tuple_like() { - /// An empty tuple-like struct. - #[derive(Reflect)] - struct Empty(); - - let r = Reflection::new::(); - let ty = r.ty(); - - assert_eq!(ty.doc().unwrap(), "An empty tuple-like struct."); - assert_eq!(r.type_name().to_string(), "Empty"); - assert!(!ty.has_body()); - } - - #[test] - fn test_empty_struct_like() { - /// An empty struct. - #[derive(Reflect)] - struct Empty {} - - let r = Reflection::new::(); - let ty = r.ty(); - - assert_eq!(ty.doc().unwrap(), "An empty struct."); - assert_eq!(r.type_name().to_string(), "Empty"); - assert!(!ty.has_body()); - } - - #[test] - fn test_struct() { - /// A struct with two fields. - #[expect(unused)] - #[derive(Reflect)] - struct Woo { - /// Foo - foo: usize, - /// Bar - bar: usize, - } - - let r = Reflection::new::(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "A struct with two fields."); - assert_eq!(r.type_name().to_string(), "Woo"); - assert!(ty.has_body()); - - let f = ty.as_aggregate().unwrap().fields().as_named().unwrap(); - assert_eq!(f.len(), 2); - - assert_eq!(f[0].name, "foo"); - assert_eq!(f[0].doc().unwrap(), "Foo"); - assert_eq!(f[0].field.type_name().to_string(), "usize"); - - assert_eq!(f[1].name, "bar"); - assert_eq!(f[1].doc().unwrap(), "Bar"); - assert_eq!(f[1].field.type_name().to_string(), "usize"); - } - - #[test] - fn test_tuple1() { - /// A tuple with one field. - #[expect(unused)] - #[derive(Reflect)] - struct Tuple1( - /// Field 0. - usize, - ); - - let r = Reflection::new::(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "A tuple with one field."); - assert_eq!(r.type_name().to_string(), "Tuple1"); - - assert!(ty.has_body()); - - let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); - assert_eq!(f.len(), 1); - assert_eq!(f[0].doc().unwrap(), "Field 0."); - assert_eq!(f[0].field().type_name().to_string(), "usize"); - } - - #[test] - fn test_tuple2() { - #[expect(unused)] - #[derive(Reflect)] - struct Tuple2( - T, - /// Field 1. - Vec, - ); - - let r = Reflection::new::>(); - let ty = r.ty(); - assert!(ty.doc().is_none()); - assert_eq!(r.type_name().to_string(), "Tuple2"); - assert!(ty.has_body()); - - let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); - assert_eq!(f.len(), 2); - - assert!(f[0].doc().is_none()); - assert_eq!(f[0].field().type_name().to_string(), "usize"); - - assert_eq!(f[1].doc().unwrap(), "Field 1."); - assert_eq!(f[1].field().type_name().to_string(), "Vec"); - } - - #[test] - fn test_empty_enum() { - /// An empty enum. - #[derive(Reflect)] - enum Empty {} - - let r = Reflection::new::(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "An empty enum."); - assert_eq!(r.type_name().to_string(), "Empty"); - assert!(!ty.has_body()); - } - - #[test] - fn test_enum_with_generics() { - /// An enum with generics. - #[expect(unused)] - #[derive(Reflect)] - enum Either { - A(Vec), - /// It's a bee! - B( - /// Buzz buzz - B, - ), - } - - let r = Reflection::new::>(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "An enum with generics."); - assert_eq!(r.type_name().to_string(), "Either"); - assert!(ty.has_body()); - - let variants = ty.as_enum().unwrap().variants(); - assert_eq!(variants.len(), 2); - - // Variant 0 - assert!(variants[0].doc().is_none()); - assert_eq!(variants[0].name(), "A"); - assert!(variants[0].fields.has_body()); - - let f = variants[0].fields.as_unnamed().unwrap(); - assert_eq!(f.len(), 1); - assert!(f[0].doc().is_none()); - assert_eq!(f[0].field().type_name().to_string(), "Vec"); - - // Variant 1 - assert_eq!(variants[1].doc().unwrap(), "It's a bee!"); - assert_eq!(variants[1].name(), "B"); - assert!(variants[1].fields.has_body()); - - let f = variants[1].fields.as_unnamed().unwrap(); - assert_eq!(f.len(), 1); - assert_eq!(f[0].doc().unwrap(), "Buzz buzz"); - assert_eq!(f[0].field().type_name().to_string(), "string"); - } - - #[test] - fn test_enum_variants() { - /// All the enums. - #[expect(unused)] - #[derive(Reflect)] - #[serde(tag = "tag", content = "content")] - enum All { - /// A unit variant. - Unit, - /// A tuple-like variant. - Tuple( - /// Field 0. - usize, - String, - ), - /// Struct-like. - Struct { - foo: usize, - /// All the bars! - bar: u32, - }, - } - - let r = Reflection::new::(); - let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "All the enums."); - assert_eq!(r.type_name().to_string(), "All"); - assert!(ty.has_body()); - - let variants = ty.as_enum().unwrap().variants(); - assert_eq!(variants.len(), 3); - - // Variant 0 - assert_eq!(variants[0].doc().unwrap(), "A unit variant."); - assert_eq!(variants[0].name(), "Unit"); - assert!(!variants[0].fields.has_body()); - assert_matches!(variants[0].fields, Fields::Unit); - - // Variant 1 - assert_eq!(variants[1].doc().unwrap(), "A tuple-like variant."); - assert_eq!(variants[1].name(), "Tuple"); - assert!(variants[1].fields.has_body()); - - let f = variants[1].fields.as_unnamed().unwrap(); - assert_eq!(f.len(), 2); - - assert_eq!(f[0].doc().unwrap(), "Field 0."); - assert_eq!(f[0].field.type_name().to_string(), "usize"); - - assert!(f[1].doc().is_none()); - assert_eq!(f[1].field.type_name().to_string(), "string"); - - // Variant 2 - assert_eq!(variants[2].name(), "Struct"); - assert_eq!(variants[2].doc().unwrap(), "Struct-like."); - let f = variants[2].fields.as_named().unwrap(); - assert_eq!(f.len(), 2); - - assert_eq!(f[0].name, "foo"); - assert!(f[0].doc().is_none()); - assert_eq!(f[0].field.type_name().to_string(), "usize"); - - assert_eq!(f[1].name, "bar"); - assert_eq!(f[1].doc().unwrap(), "All the bars!"); - assert_eq!(f[1].field.type_name().to_string(), "u32"); - } } diff --git a/diskann-benchmark-runner/src/reflect/render.rs b/diskann-benchmark-runner/src/reflect/render.rs index 003c2aee22..ac9de2cfd6 100644 --- a/diskann-benchmark-runner/src/reflect/render.rs +++ b/diskann-benchmark-runner/src/reflect/render.rs @@ -8,10 +8,8 @@ use std::fmt::{self, Write}; use crate::utils::fmt::Quote; use super::{ + tree::{Aggregate, Enum, EnumRepr, Fields, NamedField, Sequence, Type, UnnamedField, Variant}, Reflection, - tree::{ - Type, Aggregate, Enum, EnumRepr, Sequence, Fields, NamedField, UnnamedField, Variant, - } }; const INDENT: usize = 2; @@ -187,14 +185,19 @@ impl<'a> Renderer<'a> { self.blank()?; self.line("Options:")?; self.indent(|r| { - let mut first = true; + let mut previous: Option<&Variant> = None; for variant in enum_.variants().iter() { - if !first { - r.blank()?; + // Decide whether or not to put a space before this variant. + // Spaces and be skipped if the previous one was a `Unit` with no docs. + if let Some(previous) = previous { + let skip = previous.fields().is_unit() && previous.doc().is_none(); + if !skip { + r.blank()?; + } } r.render_variant(variant)?; - first = false; + previous = Some(variant) } Ok(()) @@ -299,4 +302,3 @@ impl<'a> Renderer<'a> { }) } } - diff --git a/diskann-benchmark-runner/src/reflect/test.rs b/diskann-benchmark-runner/src/reflect/test.rs new file mode 100644 index 0000000000..5e2fa99a9d --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/test.rs @@ -0,0 +1,373 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use super::{ + tree::{Fields, Type}, + Reflect, Reflection, +}; +use std::assert_matches; + +#[test] +fn test_unit() { + /// A unit struct. + #[derive(Reflect)] + struct Unit; + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(r.ty().doc().unwrap(), "A unit struct."); + assert_eq!(r.type_name().to_string(), "Unit"); + assert!(!ty.has_body()); +} + +#[test] +fn test_unit_rename() { + /// A unit struct. + #[derive(Reflect)] + #[reflect(type_name = "something-else")] + struct Unit; + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(r.ty().doc().unwrap(), "A unit struct."); + assert_eq!(r.type_name().to_string(), "something-else"); + assert!(!ty.has_body()); +} + +#[test] +fn test_unit_prefix() { + /// A unit struct. + #[derive(Reflect)] + #[reflect(prefix = "module::")] + struct Unit; + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(r.ty().doc().unwrap(), "A unit struct."); + assert_eq!(r.type_name().to_string(), "module::Unit"); + assert!(!ty.has_body()); +} + +#[test] +fn test_unit_const_generic() { + /// A unit struct. + #[derive(Reflect)] + struct Unit; + + let r = Reflection::new::>(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "A unit struct."); + assert_eq!(r.type_name().to_string(), "Unit<10>"); + assert!(!ty.has_body()); +} + +#[test] +fn test_unit_const_generic_prefix() { + /// A unit struct. + #[derive(Reflect)] + #[reflect(prefix = "module::")] + struct Unit; + + let r = Reflection::new::>(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "A unit struct."); + assert_eq!(r.type_name().to_string(), "module::Unit<10>"); + assert!(!ty.has_body()); +} + +#[test] +fn test_unit_const_generic_2() { + #[derive(Reflect)] + struct Unit; + + let r = Reflection::new::>(); + let ty = r.ty(); + + assert!(ty.doc().is_none()); + assert_eq!(r.type_name().to_string(), "Unit<10, 20>"); + assert!(!ty.has_body()); +} + +#[test] +fn test_empty_tuple_like() { + /// An empty tuple-like struct. + #[derive(Reflect)] + struct Empty(); + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "An empty tuple-like struct."); + assert_eq!(r.type_name().to_string(), "Empty"); + assert!(!ty.has_body()); +} + +#[test] +fn test_empty_struct_like() { + /// An empty struct. + #[derive(Reflect)] + struct Empty {} + + let r = Reflection::new::(); + let ty = r.ty(); + + assert_eq!(ty.doc().unwrap(), "An empty struct."); + assert_eq!(r.type_name().to_string(), "Empty"); + assert!(!ty.has_body()); +} + +#[test] +fn test_struct() { + /// A struct with two fields. + #[expect(unused)] + #[derive(Reflect)] + struct Woo { + /// Foo + foo: usize, + /// Bar + bar: usize, + } + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "A struct with two fields."); + assert_eq!(r.type_name().to_string(), "Woo"); + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_named().unwrap(); + assert_eq!(f.len(), 2); + + assert_eq!(f[0].name(), "foo"); + assert_eq!(f[0].doc().unwrap(), "Foo"); + assert_eq!(f[0].field().type_name().to_string(), "usize"); + + assert_eq!(f[1].name(), "bar"); + assert_eq!(f[1].doc().unwrap(), "Bar"); + assert_eq!(f[1].field().type_name().to_string(), "usize"); +} + +#[test] +fn test_struct_rename() { + /// A struct with three fields. + #[expect(unused)] + #[derive(Reflect)] + #[serde(rename_all = "kebab-case")] + struct Woo { + /// Foo + foo_bar: usize, + /// Bar + baz: usize, + + #[serde(rename = "oops")] + biz: usize, + } + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "A struct with three fields."); + assert_eq!(r.type_name().to_string(), "Woo"); + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_named().unwrap(); + assert_eq!(f.len(), 3); + + assert_eq!( + f[0].name(), + "foo-bar", + "fields should be renamed by `rename_all`" + ); + assert_eq!(f[0].doc().unwrap(), "Foo"); + assert_eq!(f[0].field().type_name().to_string(), "usize"); + + assert_eq!(f[1].name(), "baz"); + assert_eq!(f[1].doc().unwrap(), "Bar"); + assert_eq!(f[1].field().type_name().to_string(), "usize"); + + assert_eq!( + f[2].name(), + "oops", + "explicit rename should take precedence" + ); + assert!(f[2].doc().is_none()); + assert_eq!(f[2].field().type_name().to_string(), "usize"); +} + +#[test] +fn test_tuple1() { + /// A tuple with one field. + #[expect(unused)] + #[derive(Reflect)] + struct Tuple1( + /// Field 0. + usize, + ); + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "A tuple with one field."); + assert_eq!(r.type_name().to_string(), "Tuple1"); + + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); + assert_eq!(f.len(), 1); + assert_eq!(f[0].doc().unwrap(), "Field 0."); + assert_eq!(f[0].field().type_name().to_string(), "usize"); +} + +#[test] +fn test_tuple2() { + #[expect(unused)] + #[derive(Reflect)] + struct Tuple2( + T, + /// Field 1. + Vec, + ); + + let r = Reflection::new::>(); + let ty = r.ty(); + assert!(ty.doc().is_none()); + assert_eq!(r.type_name().to_string(), "Tuple2"); + assert!(ty.has_body()); + + let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); + assert_eq!(f.len(), 2); + + assert!(f[0].doc().is_none()); + assert_eq!(f[0].field().type_name().to_string(), "usize"); + + assert_eq!(f[1].doc().unwrap(), "Field 1."); + assert_eq!(f[1].field().type_name().to_string(), "Vec"); +} + +#[test] +fn test_empty_enum() { + /// An empty enum. + #[derive(Reflect)] + enum Empty {} + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "An empty enum."); + assert_eq!(r.type_name().to_string(), "Empty"); + assert!(!ty.has_body()); +} + +#[test] +fn test_enum_with_generics() { + /// An enum with generics. + #[expect(unused)] + #[derive(Reflect)] + enum Either { + A(Vec), + /// It's a bee! + B( + /// Buzz buzz + B, + ), + } + + let r = Reflection::new::>(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "An enum with generics."); + assert_eq!(r.type_name().to_string(), "Either"); + assert!(ty.has_body()); + + let variants = ty.as_enum().unwrap().variants(); + assert_eq!(variants.len(), 2); + + // Variant 0 + assert!(variants[0].doc().is_none()); + assert_eq!(variants[0].name(), "A"); + assert!(variants[0].fields().has_body()); + + let f = variants[0].fields().as_unnamed().unwrap(); + assert_eq!(f.len(), 1); + assert!(f[0].doc().is_none()); + assert_eq!(f[0].field().type_name().to_string(), "Vec"); + + // Variant 1 + assert_eq!(variants[1].doc().unwrap(), "It's a bee!"); + assert_eq!(variants[1].name(), "B"); + assert!(variants[1].fields().has_body()); + + let f = variants[1].fields().as_unnamed().unwrap(); + assert_eq!(f.len(), 1); + assert_eq!(f[0].doc().unwrap(), "Buzz buzz"); + assert_eq!(f[0].field().type_name().to_string(), "string"); +} + +#[test] +fn test_enum_variants() { + /// All the enums. + #[expect(unused)] + #[derive(Reflect)] + #[serde(tag = "tag", content = "content")] + enum All { + /// A unit variant. + Unit, + /// A tuple-like variant. + Tuple( + /// Field 0. + usize, + String, + ), + /// Struct-like. + Struct { + foo: usize, + /// All the bars! + bar: u32, + }, + } + + let r = Reflection::new::(); + let ty = r.ty(); + assert_eq!(ty.doc().unwrap(), "All the enums."); + assert_eq!(r.type_name().to_string(), "All"); + assert!(ty.has_body()); + + let variants = ty.as_enum().unwrap().variants(); + assert_eq!(variants.len(), 3); + + // Variant 0 + assert_eq!(variants[0].doc().unwrap(), "A unit variant."); + assert_eq!(variants[0].name(), "Unit"); + assert!(!variants[0].fields().has_body()); + assert_matches!(variants[0].fields(), Fields::Unit); + + // Variant 1 + assert_eq!(variants[1].doc().unwrap(), "A tuple-like variant."); + assert_eq!(variants[1].name(), "Tuple"); + assert!(variants[1].fields().has_body()); + + let f = variants[1].fields().as_unnamed().unwrap(); + assert_eq!(f.len(), 2); + + assert_eq!(f[0].doc().unwrap(), "Field 0."); + assert_eq!(f[0].field().type_name().to_string(), "usize"); + + assert!(f[1].doc().is_none()); + assert_eq!(f[1].field().type_name().to_string(), "string"); + + // Variant 2 + assert_eq!(variants[2].name(), "Struct"); + assert_eq!(variants[2].doc().unwrap(), "Struct-like."); + let f = variants[2].fields().as_named().unwrap(); + assert_eq!(f.len(), 2); + + assert_eq!(f[0].name(), "foo"); + assert!(f[0].doc().is_none()); + assert_eq!(f[0].field().type_name().to_string(), "usize"); + + assert_eq!(f[1].name(), "bar"); + assert_eq!(f[1].doc().unwrap(), "All the bars!"); + assert_eq!(f[1].field().type_name().to_string(), "u32"); +} diff --git a/diskann-benchmark-runner/src/reflect/tree.rs b/diskann-benchmark-runner/src/reflect/tree.rs index dfb3bcfce7..4dcf85bacb 100644 --- a/diskann-benchmark-runner/src/reflect/tree.rs +++ b/diskann-benchmark-runner/src/reflect/tree.rs @@ -57,7 +57,7 @@ impl Type { } } - fn as_aggregate(&self) -> Option<&Aggregate> { + pub(super) fn as_aggregate(&self) -> Option<&Aggregate> { if let Self::Aggregate(aggregate) = self { Some(aggregate) } else { @@ -65,7 +65,7 @@ impl Type { } } - fn as_enum(&self) -> Option<&Enum> { + pub(super) fn as_enum(&self) -> Option<&Enum> { if let Self::Enum(enum_) = self { Some(enum_) } else { @@ -169,7 +169,7 @@ impl Fields { Self::Unnamed(itr.into_iter().collect()) } - fn as_named(&self) -> Option<&[NamedField]> { + pub(super) fn as_named(&self) -> Option<&[NamedField]> { if let Self::Named(fields) = self { Some(fields) } else { @@ -177,13 +177,17 @@ impl Fields { } } - fn as_unnamed(&self) -> Option<&[UnnamedField]> { + pub(super) fn as_unnamed(&self) -> Option<&[UnnamedField]> { if let Self::Unnamed(fields) = self { Some(fields) } else { None } } + + pub(super) fn is_unit(&self) -> bool { + matches!(self, Self::Unit) + } } #[derive(Debug)] @@ -242,7 +246,6 @@ impl UnnamedField { pub(super) fn doc(&self) -> Option<&str> { self.doc.as_deref() } - } //------// @@ -381,5 +384,3 @@ impl Sequence { self.doc.as_deref() } } - - diff --git a/diskann-benchmark-runner/src/registry.rs b/diskann-benchmark-runner/src/registry.rs index 82429114f8..040e2d5b83 100644 --- a/diskann-benchmark-runner/src/registry.rs +++ b/diskann-benchmark-runner/src/registry.rs @@ -3,21 +3,34 @@ * Licensed under the MIT license. */ -use std::collections::{hash_map::Entry, HashMap}; +use std::{ + any::TypeId, + collections::{hash_map::Entry, HashMap}, +}; +use hashbrown::{hash_set, HashSet}; use thiserror::Error; use crate::{ benchmark::{self, internal::AnnotatedMatch, Benchmark, MatchContext, Regression, Score}, input, internal::visibility::Visibility, - Checkpoint, Features, Input, Output, + Checkpoint, Features, Input, Output, Reflection, }; /// A collection of registered inputs and benchmarks. pub struct Registry { - // Inputs keyed by their tag type. + /// Inputs keyed by their tag type. inputs: HashMap<&'static str, Box>, + + /// Registered types for documentation. + /// + /// The keys are generated by [`Reflection::type_name`]. We store the associated + /// [`Relflection`] and [`TypeId`] to ensure that registered type-names are unique. + name_map: HashMap, + type_ids: HashSet, + + /// The registered benchmarks. benchmarks: Vec, } @@ -26,6 +39,8 @@ impl Registry { pub fn new() -> Self { Self { inputs: HashMap::new(), + name_map: HashMap::new(), + type_ids: HashSet::new(), benchmarks: Vec::new(), } } @@ -229,6 +244,40 @@ impl Registry { let wrapper = crate::input::internal::Wrapper::::new(); match self.inputs.entry(tag) { Entry::Vacant(v) => { + // Before we insert - try to add all the reflection types. If this fails, + // then we haven't committed the input to the registry. + // + // TODO: Enable roll-back of any newly registered type. + Reflection::new::() + .visit_with(|reflection: Reflection| -> Result { + let type_id = reflection.type_id(); + + // If the type doesn't exist - try to add it via type-name. + // Only if all that succeeds to we commit the type-id. + if let hash_set::Entry::Vacant(vacant) = self.type_ids.entry(type_id) { + let type_name = reflection.type_name().to_string(); + match self.name_map.entry(type_name) { + Entry::Vacant(v) => { + v.insert((reflection, type_id)); + } + Entry::Occupied(o) => { + if o.get().1 != type_id { + panic!("eventually we will fail gracetully"); + } + } + } + + vacant.insert(); + + // Continue exploring. + Ok(true) + } else { + // Already seen this node. No need to recurse. + Ok(false) + } + }) + .unwrap(); + v.insert(Box::new(wrapper)); Ok(()) } diff --git a/diskann-benchmark-runner/src/test/dim.rs b/diskann-benchmark-runner/src/test/dim.rs index 61dc07a1b4..6d99c15d05 100644 --- a/diskann-benchmark-runner/src/test/dim.rs +++ b/diskann-benchmark-runner/src/test/dim.rs @@ -9,14 +9,14 @@ use serde::{Deserialize, Serialize}; use crate::{ benchmark::{MatchContext, PassFail, Regression, Score}, - Benchmark, Checker, Checkpoint, Input, Output, + Benchmark, Checker, Checkpoint, Input, Output, Reflect, }; /////////// // Input // /////////// -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Reflect)] pub(super) struct DimInput { dim: Option, } @@ -55,7 +55,7 @@ impl Input for DimInput { // Tolerance // /////////////// -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Reflect)] pub(super) struct Tolerance { succeed: bool, error_in_check: bool, diff --git a/diskann-benchmark-runner/src/test/gated.rs b/diskann-benchmark-runner/src/test/gated.rs index 1391c36223..0d264d65eb 100644 --- a/diskann-benchmark-runner/src/test/gated.rs +++ b/diskann-benchmark-runner/src/test/gated.rs @@ -17,6 +17,7 @@ use serde::{Deserialize, Serialize}; use crate::{ benchmark::MatchContext, benchmark::Score, Benchmark, Checker, Checkpoint, Input, Output, + Reflect, }; use super::{dim::DimInput, typed::TypeInput}; @@ -85,7 +86,7 @@ impl Benchmark for AnotherGatedBench { // Partially Gated with Input Always Registered // ////////////////////////////////////////////////// -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Reflect)] pub(super) struct SampleInput { value: String, } @@ -145,7 +146,7 @@ impl Benchmark for GatedWithIndependentInput { // The input backing the fully-gated benchmark. This is only compiled and registered when the // controlling feature set is enabled, standing in for an input whose validation would otherwise // pull in a heavy optional dependency. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Reflect)] pub(super) struct PhantomInput { value: usize, } diff --git a/diskann-benchmark-runner/src/test/typed.rs b/diskann-benchmark-runner/src/test/typed.rs index dd94f222e4..5bd900f9db 100644 --- a/diskann-benchmark-runner/src/test/typed.rs +++ b/diskann-benchmark-runner/src/test/typed.rs @@ -10,7 +10,7 @@ use serde::{Deserialize, Serialize}; use crate::{ benchmark::{MatchContext, PassFail, Regression, Score}, utils::datatype::{AsDataType, DataType}, - Benchmark, Checker, Checkpoint, Input, Output, + Benchmark, Checker, Checkpoint, Input, Output, Reflect, }; /////////// @@ -24,11 +24,11 @@ pub(crate) struct TypeInput { error_when_checked: bool, } -#[derive(Serialize, Deserialize)] +#[derive(Serialize, Deserialize, Reflect)] pub(crate) struct TypeInputRaw { data_type: DataType, dim: usize, - // Should we return an error when deserializing? + /// Should we return an error when deserializing? error_when_checked: bool, } @@ -78,9 +78,9 @@ impl Input for TypeInput { // Tolerance // /////////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] pub(super) struct Tolerance { - // Should we return an error when `from_raw` is called? + /// Should we return an error when `from_raw` is called? pub(super) error_when_checked: bool, } diff --git a/diskann-benchmark-runner/src/utils/datatype.rs b/diskann-benchmark-runner/src/utils/datatype.rs index 470251af01..5bf13d6d8c 100644 --- a/diskann-benchmark-runner/src/utils/datatype.rs +++ b/diskann-benchmark-runner/src/utils/datatype.rs @@ -6,10 +6,12 @@ use half::f16; use serde::{Deserialize, Serialize}; +use crate::Reflect; + /// An enum representation for common DiskANN data types. /// /// See also: [`AsDataType`]. -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "lowercase")] pub enum DataType { Float64, diff --git a/diskann-benchmark-runner/temp/main.rs b/diskann-benchmark-runner/temp/main.rs index 90d24fa9c6..cbdce829de 100644 --- a/diskann-benchmark-runner/temp/main.rs +++ b/diskann-benchmark-runner/temp/main.rs @@ -8,11 +8,11 @@ use diskann_benchmark_runner::{reflect, Reflection}; fn main() -> anyhow::Result<()> { - println!("{}", Reflection::new::().render()); + // println!("{}", Reflection::new::().render()); - println!("{}", Reflection::new::().render()); + // println!("{}", Reflection::new::().render()); - println!("{}", Reflection::new::().render()); + // println!("{}", Reflection::new::().render()); Ok(()) } From 979d84fc72f911a1ab73e0d25c289fbdf39120fe Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Mon, 28 Sep 2026 14:28:45 -0700 Subject: [PATCH 11/21] Serde tests. --- diskann-benchmark-runner-derive/src/lib.rs | 52 +- diskann-benchmark-runner/src/app.rs | 2 +- diskann-benchmark-runner/src/reflect/mod.rs | 77 +- .../src/reflect/render.rs | 1 + .../src/reflect/test/mod.rs | 10 + .../src/reflect/test/serde.rs | 715 ++++++++++++++++++ .../src/reflect/{test.rs => test/unit.rs} | 57 +- diskann-benchmark-runner/src/reflect/tree.rs | 42 +- diskann-benchmark-runner/src/test/typed.rs | 1 + diskann-benchmark-runner/temp/main.rs | 2 +- 10 files changed, 878 insertions(+), 81 deletions(-) create mode 100644 diskann-benchmark-runner/src/reflect/test/mod.rs create mode 100644 diskann-benchmark-runner/src/reflect/test/serde.rs rename diskann-benchmark-runner/src/reflect/{test.rs => test/unit.rs} (85%) diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index 089afa74ac..6da9748d21 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -215,7 +215,7 @@ fn generate_type_name_body( .collect(); // Check that the type-name attributes are compatible with the struct. - let prefix: Vec = match type_name { + let prefix: Option = match type_name { attributes::TypeName::Rename(rename) => { if arguments.is_empty() { let ts = quote! { @@ -228,20 +228,20 @@ fn generate_type_name_body( "The `type_name` attribute cannot be applied to types with generics", )); } - }, + } attributes::TypeName::Prefix(prefix) => { let ts = quote! { f.write_str(#prefix)?; }; - vec![ts] - }, - attributes::TypeName::None => Vec::new(), + Some(ts) + } + attributes::TypeName::None => None, }; // If there are no generics, we can dump the typename directly. if arguments.is_empty() { let ts = quote! { - #(#prefix)* + #prefix f.write_str(#name) }; Ok(ts) @@ -260,7 +260,7 @@ fn generate_type_name_body( }); let ts = quote! { - #(#prefix)* + #prefix f.write_str(#name)?; f.write_str("<")?; #(#writes)* @@ -286,7 +286,16 @@ fn build_fields( Fields::Unnamed(fields) => { add_field_bounds(generics, &fields.unnamed); let list = unnamed_fields(&fields.unnamed); - Ok(quote!(#path::tree::Fields::Unnamed(vec![#(#list),*]))) + + // Unnamed fields of length 1 become new-types instead. + let ts = if list.len() == 1 { + let new_type = &list[0]; + quote!(#path::tree::Fields::NewType(#new_type)) + } else { + quote!(#path::tree::Fields::Unnamed(vec![#(#list),*])) + }; + + Ok(ts) } Fields::Unit => Ok(quote!(#path::tree::Fields::Unit)), } @@ -309,7 +318,7 @@ where .as_ref() .expect("named fields should have identifiers"); - let name = syn::LitStr::new(&ident.to_string(), ident.span()); + let name = syn::LitStr::new(strip_raw_prefix(&ident.to_string()), ident.span()); let doc = format_docstrings(&f.attrs); let attributes::Field { rename_field } = attributes::Field::parse(&f.attrs)?; @@ -319,16 +328,19 @@ where .collect() } -fn unnamed_fields<'a, I>(fields: I) -> impl Iterator +fn unnamed_fields<'a, I>(fields: I) -> Vec where I: IntoIterator, { let path = crate_name(); - fields.into_iter().map(move |f| { - let ty = &f.ty; - let doc = format_docstrings(&f.attrs); - quote_spanned! { ty.span()=> #path::tree::UnnamedField::new::<#ty>(#doc) } - }) + fields + .into_iter() + .map(move |f| { + let ty = &f.ty; + let doc = format_docstrings(&f.attrs); + quote_spanned! { ty.span()=> #path::tree::UnnamedField::new::<#ty>(#doc) } + }) + .collect() } /// Generate the `Reflect` implementation. @@ -403,7 +415,7 @@ fn process_enum( .iter() .map(|v| -> syn::Result { let doc = format_docstrings(&v.attrs); - let name = syn::LitStr::new(&v.ident.to_string(), v.ident.span()); + let name = syn::LitStr::new(strip_raw_prefix(&v.ident.to_string()), v.ident.span()); let attributes::Variant { rename_variant, rename_variant_fields, @@ -487,3 +499,11 @@ fn extract_docs(attributes: &[syn::Attribute]) -> Option { Some(docstrings.join("\n")) } } + +//-----// +// raw // +//-----// + +fn strip_raw_prefix(s: &str) -> &str { + s.strip_prefix("r#").unwrap_or(s) +} diff --git a/diskann-benchmark-runner/src/app.rs b/diskann-benchmark-runner/src/app.rs index 1714d8959c..55dd3e75b3 100644 --- a/diskann-benchmark-runner/src/app.rs +++ b/diskann-benchmark-runner/src/app.rs @@ -215,7 +215,7 @@ impl App { writeln!(output, "{}", serde_json::to_string_pretty(&repr)?)?; // // TODO: Make a little nicer. - // writeln!(output, "{}", input.reflection().unwrap().render())?; + // writeln!(output, "{}", input.raw_reflection().unwrap().render())?; return Ok(()); } else { diff --git a/diskann-benchmark-runner/src/reflect/mod.rs b/diskann-benchmark-runner/src/reflect/mod.rs index 211972b086..93c1ab16d3 100644 --- a/diskann-benchmark-runner/src/reflect/mod.rs +++ b/diskann-benchmark-runner/src/reflect/mod.rs @@ -5,7 +5,6 @@ use std::{ any::TypeId, - collections::HashSet, fmt::{self, Write}, }; @@ -54,29 +53,6 @@ impl Reflection { Render(*self) } - // pub(crate) fn visit_all_reachable(&self, f: F) -> Result<(), E> - // where - // F: FnMut(Reflection) -> Result<(), E>, - // { - // let mut id_map = HashSet::new(); - // visit_all_reachable(*self, &mut id_map, f) - // } - - pub(crate) fn visit_unique(&self, mut f: F) -> Result<(), E> - where - F: FnMut(Reflection) -> Result, - { - let mut seen = HashSet::new(); - self.visit_with(|r: Reflection| { - if seen.insert(r.type_id()) { - f(r)?; - Ok(true) - } else { - Ok(false) - } - }) - } - pub(crate) fn visit_with(&self, f: F) -> Result<(), E> where F: FnMut(Reflection) -> Result, @@ -140,6 +116,7 @@ where let mut push_fields = |fields: &Fields| match fields { Fields::Named(named) => named.iter().for_each(|field| push(field.field())), Fields::Unnamed(unnamed) => unnamed.iter().for_each(|field| push(field.field())), + Fields::NewType(newtype) => push(newtype.field()), Fields::Unit => {} }; @@ -174,7 +151,7 @@ where T: Reflect, { fn ty() -> Type { - Type::primitive(None) + Type::primitive(tree::PrimitiveKind::Null, None) } fn format_type_name(f: &mut dyn Write) -> fmt::Result { @@ -183,10 +160,10 @@ where } macro_rules! primitive { - ($T:ty, $doc:literal, $type_name:literal) => { + ($T:ty, $kind:ident, $doc:literal, $type_name:literal) => { impl Reflect for $T { fn ty() -> Type { - Type::primitive(Some($doc.into())) + Type::primitive(tree::PrimitiveKind::$kind, Some($doc.into())) } fn format_type_name(f: &mut dyn Write) -> fmt::Result { @@ -196,12 +173,37 @@ macro_rules! primitive { }; } -primitive!((), "empty", "()"); -primitive!(usize, "A system dependent unsigned integer", "usize"); -primitive!(u32, "A 32-bit unsigned integer", "u32"); -primitive!(bool, "A value of \"true\" or \"false\"", "bool"); - -primitive!(String, "A string", "string"); +primitive!((), Null, "empty", "()"); +primitive!( + usize, + Number, + "A system dependent unsigned integer", + "usize" +); +primitive!(isize, Number, "A system dependent signed integer", "isize"); + +primitive!(u8, Number, "An 8-bit unsigned integer", "u8"); +primitive!(u16, Number, "A 16-bit unsigned integer", "u16"); +primitive!(u32, Number, "A 32-bit unsigned integer", "u32"); +primitive!(u64, Number, "A 64-bit unsigned integer", "u64"); + +primitive!(i8, Number, "An 8-bit signed integer", "i8"); +primitive!(i16, Number, "A 16-bit signed integer", "i16"); +primitive!(i32, Number, "A 32-bit signed integer", "i32"); +primitive!(i64, Number, "A 64-bit signed integer", "i64"); + +primitive!(f32, Number, "An 32-bit floating-point number", "f32"); +primitive!(f64, Number, "An 64-bit floating-point number", "f64"); + +primitive!( + std::num::NonZeroUsize, + Number, + "A system dependent, non-zero, unsigned integer", + "NonZero" +); + +primitive!(bool, Boolean, "A value of \"true\" or \"false\"", "bool"); +primitive!(String, String, "A string", "string"); impl Reflect for Option where @@ -287,12 +289,3 @@ pub(crate) mod internal { } } } - -/////////// -// Tests // -/////////// - -#[cfg(test)] -mod tests { - use super::*; -} diff --git a/diskann-benchmark-runner/src/reflect/render.rs b/diskann-benchmark-runner/src/reflect/render.rs index ac9de2cfd6..e7a7d4e8b7 100644 --- a/diskann-benchmark-runner/src/reflect/render.rs +++ b/diskann-benchmark-runner/src/reflect/render.rs @@ -236,6 +236,7 @@ impl<'a> Renderer<'a> { self.render_unnamed_field(i, field)?; } } + Fields::NewType(_) => todo!(), Fields::Unit => {} } diff --git a/diskann-benchmark-runner/src/reflect/test/mod.rs b/diskann-benchmark-runner/src/reflect/test/mod.rs new file mode 100644 index 0000000000..359f96c9af --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/test/mod.rs @@ -0,0 +1,10 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +// Directed macro unit test. +mod unit; + +// Serde compatibility. +mod serde; diff --git a/diskann-benchmark-runner/src/reflect/test/serde.rs b/diskann-benchmark-runner/src/reflect/test/serde.rs new file mode 100644 index 0000000000..b5313d2fd4 --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/test/serde.rs @@ -0,0 +1,715 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +//! [`Reflect`] macro serde-compatibility tests. +//! +//! Serde does allow some differences between serialization and deserialization. We don't +//! care to support such fanciness in our benchmark inputs and outputs, especially since we +//! need example JSONs to be round-trippable back to their serialized representations. +//! +//! As such, the main entry points ([`check_serde`] and [`check_enums`]) just take [`Serialize`] +//! bounds. + +use std::{assert_matches, borrow::Cow}; + +use hashbrown::HashSet; +use serde::Serialize; +use serde_json::Value; + +use crate::reflect::{tree, Reflect, Reflection, Type}; + +/// Check that the serialized representation of `s` in JSON matches the [`Reflection`] +/// generated for this type. +fn check_struct(s: T) +where + T: Serialize + Reflect, +{ + let r = Reflection::new::(); + let val = serde_json::to_value(s).unwrap(); + check_reflection( + r, + &val, + Context::new(&val, &val, format_args!("struct: {}", r.type_name())), + ); +} + +/// Check that the serialized representations of all the variants of the enum `T` match the +/// [`Reflection`] for `T`. +/// +/// Note that `examples` is expected to contain examples of all variants in declaration order +/// and this function will panic if this is not the case. +fn check_enums(examples: &[T], ctx: std::fmt::Arguments<'_>) +where + T: Serialize + Reflect, +{ + let r = Reflection::new::(); + let serialized: Vec<_> = examples + .iter() + .map(|e| serde_json::to_value(e).unwrap()) + .collect(); + + let Type::Enum(e) = r.ty() else { + panic!("expected an enum for type {}", r.type_name()); + }; + + check_enum_variants(&e, &serialized, ctx); +} + +//////////////////// +// Implementation // +//////////////////// + +/// A context for displaying where we are in the type tree. +/// +/// Contains the top level JSON we're working, the current JSON, and a stack of +/// [`std::fmt::Argument`]s that describe the sequence of operations that led us into a mess. +/// +/// Use the [`context`] macro for creating nexted contexts. +#[derive(Debug, Clone, Copy)] +struct Context<'a> { + top: &'a Value, + current: &'a Value, + stack: std::fmt::Arguments<'a>, +} + +impl<'a> Context<'a> { + fn new(top: &'a Value, current: &'a Value, stack: std::fmt::Arguments<'a>) -> Self { + Self { + top, + current, + stack, + } + } +} + +impl std::fmt::Display for Context<'_> { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!( + f, + "Full JSON:\n\n{}\n\nCurrent JSON:\n\n{}\n\nContext: {}", + serde_json::to_string_pretty(self.top).unwrap(), + serde_json::to_string_pretty(self.current).unwrap(), + self.stack + ) + } +} + +macro_rules! context { + ($current:ident, $next:ident, $fmt:expr) => { + Context { + top: $current.top, + current: $next, + stack: format_args!( + "{}\n- {}", + $current.stack, + format_args!($fmt), + ), + } + }; + ($current:ident, $next:ident, $fmt:expr, $($args:tt)*) => { + Context { + top: $current.top, + current: $next, + stack: format_args!( + "{}\n- {}", + $current.stack, + format_args!($fmt, $($args)*), + ), + } + }; +} + +fn check_reflection(r: Reflection, s: &Value, ctx: Context<'_>) { + check_type( + &r.ty(), + s, + context!(ctx, s, "type_name = {}", r.type_name()), + ); +} + +fn check_type(ty: &Type, s: &Value, ctx: Context<'_>) { + match ty { + Type::Primitive(p) => check_primitive(p, s, context!(ctx, s, "primitive")), + Type::Aggregate(a) => check_fields(a.fields(), s, context!(ctx, s, "aggregate")), + Type::Enum(e) => check_enum(e, s, context!(ctx, s, "enum")), + Type::Sequence(seq) => check_sequence(seq, s, context!(ctx, s, "sequence")), + } +} + +fn check_primitive(p: &tree::Primitive, s: &Value, ctx: Context<'_>) { + // Verify that `s` has the correct `PrimitiveKind`. + match p.kind() { + tree::PrimitiveKind::Null => { + if !s.is_null() { + panic!("expected `Null`\n\n{}", ctx); + } + } + tree::PrimitiveKind::Boolean => { + if !s.is_boolean() { + panic!("expected `Bool`\n\n{}", ctx); + } + } + tree::PrimitiveKind::Number => { + if !s.is_number() { + panic!("expected `Number`\n\n{}", ctx); + } + } + tree::PrimitiveKind::String => { + if !s.is_string() { + panic!("expected `String`\n\n{}", ctx); + } + } + } +} + +fn check_fields(fields: &tree::Fields, s: &Value, ctx: Context<'_>) { + match fields { + tree::Fields::Named(named_fields) => { + // All names must be unique. + assert_all_unique(named_fields.iter().map(|f| f.name()), ctx); + let map = value_as_map(s, ctx); + assert_eq!( + map.len(), + named_fields.len(), + "mismatch in number of fields\n\n{}", + ctx, + ); + + for f in named_fields { + let val = match map.get(f.name()) { + Some(val) => val, + None => panic!("Could not match field name \"{}\"\n\n{}", f.name(), ctx), + }; + + check_reflection( + f.field(), + &val, + context!(ctx, val, "named field \"{}\"", f.name()), + ); + } + } + tree::Fields::Unnamed(unnamed_fields) => { + let array = value_as_array(s, ctx); + assert_eq!( + array.len(), + unnamed_fields.len(), + "mismatch in number of fields\n\n{}", + ctx + ); + + for (i, (f, val)) in std::iter::zip(unnamed_fields.iter(), array.iter()).enumerate() { + check_reflection(f.field(), val, context!(ctx, val, "unnamed field {}", i)); + } + } + tree::Fields::NewType(newtype) => { + check_reflection(newtype.field(), s, context!(ctx, s, "newtype")) + } + tree::Fields::Unit => { + assert_matches!(s, Value::Null, "expected a unit\n{}", ctx,); + } + } +} + +fn check_enum_variants(e: &tree::Enum, s: &[Value], ctx: std::fmt::Arguments<'_>) { + let num_variants = e.variants().len(); + assert_eq!(num_variants, s.len(), "expected one example per variant"); + + for (i, (variant, example)) in std::iter::zip(e.variants().iter(), s.iter()).enumerate() { + let ctx = Context::new(example, example, ctx); + + // Verify that the extracted tag matches. + let (tag, _content) = extract_tag_and_content(e.repr(), example, ctx); + assert_eq!(variant.name(), tag, "{}", ctx); + + // Use `check_enum`. + // + // This is a little wasteful because we've already extracted the tag and the enum, + // but this is test-only code so we can afford it. + check_enum( + e, + example, + context!( + ctx, + example, + "variant = {} ({} of {})", + variant.name(), + i + 1, + num_variants + ), + ); + } +} + +/// To summarize. For a struct like +/// ``` +/// enum Foo { +/// Baz { +/// a: usize, +/// }, +/// } +/// ``` +/// +/// External looks like +/// ```json +/// { +/// "baz": { +/// "a": 10 +/// } +/// } +/// ``` +/// Internal (say "tag = "mytag") looks like +/// ```json +/// { +/// "mytag": "baz", +/// "a": 10, +/// } +/// ``` +/// Adjaceny (say "tag = "mytag", "content" = "mycontent") looks like +/// ```json +/// { +/// "mytag": "baz", +/// "mycontent": { +/// "a": 10 +/// } +/// } +/// ``` +fn extract_tag_and_content<'a>( + repr: &tree::EnumRepr, + s: &'a Value, + ctx: Context<'_>, +) -> (&'a str, Option>) { + match repr { + // When using external tagging, the associated value can look like one of two things: + // + // 1. A raw string. This is only valid if the associated payload is a unit type. + // 2. A map consisting of a single key-value pair. The `key` is the enum tag. + tree::EnumRepr::External => match s { + Value::String(s) => (s, None), + Value::Object(m) => { + assert_eq!( + m.len(), + 1, + "externally tagged enums should only have a single key-value pair\n\n{}", + ctx + ); + + let kv: Vec<_> = m.iter().collect(); + (&kv[0].0, Some(Cow::Borrowed(&kv[0].1))) + } + _ => panic!("invalid representation\n\n{}", ctx), + }, + + // For internally tagged enums - we remove the tag after revrieval. + // + // If the remaining dictionary is empty, then we change it to `None`. + // This is technically a little ambiguous between unit variants and empty struct + // variants (e.g. `Enum::Unit` and `Enum::Empty {}`. + // + // However, we teach the latter case to expect empty results with internal tagging + // for purposes of the check. + tree::EnumRepr::Internal { tag } => { + let map = value_as_map(s, ctx); + + let t = match map.get(*tag) { + Some(value) => value_as_str(value, context!(ctx, value, "tag extraction")), + None => panic!("Could not find tag \"{}\"\n\n{}", tag, ctx), + }; + + // Delete the tag field to reuse the rest of the checking infrastructure. + let mut map = map.clone(); + map.remove(*tag); + if map.is_empty() { + (t, None) + } else { + (t, Some(Cow::Owned(Value::Object(map)))) + } + } + + // For adjacent tagging, the "content" field is ommitted when the corresponding + // variant is a unit variant. + tree::EnumRepr::Adjacent { tag, content } => { + let outer = value_as_map(s, ctx); + assert!(outer.len() == 1 || outer.len() == 2); + + let t = match outer.get(*tag) { + Some(value) => value_as_str(value, context!(ctx, value, "tag extraction")), + None => panic!("expected tag field \"{}\"", tag), + }; + + if outer.len() == 1 { + (t, None) + } else { + let c = match outer.get(*content) { + Some(c) => c, + None => panic!("expected content field \"{}\"", content), + }; + + (t, Some(Cow::Borrowed(c))) + } + } + } +} + +// See the inline notes in `extract_tag_and_content`. +fn is_no_content_expected(repr: &tree::EnumRepr, variant: &tree::Variant) -> bool { + match repr { + tree::EnumRepr::External => variant.fields().is_unit(), + tree::EnumRepr::Internal { .. } => match variant.fields() { + tree::Fields::Named(named_fields) => named_fields.is_empty(), + tree::Fields::Unnamed(_) => unreachable!("unimplemented by serde"), + // Only allowed for new-type wrappers that themselves are structs. As sucn, + // we always expect a payload. + tree::Fields::NewType(_) => false, + tree::Fields::Unit => true, + }, + tree::EnumRepr::Adjacent { .. } => variant.fields().is_unit(), + } +} + +fn check_enum(e: &tree::Enum, s: &Value, ctx: Context<'_>) { + let (tag, content) = extract_tag_and_content(e.repr(), s, ctx); + assert_all_unique(e.variants().iter().map(|f| f.name()), ctx); + + let variant = match e.variants().iter().find(|v| v.name() == tag) { + Some(variant) => variant, + None => panic!("could not find variant \"{}\"\n\n{}", tag, ctx), + }; + + if is_no_content_expected(e.repr(), variant) { + assert!(content.is_none(), "{}", ctx); + } else { + let content: &Value = match content.as_deref() { + Some(content) => content, + None => panic!("No content was found when expected\n\n{}", ctx), + }; + + check_fields( + variant.fields(), + content, + context!(ctx, content, "variant \"{}\"", variant.name()), + ) + } +} + +//----------// +// sequence // +//----------// + +fn check_sequence(seq: &tree::Sequence, s: &Value, ctx: Context<'_>) { + let a = value_as_array(s, ctx); + for (i, v) in a.iter().enumerate() { + check_reflection( + seq.element(), + v, + context!(ctx, v, "element {} of {}", i + 1, a.len()), + ); + } +} + +//---------// +// Helpers // +//---------// + +fn assert_all_unique(itr: I, ctx: Context<'_>) +where + I: IntoIterator, + I::Item: std::hash::Hash + Eq + std::fmt::Debug + Clone, +{ + let mut seen = HashSet::new(); + for i in itr { + if !seen.insert(i.clone()) { + panic!("item {:?} seen multiple times\n\n{}", i, ctx); + } + } +} + +fn value_as_str<'a>(v: &'a Value, ctx: Context<'_>) -> &'a str { + if let Value::String(s) = v { + s + } else { + panic!("expected value to be a string\n\n{}", ctx); + } +} + +fn value_as_map<'a>(v: &'a Value, ctx: Context<'_>) -> &'a serde_json::value::Map { + if let Value::Object(m) = v { + m + } else { + panic!("expected value to be an object\n\n{}", ctx); + } +} + +fn value_as_array<'a>(v: &'a Value, ctx: Context<'_>) -> &'a [Value] { + if let Value::Array(a) = v { + a + } else { + panic!("expected value to be an array\n\n{}", ctx); + } +} + +/////////// +// Tests // +/////////// + +#[test] +fn test_primitives() { + check_struct(()); + check_struct(0u8); + check_struct(0u16); + check_struct(0u32); + check_struct(0u64); + check_struct(0usize); + + check_struct(0i8); + check_struct(0i16); + check_struct(0i32); + check_struct(0i64); + check_struct(0isize); +} + +#[test] +fn simple_aggregate() { + #[derive(Serialize, Reflect, Clone)] + struct Simple { + a: usize, + b: usize, + r#type: (), + } + + let s = Simple { + a: 10, + b: 20, + r#type: (), + }; + check_struct(s.clone()); + + #[derive(Serialize, Reflect)] + #[serde(rename_all = "kebab-case")] + struct Nested { + some_long_field: usize, + #[serde(rename = "s")] + manual_rename: Simple, + } + + check_struct(Nested { + some_long_field: 4, + manual_rename: s, + }); +} + +#[test] +fn simple_tuple() { + // A simple new-type tuple. + #[derive(Serialize, Reflect, Clone)] + struct NewType(String); + + let newtype = NewType(String::from("hello")); + check_struct(newtype.clone()); + + // A tuple struct of length 2 + #[derive(Serialize, Reflect, Clone)] + struct Tuple2(usize, NewType); + + let tuple2 = Tuple2(100, newtype.clone()); + check_struct(tuple2.clone()); + + // A tuple struct of length 3 with a generic + #[derive(Serialize, Reflect)] + struct Tuple3(usize, NewType, T); + + let tuple3 = Tuple3(4, NewType(String::from("foo")), tuple2.clone()); + check_struct(tuple3); + + // A tuple struct of length 3 with a generic and zero-sized type. + #[derive(Serialize, Reflect)] + struct Unit; + + #[derive(Serialize, Reflect)] + struct Tuple4(usize, Unit, NewType, T); + + let tuple3 = Tuple4(4, Unit, newtype, tuple2); + check_struct(tuple3); +} + +#[test] +fn nested_newtype() { + #[derive(Serialize, Reflect)] + struct NewType0(usize); + + #[derive(Serialize, Reflect)] + struct NewType1(NewType0); + + #[derive(Serialize, Reflect)] + struct NewType2(NewType1); + + check_struct(NewType2(NewType1(NewType0(10)))); +} + +#[test] +fn enum_externally_tagged() { + #[derive(Serialize, Reflect)] + struct Aggregate { + foo: usize, + bar: usize, + } + + #[derive(Serialize, Reflect)] + #[serde(rename_all = "kebab-case")] + enum Enum { + Unit, + NewType(T), + Tuple0(), + Tuple2(usize, usize), + Tuple3(usize, usize, Aggregate), + #[serde(rename = "longer-unit-2")] + Unit2 {}, + + #[serde(rename_all = "kebab-case")] + Struct1 { + hello_world: usize, + }, + + #[serde(rename_all = "kebab-case")] + Struct2 { + hello_world: usize, + #[serde(rename = "baz")] + foo: usize, + }, + } + + check_enums( + &[ + Enum::Unit, + Enum::NewType(0usize), + Enum::Tuple0(), + Enum::Tuple2(1, 2), + Enum::Tuple3(1, 2, Aggregate { foo: 10, bar: 20 }), + Enum::Unit2 {}, + Enum::Struct1 { hello_world: 3 }, + Enum::Struct2 { + hello_world: 4, + foo: 5, + }, + ], + format_args!("externally tagged enums"), + ); +} + +#[test] +fn enum_internally_tagged() { + #[derive(Serialize, Reflect)] + struct NewTypePayload { + value: usize, + } + + #[derive(Serialize, Reflect)] + #[serde(rename_all = "kebab-case")] + #[serde(tag = "fizzle")] + enum Enum { + Unit, + EmptyStruct {}, + NewType(NewTypePayload), + Struct1 { a: usize }, + Struct2 { a: usize, b: usize }, + Struct3 { a: usize, b: usize, c: usize }, + } + + check_enums( + &[ + Enum::Unit, + Enum::EmptyStruct {}, + Enum::NewType(NewTypePayload { value: 10 }), + Enum::Struct1 { a: 0 }, + Enum::Struct2 { a: 0, b: 1 }, + Enum::Struct3 { a: 0, b: 1, c: 3 }, + ], + format_args!("internally tagged enums"), + ); +} + +#[test] +fn enum_adjacently_tagged() { + #[derive(Serialize, Reflect)] + struct Nested { + a: usize, + b: usize, + } + + #[derive(Serialize, Reflect)] + #[serde(rename_all = "kebab-case")] + #[serde(tag = "fizzle", content = "fuzzle")] + enum Enum { + Unit, + r#RawUnit, + Empty {}, + NewType(()), + NewType2(usize), + NewType3(Nested), + Tuple0(), + Tuple1(usize), + Tuple2(usize, ()), + Struct1 { a: usize }, + Struct2 { a: usize, b: usize }, + Struct3 { a: usize, b: usize, c: usize }, + } + + check_enums( + &[ + Enum::Unit, + Enum::r#RawUnit, + Enum::Empty {}, + Enum::NewType(()), + Enum::NewType2(10), + Enum::NewType3(Nested { a: 0, b: 1 }), + Enum::Tuple0(), + Enum::Tuple1(0), + Enum::Tuple2(0, ()), + Enum::Struct1 { a: 0 }, + Enum::Struct2 { a: 0, b: 1 }, + Enum::Struct3 { a: 0, b: 1, c: 3 }, + ], + format_args!("externally tagged enums"), + ); +} + +#[test] +fn nested_enum() { + #[derive(Serialize, Reflect)] + enum Nested { + Unit, + NewType(usize), + Struct { value: usize }, + } + + #[derive(Serialize, Reflect)] + struct Outer { + nested: Nested, + } + + check_struct(Outer { + nested: Nested::Unit, + }); + check_struct(Outer { + nested: Nested::NewType(10), + }); + check_struct(Outer { + nested: Nested::Struct { value: 20 }, + }); +} + +#[test] +fn sequence_of_newtypes() { + #[derive(Serialize, Reflect)] + struct NewType(usize); + + #[derive(Serialize, Reflect)] + struct Outer { + values: Vec, + } + + check_struct(Outer { + values: vec![NewType(10), NewType(20)], + }); +} diff --git a/diskann-benchmark-runner/src/reflect/test.rs b/diskann-benchmark-runner/src/reflect/test/unit.rs similarity index 85% rename from diskann-benchmark-runner/src/reflect/test.rs rename to diskann-benchmark-runner/src/reflect/test/unit.rs index 5e2fa99a9d..7536fab5e1 100644 --- a/diskann-benchmark-runner/src/reflect/test.rs +++ b/diskann-benchmark-runner/src/reflect/test/unit.rs @@ -3,11 +3,14 @@ * Licensed under the MIT license. */ -use super::{ +//! [`Reflect`] macro unit tests. + +use std::assert_matches; + +use crate::reflect::{ tree::{Fields, Type}, Reflect, Reflection, }; -use std::assert_matches; #[test] fn test_unit() { @@ -21,6 +24,10 @@ fn test_unit() { assert_eq!(r.ty().doc().unwrap(), "A unit struct."); assert_eq!(r.type_name().to_string(), "Unit"); assert!(!ty.has_body()); + + // Unit structs should be unit aggregates. + let agg = ty.as_aggregate().unwrap(); + assert!(agg.fields().is_unit()); } #[test] @@ -36,6 +43,10 @@ fn test_unit_rename() { assert_eq!(r.ty().doc().unwrap(), "A unit struct."); assert_eq!(r.type_name().to_string(), "something-else"); assert!(!ty.has_body()); + + // Unit structs should be unit aggregates. + let agg = ty.as_aggregate().unwrap(); + assert!(agg.fields().is_unit()); } #[test] @@ -65,6 +76,10 @@ fn test_unit_const_generic() { assert_eq!(ty.doc().unwrap(), "A unit struct."); assert_eq!(r.type_name().to_string(), "Unit<10>"); assert!(!ty.has_body()); + + // Unit structs should be unit aggregates. + let agg = ty.as_aggregate().unwrap(); + assert!(agg.fields().is_unit()); } #[test] @@ -107,6 +122,11 @@ fn test_empty_tuple_like() { assert_eq!(ty.doc().unwrap(), "An empty tuple-like struct."); assert_eq!(r.type_name().to_string(), "Empty"); assert!(!ty.has_body()); + + // Unit tuple structs should be aggregates with zero unnamed fields. + let agg = ty.as_aggregate().unwrap(); + let fields = agg.fields().as_unnamed().unwrap(); + assert!(fields.is_empty()); } #[test] @@ -121,6 +141,11 @@ fn test_empty_struct_like() { assert_eq!(ty.doc().unwrap(), "An empty struct."); assert_eq!(r.type_name().to_string(), "Empty"); assert!(!ty.has_body()); + + // Unit structs should be aggregates with zero unnamed fields. + let agg = ty.as_aggregate().unwrap(); + let fields = agg.fields().as_named().unwrap(); + assert!(fields.is_empty()); } #[test] @@ -201,7 +226,7 @@ fn test_struct_rename() { #[test] fn test_tuple1() { - /// A tuple with one field. + /// A tuple with one field. This should be a "newtype". #[expect(unused)] #[derive(Reflect)] struct Tuple1( @@ -211,15 +236,17 @@ fn test_tuple1() { let r = Reflection::new::(); let ty = r.ty(); - assert_eq!(ty.doc().unwrap(), "A tuple with one field."); + assert_eq!( + ty.doc().unwrap(), + "A tuple with one field. This should be a \"newtype\"." + ); assert_eq!(r.type_name().to_string(), "Tuple1"); assert!(ty.has_body()); - let f = ty.as_aggregate().unwrap().fields().as_unnamed().unwrap(); - assert_eq!(f.len(), 1); - assert_eq!(f[0].doc().unwrap(), "Field 0."); - assert_eq!(f[0].field().type_name().to_string(), "usize"); + let f = ty.as_aggregate().unwrap().fields().as_newtype().unwrap(); + assert_eq!(f.doc().unwrap(), "Field 0."); + assert_eq!(f.field().type_name().to_string(), "usize"); } #[test] @@ -289,20 +316,18 @@ fn test_enum_with_generics() { assert_eq!(variants[0].name(), "A"); assert!(variants[0].fields().has_body()); - let f = variants[0].fields().as_unnamed().unwrap(); - assert_eq!(f.len(), 1); - assert!(f[0].doc().is_none()); - assert_eq!(f[0].field().type_name().to_string(), "Vec"); + let f = variants[0].fields().as_newtype().unwrap(); + assert!(f.doc().is_none()); + assert_eq!(f.field().type_name().to_string(), "Vec"); // Variant 1 assert_eq!(variants[1].doc().unwrap(), "It's a bee!"); assert_eq!(variants[1].name(), "B"); assert!(variants[1].fields().has_body()); - let f = variants[1].fields().as_unnamed().unwrap(); - assert_eq!(f.len(), 1); - assert_eq!(f[0].doc().unwrap(), "Buzz buzz"); - assert_eq!(f[0].field().type_name().to_string(), "string"); + let f = variants[1].fields().as_newtype().unwrap(); + assert_eq!(f.doc().unwrap(), "Buzz buzz"); + assert_eq!(f.field().type_name().to_string(), "string"); } #[test] diff --git a/diskann-benchmark-runner/src/reflect/tree.rs b/diskann-benchmark-runner/src/reflect/tree.rs index 4dcf85bacb..c0f62fd877 100644 --- a/diskann-benchmark-runner/src/reflect/tree.rs +++ b/diskann-benchmark-runner/src/reflect/tree.rs @@ -16,8 +16,8 @@ pub enum Type { } impl Type { - pub fn primitive(doc: Option) -> Self { - Self::from(Primitive::new(doc)) + pub fn primitive(kind: PrimitiveKind, doc: Option) -> Self { + Self::from(Primitive::new(kind, doc)) } pub fn aggregate(fields: Fields, doc: Option) -> Self { @@ -57,6 +57,7 @@ impl Type { } } + #[cfg(test)] pub(super) fn as_aggregate(&self) -> Option<&Aggregate> { if let Self::Aggregate(aggregate) = self { Some(aggregate) @@ -65,6 +66,7 @@ impl Type { } } + #[cfg(test)] pub(super) fn as_enum(&self) -> Option<&Enum> { if let Self::Enum(enum_) = self { Some(enum_) @@ -102,17 +104,30 @@ impl From for Type { // Primitive // //-----------// +#[derive(Debug, Clone, Copy)] +pub enum PrimitiveKind { + Null, + Boolean, + Number, + String, +} + #[derive(Debug)] pub struct Primitive { + kind: PrimitiveKind, doc: Option, } impl Primitive { - pub fn new(doc: Option) -> Self { - Self { doc } + pub fn new(kind: PrimitiveKind, doc: Option) -> Self { + Self { kind, doc } + } + + pub(super) fn kind(&self) -> PrimitiveKind { + self.kind } - fn doc(&self) -> Option<&str> { + pub(super) fn doc(&self) -> Option<&str> { self.doc.as_deref() } } @@ -149,6 +164,7 @@ impl Aggregate { pub enum Fields { Named(Vec), Unnamed(Vec), + NewType(UnnamedField), Unit, } @@ -157,6 +173,7 @@ impl Fields { match self { Self::Named(fields) => !fields.is_empty(), Self::Unnamed(fields) => !fields.is_empty(), + Self::NewType(_) => true, Self::Unit => false, } } @@ -169,6 +186,11 @@ impl Fields { Self::Unnamed(itr.into_iter().collect()) } + pub fn newtype(field: UnnamedField) -> Self { + Self::NewType(field) + } + + #[cfg(test)] pub(super) fn as_named(&self) -> Option<&[NamedField]> { if let Self::Named(fields) = self { Some(fields) @@ -177,6 +199,7 @@ impl Fields { } } + #[cfg(test)] pub(super) fn as_unnamed(&self) -> Option<&[UnnamedField]> { if let Self::Unnamed(fields) = self { Some(fields) @@ -185,6 +208,15 @@ impl Fields { } } + #[cfg(test)] + pub(super) fn as_newtype(&self) -> Option<&UnnamedField> { + if let Self::NewType(field) = self { + Some(field) + } else { + None + } + } + pub(super) fn is_unit(&self) -> bool { matches!(self, Self::Unit) } diff --git a/diskann-benchmark-runner/src/test/typed.rs b/diskann-benchmark-runner/src/test/typed.rs index 5bd900f9db..c863f3be44 100644 --- a/diskann-benchmark-runner/src/test/typed.rs +++ b/diskann-benchmark-runner/src/test/typed.rs @@ -24,6 +24,7 @@ pub(crate) struct TypeInput { error_when_checked: bool, } +/// This is a test input for testing corner cases in the benchmark runner. #[derive(Serialize, Deserialize, Reflect)] pub(crate) struct TypeInputRaw { data_type: DataType, diff --git a/diskann-benchmark-runner/temp/main.rs b/diskann-benchmark-runner/temp/main.rs index cbdce829de..d0aea2a292 100644 --- a/diskann-benchmark-runner/temp/main.rs +++ b/diskann-benchmark-runner/temp/main.rs @@ -5,7 +5,7 @@ //! Development CLI for exercising the benchmark runner with its test registry. -use diskann_benchmark_runner::{reflect, Reflection}; +use diskann_benchmark_runner::{reflect, Reflect, Reflection}; fn main() -> anyhow::Result<()> { // println!("{}", Reflection::new::().render()); From 9e81f18ad45e01d1eeda2c574a54212d8a84f652 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Mon, 28 Sep 2026 15:03:12 -0700 Subject: [PATCH 12/21] Checkpoint. --- .../src/reflect/test/serde.rs | 93 +++++++++++-------- 1 file changed, 53 insertions(+), 40 deletions(-) diff --git a/diskann-benchmark-runner/src/reflect/test/serde.rs b/diskann-benchmark-runner/src/reflect/test/serde.rs index b5313d2fd4..cd1617773a 100644 --- a/diskann-benchmark-runner/src/reflect/test/serde.rs +++ b/diskann-benchmark-runner/src/reflect/test/serde.rs @@ -9,8 +9,8 @@ //! care to support such fanciness in our benchmark inputs and outputs, especially since we //! need example JSONs to be round-trippable back to their serialized representations. //! -//! As such, the main entry points ([`check_serde`] and [`check_enums`]) just take [`Serialize`] -//! bounds. +//! As such, the main entry points ([`check_struct`] and [`check_enums`]) just take [`Serialize`] +//! bounds. We do not check round-trippability here. That's the user's problem. use std::{assert_matches, borrow::Cow}; @@ -295,8 +295,8 @@ fn extract_tag_and_content<'a>( ctx ); - let kv: Vec<_> = m.iter().collect(); - (&kv[0].0, Some(Cow::Borrowed(&kv[0].1))) + let kv = m.iter().next().unwrap(); + (&kv.0, Some(Cow::Borrowed(&kv.1))) } _ => panic!("invalid representation\n\n{}", ctx), }, @@ -320,11 +320,7 @@ fn extract_tag_and_content<'a>( // Delete the tag field to reuse the rest of the checking infrastructure. let mut map = map.clone(); map.remove(*tag); - if map.is_empty() { - (t, None) - } else { - (t, Some(Cow::Owned(Value::Object(map)))) - } + (t, Some(Cow::Owned(Value::Object(map)))) } // For adjacent tagging, the "content" field is ommitted when the corresponding @@ -352,22 +348,6 @@ fn extract_tag_and_content<'a>( } } -// See the inline notes in `extract_tag_and_content`. -fn is_no_content_expected(repr: &tree::EnumRepr, variant: &tree::Variant) -> bool { - match repr { - tree::EnumRepr::External => variant.fields().is_unit(), - tree::EnumRepr::Internal { .. } => match variant.fields() { - tree::Fields::Named(named_fields) => named_fields.is_empty(), - tree::Fields::Unnamed(_) => unreachable!("unimplemented by serde"), - // Only allowed for new-type wrappers that themselves are structs. As sucn, - // we always expect a payload. - tree::Fields::NewType(_) => false, - tree::Fields::Unit => true, - }, - tree::EnumRepr::Adjacent { .. } => variant.fields().is_unit(), - } -} - fn check_enum(e: &tree::Enum, s: &Value, ctx: Context<'_>) { let (tag, content) = extract_tag_and_content(e.repr(), s, ctx); assert_all_unique(e.variants().iter().map(|f| f.name()), ctx); @@ -377,20 +357,38 @@ fn check_enum(e: &tree::Enum, s: &Value, ctx: Context<'_>) { None => panic!("could not find variant \"{}\"\n\n{}", tag, ctx), }; - if is_no_content_expected(e.repr(), variant) { - assert!(content.is_none(), "{}", ctx); - } else { - let content: &Value = match content.as_deref() { - Some(content) => content, - None => panic!("No content was found when expected\n\n{}", ctx), - }; - - check_fields( - variant.fields(), - content, - context!(ctx, content, "variant \"{}\"", variant.name()), - ) - } + let next: &Value = match (e.repr(), &content) { + (tree::EnumRepr::External, None) => { + assert!( + variant.fields().is_unit(), + "content may only be excluded for unit variants\n\n{}", ctx + ); + return; + }, + (tree::EnumRepr::External, Some(c)) => &c, + (tree::EnumRepr::Internal { .. }, None) => unreachable!("internal always returns content"), + (tree::EnumRepr::Internal { .. }, Some(c)) => { + if variant.fields().is_unit() { + let map = value_as_map(c, ctx); + assert!(map.is_empty(), "unit enums should have no remaining values"); + return; + } else { + c + } + }, + (tree::EnumRepr::Adjacent { .. }, None) => { + assert!(variant.fields().is_unit(), + "content may only be excluded for unit variants\n\n{}", ctx); + return; + } + (tree::EnumRepr::Adjacent { .. }, Some(c)) => &c, + }; + + check_fields( + variant.fields(), + next, + context!(ctx, next, "variant \"{}\"", variant.name()), + ) } //----------// @@ -456,6 +454,9 @@ fn value_as_array<'a>(v: &'a Value, ctx: Context<'_>) -> &'a [Value] { #[test] fn test_primitives() { check_struct(()); + + check_struct(false); + check_struct(0u8); check_struct(0u16); check_struct(0u32); @@ -467,6 +468,13 @@ fn test_primitives() { check_struct(0i32); check_struct(0i64); check_struct(0isize); + + check_struct(0f32); + check_struct(0f64); + + check_struct(std::marker::PhantomData::); + + check_struct(String::from("hello")); } #[test] @@ -604,6 +612,9 @@ fn enum_internally_tagged() { value: usize, } + #[derive(Serialize, Reflect)] + struct NewTypeEmpty {} + #[derive(Serialize, Reflect)] #[serde(rename_all = "kebab-case")] #[serde(tag = "fizzle")] @@ -611,6 +622,7 @@ fn enum_internally_tagged() { Unit, EmptyStruct {}, NewType(NewTypePayload), + NewTypeEmpty(NewTypeEmpty), Struct1 { a: usize }, Struct2 { a: usize, b: usize }, Struct3 { a: usize, b: usize, c: usize }, @@ -621,6 +633,7 @@ fn enum_internally_tagged() { Enum::Unit, Enum::EmptyStruct {}, Enum::NewType(NewTypePayload { value: 10 }), + Enum::NewTypeEmpty(NewTypeEmpty {}), Enum::Struct1 { a: 0 }, Enum::Struct2 { a: 0, b: 1 }, Enum::Struct3 { a: 0, b: 1, c: 3 }, @@ -670,7 +683,7 @@ fn enum_adjacently_tagged() { Enum::Struct2 { a: 0, b: 1 }, Enum::Struct3 { a: 0, b: 1, c: 3 }, ], - format_args!("externally tagged enums"), + format_args!("adjacently tagged enums"), ); } From 288505f22bd9383b6e9be930ee427308ebac56b2 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Mon, 28 Sep 2026 16:47:43 -0700 Subject: [PATCH 13/21] More progress? --- diskann-benchmark-runner/src/app.rs | 50 +-- diskann-benchmark-runner/src/reflect/mod.rs | 21 +- .../src/reflect/render.rs | 322 +++++++++++++++++- .../src/reflect/test/serde.rs | 83 ++++- diskann-benchmark-runner/src/reflect/tree.rs | 51 +++ diskann-benchmark-runner/src/ux.rs | 44 ++- .../tests/rendered_reflections.txt | 164 +++++++++ 7 files changed, 650 insertions(+), 85 deletions(-) create mode 100644 diskann-benchmark-runner/tests/rendered_reflections.txt diff --git a/diskann-benchmark-runner/src/app.rs b/diskann-benchmark-runner/src/app.rs index 55dd3e75b3..d355d988fd 100644 --- a/diskann-benchmark-runner/src/app.rs +++ b/diskann-benchmark-runner/src/app.rs @@ -545,8 +545,6 @@ mod tests { use crate::{registry, test::TestConfig, ux}; - const ENV: &str = "DISKANN_TEST"; - // Expected I/O files. const STDIN: &str = "stdin.txt"; const STDOUT: &str = "stdout.txt"; @@ -564,42 +562,6 @@ mod tests { const ALL_GENERATED_OUTPUTS: [&str; 2] = [OUTPUT_FILE, CHECK_OUTPUT_FILE]; - // Read the entire contents of a file to a string. - fn read_to_string>(path: P, ctx: &str) -> String { - match std::fs::read_to_string(path.as_ref()) { - Ok(s) => ux::normalize(s), - Err(err) => panic!( - "failed to read {} {:?} with error: {}", - ctx, - path.as_ref(), - err - ), - } - } - - // Check if `DISKANN_TEST=overwrite` is configured. Return `true` if so - otherwise - // return `false`. - // - // If `DISKANN_TEST` is set but its value is not `overwrite` - panic. - fn overwrite() -> bool { - match std::env::var(ENV) { - Ok(v) => { - if v == "overwrite" { - true - } else { - panic!( - "Unknown value for {}: \"{}\". Expected \"overwrite\"", - ENV, v - ); - } - } - Err(std::env::VarError::NotPresent) => false, - Err(std::env::VarError::NotUnicode(_)) => { - panic!("Value for {} is not unicode", ENV); - } - } - } - // Test Runner struct Test { dir: PathBuf, @@ -610,7 +572,7 @@ mod tests { fn new(dir: &Path) -> Self { Self { dir: dir.into(), - overwrite: overwrite(), + overwrite: ux::overwrite(), } } @@ -618,7 +580,7 @@ mod tests { let path = self.dir.join(STDIN); // Read the standard input file to a string. - let stdin = read_to_string(&path, "standard input"); + let stdin = ux::read_to_string(&path, "standard input"); let output: Vec = stdin .lines() @@ -727,7 +689,7 @@ mod tests { if self.overwrite { std::fs::write(output, stdout).unwrap(); } else { - let expected = read_to_string(&output, "expected standard output"); + let expected = ux::read_to_string(&output, "expected standard output"); if stdout != expected { panic!("Got:\n--\n{}\n--\nExpected:\n--\n{}\n--", stdout, expected); } @@ -769,9 +731,9 @@ mod tests { } else { match (was_generated, is_expected) { (true, true) => { - let output_contents = read_to_string(generated_path, "generated"); + let output_contents = ux::read_to_string(generated_path, "generated"); - let expected_contents = read_to_string(expected_path, "expected"); + let expected_contents = ux::read_to_string(expected_path, "expected"); if output_contents != expected_contents { panic!( @@ -781,7 +743,7 @@ mod tests { } } (true, false) => { - let output_contents = read_to_string(generated_path, "generated"); + let output_contents = ux::read_to_string(generated_path, "generated"); panic!( "{} was generated when none was expected. Contents:\n\n{}", diff --git a/diskann-benchmark-runner/src/reflect/mod.rs b/diskann-benchmark-runner/src/reflect/mod.rs index 93c1ab16d3..2b7bc836da 100644 --- a/diskann-benchmark-runner/src/reflect/mod.rs +++ b/diskann-benchmark-runner/src/reflect/mod.rs @@ -93,7 +93,7 @@ pub struct Render(Reflection); impl std::fmt::Display for Render { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - let mut r = render::Renderer::new(f, 3); + let mut r = render::Renderer::new(f, 2); r.render_subject(self.0) } } @@ -130,6 +130,7 @@ where .iter() .for_each(|variant| push_fields(variant.fields())), Type::Sequence(seq) => push(seq.element()), + Type::Optional(opt) => push(opt.value()), } } @@ -210,23 +211,7 @@ where T: Reflect, { fn ty() -> Type { - Type::enum_( - // TODO: Untagged - tree::EnumRepr::External, - [ - tree::Variant::new( - "", - tree::Fields::Unit, - Some("Use `null` to indicate that this value does not exist".into()), - ), - tree::Variant::new( - "", - tree::Fields::unnamed([tree::UnnamedField::new::(None)]), - Some("Presence indicates the value is present".into()), - ), - ], - Some("An optional configuration".into()), - ) + Type::optional::(Some("An optional type".into())) } fn format_type_name(f: &mut dyn Write) -> fmt::Result { diff --git a/diskann-benchmark-runner/src/reflect/render.rs b/diskann-benchmark-runner/src/reflect/render.rs index e7a7d4e8b7..e6d2bdb876 100644 --- a/diskann-benchmark-runner/src/reflect/render.rs +++ b/diskann-benchmark-runner/src/reflect/render.rs @@ -8,7 +8,10 @@ use std::fmt::{self, Write}; use crate::utils::fmt::Quote; use super::{ - tree::{Aggregate, Enum, EnumRepr, Fields, NamedField, Sequence, Type, UnnamedField, Variant}, + tree::{ + Aggregate, Enum, EnumRepr, Fields, NamedField, Optional, Sequence, Type, UnnamedField, + Variant, + }, Reflection, }; @@ -160,9 +163,10 @@ impl<'a> Renderer<'a> { fn render_body(&mut self, tagged: &Tagged) -> fmt::Result { match tagged.ty() { Type::Primitive(_) => Ok(()), - Type::Aggregate(aggregate) => self.render_aggregate(&aggregate), + Type::Aggregate(aggregate) => self.render_aggregate(aggregate), Type::Enum(enum_) => self.render_enum(enum_), - Type::Sequence(sequence) => self.render_sequence(&sequence), + Type::Sequence(sequence) => self.render_sequence(sequence), + Type::Optional(opt) => self.render_optional(opt), } } @@ -212,6 +216,11 @@ impl<'a> Renderer<'a> { ) } + fn render_optional(&mut self, op: &Optional) -> fmt::Result { + self.line("May be `null`.")?; + self.next(|r| r.render_body(&Tagged::new(op.value()))) + } + //--------// // Fields // //--------// @@ -233,10 +242,10 @@ impl<'a> Renderer<'a> { if i != 0 { self.blank()?; } - self.render_unnamed_field(i, field)?; + self.render_unnamed_field(Some(i), field)?; } } - Fields::NewType(_) => todo!(), + Fields::NewType(newtype) => self.render_unnamed_field(None, newtype)?, Fields::Unit => {} } @@ -264,15 +273,17 @@ impl<'a> Renderer<'a> { Ok(()) } - fn render_unnamed_field(&mut self, index: usize, field: &UnnamedField) -> fmt::Result { + fn render_unnamed_field(&mut self, index: Option, field: &UnnamedField) -> fmt::Result { let tagged = Tagged::new(field.field()); let will_render_body = self.will_render(tagged.ty()); - self.line(format_args!( - "{}: {}", - index, - tagged.reflection().type_name() - ))?; + if let Some(index) = index { + self.line(format_args!( + "{}: {}", + index, + tagged.reflection().type_name() + ))?; + } let rendered_doc = self.indent(|r| r.render_doc(field.doc()))?; if rendered_doc && will_render_body { @@ -303,3 +314,292 @@ impl<'a> Renderer<'a> { }) } } + +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + use std::{ + fs::File, + io::{self, BufRead, BufReader, Write}, + path::{Path, PathBuf}, + }; + + use serde::{Deserialize, Serialize}; + + use crate::{ux, Reflect}; + + // For these tests, we use a variation of baseline tests where all the expected results + // are put into a single file, mainly to keep from generating a bunch of files for the + // relatively small tests. + + fn baseline_path() -> PathBuf { + format!( + "{}/tests/rendered_reflections.txt", + env!("CARGO_MANIFEST_DIR") + ) + .into() + } + + fn overwrite_hint() -> &'static str { + "Baselines can be regenerated by running tests with `DISKANN_TEST=overwrite`" + } + + #[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] + struct Content { + name: String, + description: String, + } + + const CASE_SEPARATOR: &'static str = "========"; + const OUTPUT_SEPARATOR: &'static str = "--------"; + + #[derive(Debug)] + struct Case { + content: Content, + r: Reflection, + } + + impl Case { + fn new(name: &str, description: &str) -> Self + where + T: Reflect, + { + Self { + content: Content { + name: name.into(), + description: description.into(), + }, + r: Reflection::new::(), + } + } + } + + #[derive(Debug)] + struct Rendered { + content: Content, + rendered: String, + } + + #[derive(Default, Debug, PartialEq)] + enum State { + #[default] + Content, + Rendered, + Done, + } + + impl State { + fn to_rendered(&mut self) { + assert_eq!(*self, Self::Content); + *self = Self::Rendered; + } + + fn to_done(&mut self) { + assert_eq!(*self, Self::Rendered); + *self = Self::Done; + } + + fn assert_is_done(self) { + assert_eq!(self, Self::Done); + } + } + + impl Rendered { + fn parse_if_present(lines: &mut std::io::Lines) -> Option + where + B: std::io::BufRead, + { + match lines.next() { + Some(ln) => assert_eq!(ln.unwrap(), CASE_SEPARATOR), + None => return None, + }; + + let mut content = String::new(); + let mut rendered = String::new(); + + let mut state = State::default(); + + while let Some(ln) = lines.next() { + let ln: &str = &ln.unwrap(); + match ln { + CASE_SEPARATOR => { + state.to_done(); + break; + } + OUTPUT_SEPARATOR => { + state.to_rendered(); + } + ln => match state { + State::Content => content.push_str(ln), + State::Rendered => { + rendered.push('\n'); + rendered.push_str(ln); + } + State::Done => panic!("invalid state"), + }, + } + } + + state.assert_is_done(); + + Some(Self { + content: serde_json::from_str(&content).unwrap(), + rendered: rendered.trim().to_string(), + }) + } + } + + fn parse_baselines(path: &Path) -> Vec { + let file = match File::open(path) { + Ok(file) => file, + Err(err) => panic!( + "Opening path \"{}\" failed with {}. {}", + path.display(), + err, + overwrite_hint() + ), + }; + let mut lines = BufReader::new(file).lines(); + let mut baselines = Vec::new(); + while let Some(baseline) = Rendered::parse_if_present(&mut lines) { + baselines.push(baseline); + } + + baselines + } + + fn write_baselines(rendered: &[Rendered], path: &Path) { + let mut io = File::create(path).unwrap(); + + for r in rendered.iter() { + let Rendered { content, rendered } = r; + writeln!(io, "{}", CASE_SEPARATOR).unwrap(); + writeln!(io, "{}", serde_json::to_string_pretty(content).unwrap()).unwrap(); + writeln!(io, "{}", OUTPUT_SEPARATOR).unwrap(); + writeln!(io, "{}", ux::normalize(rendered.clone())).unwrap(); + writeln!(io, "{}", CASE_SEPARATOR).unwrap(); + } + } + + fn run_tests_inner(cases: &[Case], path: &Path, overwrite: bool) { + // Generate the current set of baseline. + let current: Vec<_> = cases + .iter() + .map(|case| { + let rendered = ux::normalize(case.r.render().to_string()); + Rendered { + content: case.content.clone(), + rendered, + } + }) + .collect(); + + if overwrite { + write_baselines(¤t, path); + } else { + let expected = parse_baselines(path); + assert_eq!( + current.len(), + expected.len(), + "Number of baseline cases differs. {}", + overwrite_hint(), + ); + + for (current, expected) in std::iter::zip(current.iter(), expected.iter()) { + assert_eq!( + current.content, + expected.content, + "Baseline headers differ. {}", + overwrite_hint(), + ); + + if current.rendered != expected.rendered { + panic!( + "Difference for case name {}\n\nEXPECTED\n\n{}\n\nGOT\n\n{}\n\n{}", + expected.content.name, + expected.rendered, + current.rendered, + overwrite_hint() + ); + } + } + } + } + + /// Select the data type please. + #[derive(Reflect)] + #[serde(rename_all = "kebab-case")] + #[reflect(prefix = "render::")] + enum SimpleDataType { + Float32, + Float16, + } + + /// Select the data type please. + #[derive(Reflect)] + #[serde(rename_all = "kebab-case")] + #[reflect(prefix = "render::")] + enum AnnotatedDataType { + /// Use high-precision. + Float32, + /// Use lower precision. + Float16, + Int8, + } + + #[derive(Reflect)] + #[serde(rename_all = "snake_case")] + #[reflect(prefix = "render::")] + enum Source { + /// Build from scratch. + Build { + /// The type of the input data. + data_type: SimpleDataType, + + /// Input data in the `.bin` binary format. + file: String, + + /// Output file. + /// + /// If provided, saved data will go here. + output: Option, + }, + /// Run from a previously generated output. + FromPrevious { + data_type: AnnotatedDataType, + /// The previously generated output. + file: String, + }, + } + + /// A new-type wrapper around `Source`. + #[derive(Reflect)] + struct SourceWrapper(Source); + + /// A top level config. + #[derive(Reflect)] + struct Config { + source: SourceWrapper, + /// This does one thing. + param1: String, + /// This does another. + param2: Vec, + } + + #[test] + fn run_tests() { + let cases = [ + Case::new::("usize", "a simple test"), + Case::new::("source", "a configurable source enum"), + Case::new::("source wrapper", "render a newtype wrapper"), + Case::new::("config", "a sample struct config."), + ]; + + run_tests_inner(&cases, &baseline_path(), ux::overwrite()); + } +} diff --git a/diskann-benchmark-runner/src/reflect/test/serde.rs b/diskann-benchmark-runner/src/reflect/test/serde.rs index cd1617773a..917b95a4d0 100644 --- a/diskann-benchmark-runner/src/reflect/test/serde.rs +++ b/diskann-benchmark-runner/src/reflect/test/serde.rs @@ -64,9 +64,9 @@ where /// A context for displaying where we are in the type tree. /// /// Contains the top level JSON we're working, the current JSON, and a stack of -/// [`std::fmt::Argument`]s that describe the sequence of operations that led us into a mess. +/// [`std::fmt::Arguments`] that describe the sequence of operations that led us into a mess. /// -/// Use the [`context`] macro for creating nexted contexts. +/// Use the [`context`] macro for creating nested contexts. #[derive(Debug, Clone, Copy)] struct Context<'a> { top: &'a Value, @@ -135,6 +135,7 @@ fn check_type(ty: &Type, s: &Value, ctx: Context<'_>) { Type::Aggregate(a) => check_fields(a.fields(), s, context!(ctx, s, "aggregate")), Type::Enum(e) => check_enum(e, s, context!(ctx, s, "enum")), Type::Sequence(seq) => check_sequence(seq, s, context!(ctx, s, "sequence")), + Type::Optional(opt) => check_optional(opt, s, context!(ctx, s, "optional")), } } @@ -266,7 +267,7 @@ fn check_enum_variants(e: &tree::Enum, s: &[Value], ctx: std::fmt::Arguments<'_> /// "a": 10, /// } /// ``` -/// Adjaceny (say "tag = "mytag", "content" = "mycontent") looks like +/// Adjacency (say "tag = "mytag", "content" = "mycontent") looks like /// ```json /// { /// "mytag": "baz", @@ -301,7 +302,7 @@ fn extract_tag_and_content<'a>( _ => panic!("invalid representation\n\n{}", ctx), }, - // For internally tagged enums - we remove the tag after revrieval. + // For internally tagged enums - we remove the tag after retrieval. // // If the remaining dictionary is empty, then we change it to `None`. // This is technically a little ambiguous between unit variants and empty struct @@ -323,7 +324,7 @@ fn extract_tag_and_content<'a>( (t, Some(Cow::Owned(Value::Object(map)))) } - // For adjacent tagging, the "content" field is ommitted when the corresponding + // For adjacent tagging, the "content" field is omitted when the corresponding // variant is a unit variant. tree::EnumRepr::Adjacent { tag, content } => { let outer = value_as_map(s, ctx); @@ -361,10 +362,11 @@ fn check_enum(e: &tree::Enum, s: &Value, ctx: Context<'_>) { (tree::EnumRepr::External, None) => { assert!( variant.fields().is_unit(), - "content may only be excluded for unit variants\n\n{}", ctx + "content may only be excluded for unit variants\n\n{}", + ctx ); return; - }, + } (tree::EnumRepr::External, Some(c)) => &c, (tree::EnumRepr::Internal { .. }, None) => unreachable!("internal always returns content"), (tree::EnumRepr::Internal { .. }, Some(c)) => { @@ -375,10 +377,13 @@ fn check_enum(e: &tree::Enum, s: &Value, ctx: Context<'_>) { } else { c } - }, + } (tree::EnumRepr::Adjacent { .. }, None) => { - assert!(variant.fields().is_unit(), - "content may only be excluded for unit variants\n\n{}", ctx); + assert!( + variant.fields().is_unit(), + "content may only be excluded for unit variants\n\n{}", + ctx + ); return; } (tree::EnumRepr::Adjacent { .. }, Some(c)) => &c, @@ -406,6 +411,16 @@ fn check_sequence(seq: &tree::Sequence, s: &Value, ctx: Context<'_>) { } } +//----------// +// optional // +//----------// + +fn check_optional(opt: &tree::Optional, s: &Value, ctx: Context<'_>) { + if !s.is_null() { + check_reflection(opt.value(), s, context!(ctx, s, "present optional")) + } +} + //---------// // Helpers // //---------// @@ -509,7 +524,7 @@ fn simple_aggregate() { #[test] fn simple_tuple() { - // A simple new-type tuple. + // A simple newtype tuple. #[derive(Serialize, Reflect, Clone)] struct NewType(String); @@ -726,3 +741,49 @@ fn sequence_of_newtypes() { values: vec![NewType(10), NewType(20)], }); } + +#[test] +fn test_optionals() { + #[derive(Serialize, Reflect)] + #[serde(tag = "tag", content = "content", rename_all = "snake_case")] + enum CasesAdjacent { + NewType(Option), + Struct { val: Option }, + }; + + #[derive(Serialize, Reflect)] + struct NewType(Option); + + #[derive(Serialize, Reflect)] + struct More { + a: usize, + b: usize, + } + + #[derive(Serialize, Reflect)] + struct Struct { + e0: CasesAdjacent, + e1: CasesAdjacent, + more: Option, + opt: Option, + newtype: NewType, + } + + check_struct(Struct { + e0: CasesAdjacent::NewType(None), + e1: CasesAdjacent::Struct { val: None }, + more: None, + opt: None, + newtype: NewType(None), + }); + + let some = Some(1); + + check_struct(Struct { + e0: CasesAdjacent::NewType(some), + e1: CasesAdjacent::Struct { val: some }, + more: Some(More { a: 0, b: 1 }), + opt: some, + newtype: NewType(some), + }); +} diff --git a/diskann-benchmark-runner/src/reflect/tree.rs b/diskann-benchmark-runner/src/reflect/tree.rs index c0f62fd877..d92df1ae25 100644 --- a/diskann-benchmark-runner/src/reflect/tree.rs +++ b/diskann-benchmark-runner/src/reflect/tree.rs @@ -13,6 +13,7 @@ pub enum Type { Aggregate(Aggregate), Enum(Enum), Sequence(Sequence), + Optional(Optional), } impl Type { @@ -39,12 +40,22 @@ impl Type { Self::from(Sequence::new::(doc)) } + /// Keep the constructor private since we don't want users constructing the very special + /// `Optional` type for their own types. + pub(crate) fn optional(doc: Option) -> Self + where + T: Reflect, + { + Self::from(Optional::new::(doc)) + } + pub fn doc(&self) -> Option<&str> { match self { Self::Primitive(p) => p.doc(), Self::Aggregate(a) => a.doc(), Self::Enum(e) => e.doc(), Self::Sequence(s) => s.doc(), + Self::Optional(o) => o.doc(), } } @@ -54,6 +65,7 @@ impl Type { Self::Aggregate(a) => a.has_body(), Self::Enum(e) => e.has_body(), Self::Sequence(_) => true, + Self::Optional(_) => true, } } @@ -100,6 +112,12 @@ impl From for Type { } } +impl From for Type { + fn from(s: Optional) -> Self { + Self::Optional(s) + } +} + //-----------// // Primitive // //-----------// @@ -322,6 +340,7 @@ impl Enum { } #[derive(Debug)] +#[non_exhaustive] pub enum EnumRepr { /// Enums are tagged as the key in a collection. External, @@ -416,3 +435,35 @@ impl Sequence { self.doc.as_deref() } } + +//----------// +// Optional // +//----------// + +#[derive(Debug)] +pub struct Optional { + element: Reflection, + doc: Option, +} + +impl Optional { + /// Keep the constructor private since we don't want users constructing the very special + /// `Optional` type for their own types. + pub(super) fn new(doc: Option) -> Self + where + T: Reflect, + { + Self { + element: Reflection::new::(), + doc, + } + } + + pub(super) fn value(&self) -> Reflection { + self.element + } + + pub(super) fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } +} diff --git a/diskann-benchmark-runner/src/ux.rs b/diskann-benchmark-runner/src/ux.rs index 98e0a6b42b..33693bf35a 100644 --- a/diskann-benchmark-runner/src/ux.rs +++ b/diskann-benchmark-runner/src/ux.rs @@ -3,7 +3,7 @@ * Licensed under the MIT license. */ -use std::sync::LazyLock; +use std::{path::Path, sync::LazyLock}; /// Normalize a string for comparison. /// @@ -90,3 +90,45 @@ pub fn strip_backtrace(s: String) -> String { lines.join("\n") } + +//--------------// +// Crate Shared // +//--------------// + +// Read the entire contents of a file to a string. +pub(crate) fn read_to_string>(path: P, ctx: &str) -> String { + match std::fs::read_to_string(path.as_ref()) { + Ok(s) => normalize(s), + Err(err) => panic!( + "failed to read {} {:?} with error: {}", + ctx, + path.as_ref(), + err + ), + } +} + +const ENV: &str = "DISKANN_TEST"; + +// Check if `DISKANN_TEST=overwrite` is configured. Return `true` if so - otherwise +// return `false`. +// +// If `DISKANN_TEST` is set but its value is not `overwrite` - panic. +pub(crate) fn overwrite() -> bool { + match std::env::var(ENV) { + Ok(v) => { + if v == "overwrite" { + true + } else { + panic!( + "Unknown value for {}: \"{}\". Expected \"overwrite\"", + ENV, v + ); + } + } + Err(std::env::VarError::NotPresent) => false, + Err(std::env::VarError::NotUnicode(_)) => { + panic!("Value for {} is not unicode", ENV); + } + } +} diff --git a/diskann-benchmark-runner/tests/rendered_reflections.txt b/diskann-benchmark-runner/tests/rendered_reflections.txt new file mode 100644 index 0000000000..a2770a7e17 --- /dev/null +++ b/diskann-benchmark-runner/tests/rendered_reflections.txt @@ -0,0 +1,164 @@ +======== +{ + "name": "usize", + "description": "a simple test" +} +-------- +usize + A system dependent unsigned integer +======== +======== +{ + "name": "source", + "description": "a configurable source enum" +} +-------- +render::Source + Representation: externally tagged + + Options: + "build" + Build from scratch. + + "data_type": render::SimpleDataType + The type of the input data. + + Representation: externally tagged + + Options: + "float32" + "float16" + + "file": string + Input data in the `.bin` binary format. + + "output": Option + Output file. + + If provided, saved data will go here. + + May be `null`. + + "from_previous" + Run from a previously generated output. + + "data_type": render::AnnotatedDataType + Representation: externally tagged + + Options: + "float32" + Use high-precision. + + "float16" + Use lower precision. + + "int8" + + "file": string + The previously generated output. +======== +======== +{ + "name": "source wrapper", + "description": "render a newtype wrapper" +} +-------- +SourceWrapper + A new-type wrapper around `Source`. + + Representation: externally tagged + + Options: + "build" + Build from scratch. + + "data_type": render::SimpleDataType + The type of the input data. + + Representation: externally tagged + + Options: + "float32" + "float16" + + "file": string + Input data in the `.bin` binary format. + + "output": Option + Output file. + + If provided, saved data will go here. + + May be `null`. + + "from_previous" + Run from a previously generated output. + + "data_type": render::AnnotatedDataType + Representation: externally tagged + + Options: + "float32" + Use high-precision. + + "float16" + Use lower precision. + + "int8" + + "file": string + The previously generated output. +======== +======== +{ + "name": "config", + "description": "a sample struct config." +} +-------- +Config + A top level config. + + "source": SourceWrapper + Representation: externally tagged + + Options: + "build" + Build from scratch. + + "data_type": render::SimpleDataType + The type of the input data. + + "file": string + Input data in the `.bin` binary format. + + "output": Option + Output file. + + If provided, saved data will go here. + + "from_previous" + Run from a previously generated output. + + "data_type": render::AnnotatedDataType + + "file": string + The previously generated output. + + "param1": string + This does one thing. + + "param2": Vec + This does another. + + Elements: render::AnnotatedDataType + Representation: externally tagged + + Options: + "float32" + Use high-precision. + + "float16" + Use lower precision. + + "int8" +======== From 8b2d91e2e8ea5fcc942b7f2ba7b2b9d5b57cdbbe Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Mon, 28 Sep 2026 16:59:48 -0700 Subject: [PATCH 14/21] `diskann-benchmark-simd`. --- diskann-benchmark-runner/src/utils/num.rs | 4 +- diskann-benchmark-simd/src/lib.rs | 52 ++++++++++++++++++++--- 2 files changed, 49 insertions(+), 7 deletions(-) diff --git a/diskann-benchmark-runner/src/utils/num.rs b/diskann-benchmark-runner/src/utils/num.rs index 049c2800f9..a769278259 100644 --- a/diskann-benchmark-runner/src/utils/num.rs +++ b/diskann-benchmark-runner/src/utils/num.rs @@ -8,6 +8,8 @@ use serde::{Deserialize, Deserializer, Serialize, Serializer}; use thiserror::Error; +use crate::Reflect; + /// Compute the relative change from `before` to `after`. /// /// This helper is intentionally opinionated for benchmark-style metrics: @@ -54,7 +56,7 @@ pub enum RelativeChangeError { } /// A finite floating-point value that is greater than or equal to zero. -#[derive(Debug, Clone, Copy, PartialEq, PartialOrd)] +#[derive(Debug, Clone, Copy, PartialEq, PartialOrd, Reflect)] pub struct NonNegativeFinite(f64); impl NonNegativeFinite { diff --git a/diskann-benchmark-simd/src/lib.rs b/diskann-benchmark-simd/src/lib.rs index 25b021d8a9..ba03f41455 100644 --- a/diskann-benchmark-simd/src/lib.rs +++ b/diskann-benchmark-simd/src/lib.rs @@ -26,7 +26,7 @@ use diskann_benchmark_runner::{ num::{relative_change, NonNegativeFinite}, percentiles, MicroSeconds, }, - Benchmark, Checker, Input, Registry, + Benchmark, Checker, Input, Registry, Reflect, }; //////////////// @@ -55,7 +55,7 @@ impl std::ops::Deref for DisplayWrapper<'_, T> { // Inputs // //////////// -#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "snake_case")] pub enum SimilarityMeasure { SquaredL2, @@ -74,17 +74,40 @@ impl std::fmt::Display for SimilarityMeasure { } } -#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +/// DiskANN provides specialization for different micro-architectures, allowing binaries +/// compiled for an older machine to still execute accelerated kernels when the runtime +/// CPU allows. +/// +/// This enum selects the target micro-architecture's implementation to benchmark. +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] enum Arch { + /// Target AVX-512 with additional VNNI and population-count instructions. + /// + /// This is generally IceLake server or better. + /// + /// Only usable when compiling for x86-64. #[serde(rename = "x86-64-v4")] #[allow(non_camel_case_types)] X86_64_V4, + + /// Target AVX2. + /// + /// Only usable when compiling for x86-64. #[serde(rename = "x86-64-v3")] #[allow(non_camel_case_types)] X86_64_V3, + /// Target the Aarch64 Neon instruction set. + /// + /// Only usable when compiling for aarch64. Neon, + + /// Use the `diskann-vector` scalar fallbacks. + /// + /// These use a mixture of actual scalar code and auto-vectorization friendly loops. Scalar, + + /// Reference operations using purely scalar loops. Reference, } @@ -101,20 +124,37 @@ impl std::fmt::Display for Arch { } } -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +/// Parameters controlling the benchmark for an individual kernel. +/// +/// Kernels work by using a fixed `query` vector and a variable of `data` vectors. +/// Internal timers measure how long it takes to compute all distances from `query` to each +/// data vector. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Reflect)] struct Run { + /// The kernel type to test. distance: SimilarityMeasure, + /// The number of dimensions in each vector. dim: NonZeroUsize, + /// The number of data points to loop over. num_points: NonZeroUsize, + /// The number of loops to run between timing samples. loops_per_measurement: NonZeroUsize, + /// The total number of measurements to take. num_measurements: NonZeroUsize, } -#[derive(Debug, Serialize, Deserialize)] +/// A full-precision SIMD Accelerated kernel. +#[derive(Debug, Serialize, Deserialize, Reflect)] pub struct SimdOp { + /// The data type for the left-hand side. + /// + /// This is typically the query in KNN scenarios. query_type: DataType, + /// The data type for the right-hand side. data_type: DataType, + /// The micro-architecture specialization. arch: Arch, + /// Kernel configurations. runs: Vec, } @@ -205,7 +245,7 @@ impl Input for SimdOp { /// /// Each field specifies the maximum allowed relative increase in the corresponding metric. /// For example, a value of `0.10` means a 10% increase is tolerated. -#[derive(Debug, Clone, Copy, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, Serialize, Deserialize, Reflect)] struct SimdTolerance { min_time_regression: NonNegativeFinite, } From 900cc263ae6a10164f1acf7bc845a313c2e60279 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Tue, 29 Sep 2026 10:23:37 -0700 Subject: [PATCH 15/21] Registry tests. --- diskann-benchmark-runner/src/reflect/mod.rs | 63 ++- .../src/reflect/render.rs | 7 +- .../src/reflect/test/serde.rs | 2 +- .../src/reflect/test/unit.rs | 5 +- diskann-benchmark-runner/src/registry.rs | 527 ++++++++++++++++-- 5 files changed, 508 insertions(+), 96 deletions(-) diff --git a/diskann-benchmark-runner/src/reflect/mod.rs b/diskann-benchmark-runner/src/reflect/mod.rs index 2b7bc836da..aca1f5872a 100644 --- a/diskann-benchmark-runner/src/reflect/mod.rs +++ b/diskann-benchmark-runner/src/reflect/mod.rs @@ -24,7 +24,7 @@ pub trait Reflect: 'static { #[derive(Clone, Copy)] pub struct Reflection { - reflection: &'static dyn internal::Reflect, + reflection: &'static internal::VTable, } impl Reflection { @@ -33,12 +33,12 @@ impl Reflection { T: Reflect, { Self { - reflection: &internal::Wrapper::::INSTANCE, + reflection: internal::VTable::new::(), } } pub fn ty(&self) -> Type { - self.reflection.ty() + (self.reflection.ty)() } pub fn type_name(&self) -> TypeName { @@ -46,7 +46,7 @@ impl Reflection { } pub fn type_id(&self) -> TypeId { - self.reflection.type_id() + (self.reflection.type_id)() } pub fn render(&self) -> Render { @@ -73,7 +73,7 @@ pub struct TypeName(Reflection); impl TypeName { fn format_type_name(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { - self.0.reflection.format_type_name(f) + (self.0.reflection.format_type_name)(f) } } @@ -238,39 +238,44 @@ where // Internal // ////////////// -pub(crate) mod internal { - use std::marker::PhantomData; - - pub(crate) trait Reflect { - fn ty(&self) -> super::Type; - fn format_type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result; - fn type_id(&self) -> std::any::TypeId; +pub mod internal { + pub(super) struct VTable { + pub(super) ty: fn() -> super::Type, + pub(super) format_type_name: fn(&mut dyn std::fmt::Write) -> std::fmt::Result, + pub(super) type_id: fn() -> std::any::TypeId, } - pub(crate) struct Wrapper(PhantomData); - - impl Wrapper { - pub(crate) const INSTANCE: Self = Self::new(); - - pub(crate) const fn new() -> Self { - Self(PhantomData) + impl VTable { + pub(super) const fn new() -> &'static Self + where + T: super::Reflect, + { + &Self { + ty: ty::, + format_type_name: format_type_name::, + type_id: type_id::, + } } } - impl Reflect for Wrapper + fn ty() -> super::Type where T: super::Reflect, { - fn ty(&self) -> super::Type { - ::ty() - } + ::ty() + } - fn format_type_name(&self, f: &mut dyn std::fmt::Write) -> std::fmt::Result { - ::format_type_name(f) - } + fn format_type_name(f: &mut dyn std::fmt::Write) -> std::fmt::Result + where + T: super::Reflect, + { + ::format_type_name(f) + } - fn type_id(&self) -> std::any::TypeId { - std::any::TypeId::of::() - } + fn type_id() -> std::any::TypeId + where + T: super::Reflect, + { + std::any::TypeId::of::() } } diff --git a/diskann-benchmark-runner/src/reflect/render.rs b/diskann-benchmark-runner/src/reflect/render.rs index e6d2bdb876..59f64ce76e 100644 --- a/diskann-benchmark-runner/src/reflect/render.rs +++ b/diskann-benchmark-runner/src/reflect/render.rs @@ -325,7 +325,7 @@ mod tests { use std::{ fs::File, - io::{self, BufRead, BufReader, Write}, + io::{BufRead, BufReader, Write}, path::{Path, PathBuf}, }; @@ -535,6 +535,7 @@ mod tests { #[derive(Reflect)] #[serde(rename_all = "kebab-case")] #[reflect(prefix = "render::")] + #[expect(unused, reason = "testing")] enum SimpleDataType { Float32, Float16, @@ -544,6 +545,7 @@ mod tests { #[derive(Reflect)] #[serde(rename_all = "kebab-case")] #[reflect(prefix = "render::")] + #[expect(unused, reason = "testing")] enum AnnotatedDataType { /// Use high-precision. Float32, @@ -555,6 +557,7 @@ mod tests { #[derive(Reflect)] #[serde(rename_all = "snake_case")] #[reflect(prefix = "render::")] + #[expect(unused, reason = "testing")] enum Source { /// Build from scratch. Build { @@ -579,10 +582,12 @@ mod tests { /// A new-type wrapper around `Source`. #[derive(Reflect)] + #[expect(unused, reason = "testing")] struct SourceWrapper(Source); /// A top level config. #[derive(Reflect)] + #[expect(unused, reason = "testing")] struct Config { source: SourceWrapper, /// This does one thing. diff --git a/diskann-benchmark-runner/src/reflect/test/serde.rs b/diskann-benchmark-runner/src/reflect/test/serde.rs index 917b95a4d0..170d7278d5 100644 --- a/diskann-benchmark-runner/src/reflect/test/serde.rs +++ b/diskann-benchmark-runner/src/reflect/test/serde.rs @@ -749,7 +749,7 @@ fn test_optionals() { enum CasesAdjacent { NewType(Option), Struct { val: Option }, - }; + } #[derive(Serialize, Reflect)] struct NewType(Option); diff --git a/diskann-benchmark-runner/src/reflect/test/unit.rs b/diskann-benchmark-runner/src/reflect/test/unit.rs index 7536fab5e1..cdc046d6a6 100644 --- a/diskann-benchmark-runner/src/reflect/test/unit.rs +++ b/diskann-benchmark-runner/src/reflect/test/unit.rs @@ -7,10 +7,7 @@ use std::assert_matches; -use crate::reflect::{ - tree::{Fields, Type}, - Reflect, Reflection, -}; +use crate::reflect::{tree::Fields, Reflect, Reflection}; #[test] fn test_unit() { diff --git a/diskann-benchmark-runner/src/registry.rs b/diskann-benchmark-runner/src/registry.rs index 040e2d5b83..fb1e0c07d2 100644 --- a/diskann-benchmark-runner/src/registry.rs +++ b/diskann-benchmark-runner/src/registry.rs @@ -25,9 +25,10 @@ pub struct Registry { /// Registered types for documentation. /// - /// The keys are generated by [`Reflection::type_name`]. We store the associated - /// [`Relflection`] and [`TypeId`] to ensure that registered type-names are unique. - name_map: HashMap, + /// The keys are generated by [`Reflection::type_name`]. + name_map: HashMap, + + /// The type IDs present in `name_map`. type_ids: HashSet, /// The registered benchmarks. @@ -45,6 +46,10 @@ impl Registry { } } + //-------// + // Input // + //-------// + /// Return the input with the registered `tag` if present. Otherwise, return `None`. /// /// Inputs are automatically registered as a side-effect of: @@ -59,6 +64,15 @@ impl Registry { self.inputs.values().map(|v| input::Registered(&**v)) } + //-----------// + // Type Info // + //-----------// + + /// Return the [`Reflection`] type information for the provided type-name. + pub fn type_info(&self, type_name: &str) -> Option { + self.name_map.get(type_name).copied() + } + //--------------// // Registration // //--------------// @@ -246,37 +260,11 @@ impl Registry { Entry::Vacant(v) => { // Before we insert - try to add all the reflection types. If this fails, // then we haven't committed the input to the registry. - // - // TODO: Enable roll-back of any newly registered type. - Reflection::new::() - .visit_with(|reflection: Reflection| -> Result { - let type_id = reflection.type_id(); - - // If the type doesn't exist - try to add it via type-name. - // Only if all that succeeds to we commit the type-id. - if let hash_set::Entry::Vacant(vacant) = self.type_ids.entry(type_id) { - let type_name = reflection.type_name().to_string(); - match self.name_map.entry(type_name) { - Entry::Vacant(v) => { - v.insert((reflection, type_id)); - } - Entry::Occupied(o) => { - if o.get().1 != type_id { - panic!("eventually we will fail gracetully"); - } - } - } - - vacant.insert(); - - // Continue exploring. - Ok(true) - } else { - // Already seen this node. No need to recurse. - Ok(false) - } - }) - .unwrap(); + Self::register_reflection( + Reflection::new::(), + &mut self.name_map, + &mut self.type_ids, + )?; v.insert(Box::new(wrapper)); Ok(()) @@ -291,17 +279,17 @@ impl Registry { .as_any() .downcast_ref::() { - Err(RegistryError { + Err(RegistryError::input_conflict( tag, - existing: Kind::Gated(existing.features().to_string()), - new: Kind::Available(wrapper.type_name()), - }) + Kind::Gated(existing.features().to_string()), + Kind::Available(wrapper.type_name()), + )) } else { - Err(RegistryError { + Err(RegistryError::input_conflict( tag, - existing: Kind::Available(o.get().type_name()), - new: Kind::Available(wrapper.type_name()), - }) + Kind::Available(o.get().type_name()), + Kind::Available(wrapper.type_name()), + )) } } } @@ -323,23 +311,118 @@ impl Registry { .downcast_ref::() { if existing.features() != input.features() { - Err(RegistryError { - tag: input.tag(), - existing: Kind::Gated(existing.features().to_string()), - new: Kind::Gated(input.features().to_string()), - }) + Err(RegistryError::input_conflict( + input.tag(), + Kind::Gated(existing.features().to_string()), + Kind::Gated(input.features().to_string()), + )) } else { Ok(()) } } else { let type_name = o.get().type_name(); - Err(RegistryError { - tag: input.tag(), - existing: Kind::Available(type_name), - new: Kind::Gated(input.features().to_string()), - }) + Err(RegistryError::input_conflict( + input.tag(), + Kind::Available(type_name), + Kind::Gated(input.features().to_string()), + )) + } + } + } + } + + //----------------// + // Type Catalogue // + //----------------// + + #[cfg(test)] + fn test_register_reflection(&mut self) -> Result<(), RegistryError> + where + T: crate::Reflect, + { + Self::register_reflection( + Reflection::new::(), + &mut self.name_map, + &mut self.type_ids, + ) + } + + /// Register `reflection` and all types reachable from it into the type catalogue. + /// + /// This is called while an entry to `self.inputs` is held, so if modeled as an associated + /// function instead of a method. + fn register_reflection( + reflection: Reflection, + name_map: &mut HashMap, + type_ids: &mut HashSet, + ) -> Result<(), RegistryError> { + struct Conflict { + type_name: String, + } + + // To implement roll-back, we record the `Reflections` added this round. + // + // If all goes well, this can be discarded at the end. But if there is an error, + // we use this type to undo additions. + // + // This makes the error path more expensive but this is expected to be less common + // and means that we have less work to do on the happy path. + let mut added = Vec::::new(); + + let result: Result<(), Conflict> = + reflection.visit_with(|r: Reflection| -> Result { + let type_id = r.type_id(); + + // If the type doesn't exist - try to add it via type-name. + // Only if all that succeeds to we commit the type-id. + if let hash_set::Entry::Vacant(vacant_type_id) = type_ids.entry(type_id) { + let type_name = r.type_name().to_string(); + match name_map.entry(type_name) { + // Everything checks out, we're good to go. + Entry::Vacant(vacant_type_name) => { + vacant_type_name.insert(r); + vacant_type_id.insert(); + added.push(type_id); + } + Entry::Occupied(occupied_type_name) => { + // This branch is only reachable if `type_id` did not exist + // in `type_ids`, which means it never could have entered into + // `name_map` anyways. + debug_assert_ne!(occupied_type_name.get().type_id(), type_id,); + + return Err(Conflict { + type_name: occupied_type_name.key().clone(), + }); + } + } + + // Continue exploring. + Ok(true) + } else { + // Already seen this node. No need to recurse. + Ok(false) } + }); + + // Slow error recovery path. + // + // We first roll-back the type IDs that were added, which is relatively easy. + // Then, we cannot rely on `Reflect::type_name` returning the same name every time. + // So instead, we have to loop over all entries in `name_map` and inspect the + // `TypeId`s directly (which we control) to determing removal. + // + // Fortunately, `extract_if` does exactly what we want. + if let Err(conflict) = result { + let mut added_type_ids = HashSet::new(); + for type_id in added { + type_ids.remove(&type_id); + added_type_ids.insert(type_id); } + + name_map.retain(|_, v| !added_type_ids.contains(&v.type_id())); + Err(RegistryError::type_name_conflict(conflict.type_name)) + } else { + Ok(()) } } @@ -528,16 +611,43 @@ impl std::fmt::Display for Display<'_> { /// Error for [`Registry::register`] or [`Registry::register_regression`]. #[derive(Debug, Error)] -#[error( - "A different input with tag \"{}\" was already registered. Existing {}. New {}", - self.tag, - self.existing, - self.new, -)] +#[error(transparent)] pub struct RegistryError { - tag: &'static str, - existing: Kind, - new: Kind, + inner: RegistryErrorInner, +} + +impl RegistryError { + fn input_conflict(tag: &'static str, existing: Kind, new: Kind) -> Self { + Self { + inner: RegistryErrorInner::InputConflict { tag, existing, new }, + } + } + + fn type_name_conflict(type_name: String) -> Self { + Self { + inner: RegistryErrorInner::TypeNameConflict { type_name }, + } + } +} + +#[derive(Debug, Error)] +enum RegistryErrorInner { + #[error( + "A different input with tag \"{}\" was already registered. Existing {}. New {}", + tag, + existing, + new + )] + InputConflict { + tag: &'static str, + existing: Kind, + new: Kind, + }, + #[error( + "A different type with the type name \"{}\" was already registered", + type_name + )] + TypeNameConflict { type_name: String }, } #[derive(Debug)] @@ -810,3 +920,298 @@ mod tests { ); } } + +#[cfg(test)] +mod test_type_catalog { + use super::*; + + use serde::{Deserialize, Serialize}; + + use crate::{Checker, Reflect}; + + // The types `A`, `B`, `C`, etc. are the valid types for testing registration. + + #[derive(Reflect)] + #[expect(unused, reason = "testing")] + struct A { + a: usize, + b: Option, + } + + #[derive(Reflect)] + #[expect(unused, reason = "testing")] + struct B(A, f32); + + #[derive(Reflect)] + #[expect(unused, reason = "testing")] + struct C { + a: f64, + b: isize, + } + + #[derive(Reflect)] + #[expect(unused, reason = "testing")] + enum D { + V0(A), + V1 { first: B, second: B }, + V2, + V3(Vec>), + } + + // The type tree `Bad0`, `Bad1`, etc. are the types we use to test the erroring path. + // + // The idea is that somewhere at the bottom of this path is a type with a type-name + // that conflicts with one of the "good" types above. + #[derive(Reflect)] + #[reflect(type_name = "usize")] + struct Bad0; + + #[derive(Reflect)] + #[expect(unused, reason = "testing")] + struct Bad1 { + a: usize, + b: D, + c: i32, + d: Option, + } + + #[derive(Reflect)] + #[expect(unused, reason = "testing")] + struct Bad2(u32, Vec); + + #[derive(Reflect)] + #[expect(unused, reason = "testing")] + enum Bad3 { + Unit, + S(String), + B { field: Bad2 }, + } + + fn type_name_for() -> String + where + T: Reflect, + { + Reflection::new::().type_name().to_string() + } + + fn assert_is_present(registry: &Registry) + where + T: Reflect, + { + let type_name = type_name_for::(); + let r = registry.type_info(&type_name).unwrap(); + + assert_eq!(r.type_id(), TypeId::of::()); + assert_eq!(r.type_name().to_string(), type_name); + } + + fn assert_is_absent(registry: &Registry) + where + T: Reflect, + { + let type_name = type_name_for::(); + if let Some(r) = registry.type_info(&type_name) { + assert_ne!( + r.type_id(), + Reflection::new::().type_id(), + "Did not expect \"{}\" to be registered", + std::any::type_name::(), + ); + } + } + + // We expect the following types to be present. + // + // * usize + // * String + // * Option + // * f32 + // * f64 + // * isize + // * A + // * B + // * C + // * Option + // * Vec> + // * D + // + // The following types should *not* be present. + // + // * i32 + // * Bad0 + // * Option + // * Bad1 + // * u32 + // * Vec + // * Bad2 + // * Bad3 + fn check_types_are_present(r: &Registry) { + assert_is_present::(r); + assert_is_present::(r); + assert_is_present::>(r); + assert_is_present::(r); + assert_is_present::(r); + assert_is_present::(r); + assert_is_present::(r); + assert_is_present::(r); + assert_is_present::(r); + assert_is_present::>(r); + assert_is_present::>>(r); + assert_is_present::(r); + + assert_is_absent::(r); + assert_is_absent::(r); + assert_is_absent::>(r); + assert_is_absent::(r); + assert_is_absent::(r); + assert_is_absent::>(r); + assert_is_absent::(r); + assert_is_absent::(r); + } + + // Register the entire type-tree at once. + #[test] + fn test_registration_all() { + let mut registry = Registry::new(); + registry.test_register_reflection::().unwrap(); + check_types_are_present(®istry); + + // Registering again should be fine. + registry.test_register_reflection::().unwrap(); + check_types_are_present(®istry); + } + + // Register portions of the type-tree. + #[test] + fn test_incremental_registration() { + let mut registry = Registry::new(); + registry.test_register_reflection::().unwrap(); + registry.test_register_reflection::().unwrap(); + registry.test_register_reflection::().unwrap(); + registry.test_register_reflection::().unwrap(); + + check_types_are_present(®istry); + + registry.test_register_reflection::().unwrap(); + check_types_are_present(®istry); + } + + // Immediate errors do not leave invalid state in the registry. + #[test] + fn test_errors_are_recoverable() { + let mut registry = Registry::new(); + + // Registering the "bad" type should eventually hit a naming conflict. + // Everything should be rolled back and adding the "good" type should still work. + let err = registry.test_register_reflection::().unwrap_err(); + let msg = err.to_string(); + assert_eq!( + msg, + "A different type with the type name \"usize\" was already registered" + ); + + assert!( + registry.name_map.is_empty(), + "invalid registration should not commit registered items" + ); + + assert!( + registry.type_ids.is_empty(), + "invalid registration should not commit registered items" + ); + + // After this, we should succeed in adding more types correctly. + registry.test_register_reflection::().unwrap(); + check_types_are_present(®istry); + + // Again, bad registration should not corrupt the + registry.test_register_reflection::().unwrap_err(); + check_types_are_present(®istry); + } + + // Test that type registration happens before input registration, and that if type + // registration fails, the input is not registered. + #[test] + fn test_input_registration_aborts_correct() { + #[derive(Serialize, Deserialize, Reflect, Debug)] + struct GoodInput { + a: usize, + b: isize, + } + + impl Input for GoodInput { + type Raw = Self; + fn tag() -> &'static str { + "good-input" + } + fn from_raw(_raw: Self::Raw, _checker: &mut Checker) -> anyhow::Result { + unimplemented!("this struct is for test only"); + } + fn serialize(&self) -> anyhow::Result { + unimplemented!("this struct is for test only"); + } + fn example() -> Self::Raw { + unimplemented!("this struct is for test only"); + } + } + + #[derive(Serialize, Deserialize, Reflect, Debug)] + #[reflect(type_name = "isize")] + struct Boom; + + #[derive(Serialize, Deserialize, Reflect, Debug)] + struct BadInput { + /// This type should not be registered on failure. + a: f32, + b: Boom, + /// Put another `f32` on the other side of `Boom` so no matter the expansion + /// order, a `f32` is registered before `Boom`. + c: f32, + } + + impl Input for BadInput { + type Raw = Self; + fn tag() -> &'static str { + "bad-input" + } + fn from_raw(_raw: Self::Raw, _checker: &mut Checker) -> anyhow::Result { + unimplemented!("this struct is for test only"); + } + fn serialize(&self) -> anyhow::Result { + unimplemented!("this struct is for test only"); + } + fn example() -> Self::Raw { + unimplemented!("this struct is for test only"); + } + } + + let mut registry = Registry::new(); + registry.register_input::().unwrap(); + + assert!(registry.input("good-input").is_some()); + assert!(registry.input("bad-input").is_none()); + + assert_is_present::(®istry); + assert_is_present::(®istry); + assert_is_present::(®istry); + + assert_is_absent::(®istry); + assert_is_absent::(®istry); + + // This should hit a type-conflict. + registry.register_input::().unwrap_err(); + + assert!(registry.input("good-input").is_some()); + assert!( + registry.input("bad-input").is_none(), + "bad input should not be registered on failure" + ); + + assert_is_present::(®istry); + assert_is_present::(®istry); + assert_is_present::(®istry); + + assert_is_absent::(®istry); + assert_is_absent::(®istry); + } +} From 0530df6c8d2d899a6a3bc5db1f58b19bffb213ac Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Tue, 29 Sep 2026 17:23:43 -0700 Subject: [PATCH 16/21] trybuild test and docs. --- Cargo.lock | 2 + diskann-benchmark-runner-derive/Cargo.toml | 4 + diskann-benchmark-runner-derive/SERDE_TODO.md | 77 +++++++--- .../src/attributes.rs | 32 +++-- .../tests/compile.rs | 11 ++ .../tests/ui/fail/content_without_tag.rs | 14 ++ .../tests/ui/fail/content_without_tag.stderr | 5 + .../tests/ui/fail/duplicate_content.rs | 14 ++ .../tests/ui/fail/duplicate_content.stderr | 5 + .../tests/ui/fail/duplicate_field_rename.rs | 14 ++ .../ui/fail/duplicate_field_rename.stderr | 5 + .../tests/ui/fail/duplicate_prefix.rs | 12 ++ .../tests/ui/fail/duplicate_prefix.stderr | 5 + .../tests/ui/fail/duplicate_rename_all.rs | 14 ++ .../tests/ui/fail/duplicate_rename_all.stderr | 5 + .../tests/ui/fail/duplicate_tag.rs | 14 ++ .../tests/ui/fail/duplicate_tag.stderr | 5 + .../tests/ui/fail/duplicate_type_name.rs | 12 ++ .../tests/ui/fail/duplicate_type_name.stderr | 5 + .../tests/ui/fail/duplicate_variant_rename.rs | 14 ++ .../ui/fail/duplicate_variant_rename.stderr | 5 + .../tests/ui/fail/generic_type_name.rs | 14 ++ .../tests/ui/fail/generic_type_name.stderr | 5 + .../tests/ui/fail/prefix_and_type_name.rs | 12 ++ .../tests/ui/fail/prefix_and_type_name.stderr | 5 + .../tests/ui/fail/tag_content_on_struct.rs | 14 ++ .../ui/fail/tag_content_on_struct.stderr | 5 + .../tests/ui/fail/tag_on_struct.rs | 14 ++ .../tests/ui/fail/tag_on_struct.stderr | 5 + .../tests/ui/fail/union.rs | 13 ++ .../tests/ui/fail/union.stderr | 7 + .../fail/unsupported_container_attribute.rs | 14 ++ .../unsupported_container_attribute.stderr | 5 + .../ui/fail/unsupported_field_attribute.rs | 14 ++ .../fail/unsupported_field_attribute.stderr | 5 + .../ui/fail/unsupported_reflect_attribute.rs | 12 ++ .../fail/unsupported_reflect_attribute.stderr | 5 + .../tests/ui/fail/unsupported_rename_all.rs | 14 ++ .../ui/fail/unsupported_rename_all.stderr | 5 + .../ui/fail/unsupported_variant_attribute.rs | 14 ++ .../fail/unsupported_variant_attribute.stderr | 5 + .../tests/ui/pass/supported.rs | 31 ++++ diskann-benchmark-runner/src/reflect/mod.rs | 132 +++++++++++++++++- diskann-benchmark-runner/src/reflect/tree.rs | 86 +++++++----- diskann-benchmark-simd/src/lib.rs | 2 +- 45 files changed, 640 insertions(+), 72 deletions(-) create mode 100644 diskann-benchmark-runner-derive/tests/compile.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/content_without_tag.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/content_without_tag.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_content.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_content.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_field_rename.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_field_rename.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_prefix.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_prefix.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_rename_all.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_rename_all.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_tag.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_tag.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_type_name.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_type_name.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_variant_rename.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/duplicate_variant_rename.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/generic_type_name.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/generic_type_name.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/prefix_and_type_name.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/prefix_and_type_name.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/tag_content_on_struct.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/tag_content_on_struct.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/tag_on_struct.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/tag_on_struct.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/union.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/union.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_container_attribute.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_container_attribute.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_field_attribute.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_field_attribute.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_reflect_attribute.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_reflect_attribute.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_rename_all.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_rename_all.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_variant_attribute.rs create mode 100644 diskann-benchmark-runner-derive/tests/ui/fail/unsupported_variant_attribute.stderr create mode 100644 diskann-benchmark-runner-derive/tests/ui/pass/supported.rs diff --git a/Cargo.lock b/Cargo.lock index 39f0338718..79f9e7cbc5 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -531,9 +531,11 @@ dependencies = [ name = "diskann-benchmark-runner-derive" version = "0.59.0" dependencies = [ + "diskann-benchmark-runner", "proc-macro2", "quote", "syn 2.0.117", + "trybuild", ] [[package]] diff --git a/diskann-benchmark-runner-derive/Cargo.toml b/diskann-benchmark-runner-derive/Cargo.toml index 9672dbebcb..8d66decc3a 100644 --- a/diskann-benchmark-runner-derive/Cargo.toml +++ b/diskann-benchmark-runner-derive/Cargo.toml @@ -14,5 +14,9 @@ syn = { version = "2", features = ["full"] } quote = "1" proc-macro2 = "1" +[dev-dependencies] +diskann-benchmark-runner = { workspace = true } +trybuild = "1.0.120" + [lints] workspace = true diff --git a/diskann-benchmark-runner-derive/SERDE_TODO.md b/diskann-benchmark-runner-derive/SERDE_TODO.md index b8c19c43bc..8b965f78f5 100644 --- a/diskann-benchmark-runner-derive/SERDE_TODO.md +++ b/diskann-benchmark-runner-derive/SERDE_TODO.md @@ -12,39 +12,74 @@ enum representations. `x_m_l_http_request` under `snake_case`, while `heck` produces `xml_http_request`. - For fields, `snake_case` is identity and `kebab-case` replaces underscores. - - Implementing the supported transformations directly should allow removal of the `heck` - dependency. -- [ ] Reject tuple variants on internally tagged enums with an error at the variant. -- [ ] Confirm the intended treatment of internally tagged newtype variants. Serde accepts - some forms syntactically, but compatibility depends on the wrapped value's serialized - shape. -- [ ] Check raw identifiers such as `r#type`; reflected names must match Serde's wire names - rather than include the raw-identifier prefix. + - The `heck` dependency has been removed. +- [X] Confirm the intended treatment of internally tagged newtype variants. + - Newtypes remain supported because Serde accepts newtypes containing struct/map-like + values. + - Serde compatibility tests cover non-empty and empty struct payloads. +- [X] Strip the `r#` prefix from raw field and variant identifiers so reflected names match + Serde's wire names. + +## Attribute validation + +- [ ] Reject internally tagged tuple variants at the variant span. + - Newtype variants must remain supported because Serde permits newtypes containing + struct/map-like values. + - Do not classify syntactic newtypes as tuple variants. +- [ ] Reject `#[reflect(type_name = "...")]` on every generic type, including types whose + only generic parameters are lifetimes. + - Check `input.generics.params.is_empty()` rather than the generated list of displayed + type and const arguments. +- [ ] Reject adjacent enum representations whose `tag` and `content` names are equal. +- [ ] Reject fields in internally tagged variants whose effective serialized name conflicts + with the internal tag. + - Compare names after applying field `rename` and variant-level `rename_all`. +- [ ] Reject duplicate effective serialized names: + - enum variants after container `rename_all` and variant `rename`; + - named struct fields after container `rename_all` and field `rename`; + - named variant fields after variant-level `rename_all` and field `rename`. +- [ ] Reject misplaced `reflect` attributes instead of silently ignoring them. + - Until field-level features such as `#[reflect(opaque)]` exist, any `reflect` attribute + on a field or variant should produce an unsupported/misplaced-attribute diagnostic. +- [ ] Add deliberate diagnostics for asymmetric Serde naming syntax instead of relying on a + lower-level parse error: + - `rename(serialize = "...", deserialize = "...")`; + - `rename_all(serialize = "...", deserialize = "...")`. +- [X] Keep unsupported representation-changing attributes rejected until the reflection + model explicitly supports them, including `untagged`, `flatten`, `skip*`, `default`, + `alias`, `with`, `remote`, `from`, and `try_from`. ## Tests -- [ ] Update the enum smoke test to expect renamed variants such as `"unit"` rather than +- [X] Update enum compatibility coverage to expect renamed variants such as `"unit"` rather than `"Unit"`. -- [ ] Assert the generated enum representation, including both `tag` and `content` for an - adjacently tagged enum. +- [X] Verify the generated enum representation, including both `tag` and `content` for an + adjacently tagged enum, against serialized JSON. - [ ] Add naming tests that compare reflection metadata with `serde_json`, covering: - - explicit field and variant `rename`; - - struct field `rename_all`; - - enum variant `rename_all`; - - variant-level `rename_all` for struct-variant fields; - - acronym-heavy variants such as `XMLHttpRequest`; - - explicit `rename` taking precedence over `rename_all`. -- [ ] Add representation tests for external, internal, and adjacent tagging. + - [X] explicit field and variant `rename`; + - [X] struct field `rename_all`; + - [X] enum variant `rename_all`; + - [X] variant-level `rename_all` for struct-variant fields; + - [ ] acronym-heavy variants such as `XMLHttpRequest`; + - [X] explicit `rename` taking precedence over `rename_all`. +- [X] Add representation tests for external, internal, and adjacent tagging. - [ ] Add compile-fail tests for duplicate attributes, `content` without `tag`, enum-only attributes on structs, internally tagged tuple variants, unsupported rename rules, and unsupported Serde attributes. + - [X] duplicate attributes; + - [X] `content` without `tag`; + - [X] enum-only attributes on structs; + - [ ] internally tagged tuple variants; + - [X] unsupported rename rules; + - [X] unsupported Serde attributes. + - Add fixtures for the remaining validation rules above as their implementations land. ## Cleanup and validation -- [ ] Fix the `generate_type_name_body` doctest by returning the final +- [X] Fix the `generate_type_name_body` doctest by returning the final `f.write_str(">")` result instead of discarding it with a semicolon. -- [ ] Run `cargo fmt --all`. -- [ ] Run the targeted derive and runner tests. +- [X] Run `cargo fmt --all`. +- [X] Run the targeted derive and runner tests. - [ ] Run Clippy with warnings denied once the implementation and tests settle. ## Deferred Serde features diff --git a/diskann-benchmark-runner-derive/src/attributes.rs b/diskann-benchmark-runner-derive/src/attributes.rs index b82d7b9ba1..685f89f4fd 100644 --- a/diskann-benchmark-runner-derive/src/attributes.rs +++ b/diskann-benchmark-runner-derive/src/attributes.rs @@ -319,7 +319,7 @@ impl RenameOnce { } } -/// Strategy for generting type-names. +/// Strategy for generating type-names. #[derive(Default, Clone)] pub(crate) enum TypeName { #[default] @@ -335,18 +335,28 @@ enum TypeNameKind { impl TypeName { fn set_unique(&mut self, kind: TypeNameKind, value: syn::LitStr) -> syn::Result<()> { - if !matches!(self, Self::None) { - Err(syn::Error::new_spanned( - value, - "reflect attribute `prefix` found multiple times", - )) - } else { - match kind { - TypeNameKind::Prefix => *self = Self::Prefix(value), - TypeNameKind::Rename => *self = Self::Rename(value), + let error = match (&*self, &kind) { + (Self::Prefix(_), TypeNameKind::Prefix) => { + Some("reflect attribute `prefix` found multiple times") + } + (Self::Rename(_), TypeNameKind::Rename) => { + Some("reflect attribute `type_name` found multiple times") + } + (Self::Prefix(_), TypeNameKind::Rename) | (Self::Rename(_), TypeNameKind::Prefix) => { + Some("reflect attributes `prefix` and `type_name` are mutually exclusive") } - Ok(()) + (Self::None, _) => None, + }; + + if let Some(error) = error { + return Err(syn::Error::new_spanned(value, error)); + } + + match kind { + TypeNameKind::Prefix => *self = Self::Prefix(value), + TypeNameKind::Rename => *self = Self::Rename(value), } + Ok(()) } } diff --git a/diskann-benchmark-runner-derive/tests/compile.rs b/diskann-benchmark-runner-derive/tests/compile.rs new file mode 100644 index 0000000000..f79da56766 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/compile.rs @@ -0,0 +1,11 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +#[test] +fn compile_tests() { + let t = trybuild::TestCases::new(); + t.pass("tests/ui/pass/*.rs"); + t.compile_fail("tests/ui/fail/*.rs"); +} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/content_without_tag.rs b/diskann-benchmark-runner-derive/tests/ui/fail/content_without_tag.rs new file mode 100644 index 0000000000..a14122e252 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/content_without_tag.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[serde(content = "value")] +enum ContentWithoutTag { + Unit, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/content_without_tag.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/content_without_tag.stderr new file mode 100644 index 0000000000..199396823a --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/content_without_tag.stderr @@ -0,0 +1,5 @@ +error: serde attribute `content` provided without a `tag` + --> tests/ui/fail/content_without_tag.rs:9:19 + | +9 | #[serde(content = "value")] + | ^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_content.rs b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_content.rs new file mode 100644 index 0000000000..f7ff6a9d94 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_content.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[serde(tag = "kind", content = "first", content = "second")] +enum DuplicateContent { + Unit, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_content.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_content.stderr new file mode 100644 index 0000000000..43a4c717b5 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_content.stderr @@ -0,0 +1,5 @@ +error: serde attribute `content` found multiple times + --> tests/ui/fail/duplicate_content.rs:9:52 + | +9 | #[serde(tag = "kind", content = "first", content = "second")] + | ^^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_field_rename.rs b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_field_rename.rs new file mode 100644 index 0000000000..b6d58d9260 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_field_rename.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +struct DuplicateFieldRename { + #[serde(rename = "first", rename = "second")] + value: usize, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_field_rename.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_field_rename.stderr new file mode 100644 index 0000000000..450d655e75 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_field_rename.stderr @@ -0,0 +1,5 @@ +error: serde attribute `rename` found multiple times + --> tests/ui/fail/duplicate_field_rename.rs:10:40 + | +10 | #[serde(rename = "first", rename = "second")] + | ^^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_prefix.rs b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_prefix.rs new file mode 100644 index 0000000000..6ca1d1b106 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_prefix.rs @@ -0,0 +1,12 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[reflect(prefix = "one::", prefix = "two::")] +struct DuplicatePrefix; + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_prefix.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_prefix.stderr new file mode 100644 index 0000000000..adfb701b34 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_prefix.stderr @@ -0,0 +1,5 @@ +error: reflect attribute `prefix` found multiple times + --> tests/ui/fail/duplicate_prefix.rs:9:38 + | +9 | #[reflect(prefix = "one::", prefix = "two::")] + | ^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_rename_all.rs b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_rename_all.rs new file mode 100644 index 0000000000..1813af8ba0 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_rename_all.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[serde(rename_all = "snake_case", rename_all = "kebab-case")] +struct DuplicateRenameAll { + field_name: usize, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_rename_all.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_rename_all.stderr new file mode 100644 index 0000000000..826ecb0e67 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_rename_all.stderr @@ -0,0 +1,5 @@ +error: serde attribute `rename_all` found multiple times + --> tests/ui/fail/duplicate_rename_all.rs:9:49 + | +9 | #[serde(rename_all = "snake_case", rename_all = "kebab-case")] + | ^^^^^^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_tag.rs b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_tag.rs new file mode 100644 index 0000000000..1582ac8c02 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_tag.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[serde(tag = "first", tag = "second")] +enum DuplicateTag { + Unit, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_tag.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_tag.stderr new file mode 100644 index 0000000000..af85798900 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_tag.stderr @@ -0,0 +1,5 @@ +error: serde attribute `tag` found multiple times + --> tests/ui/fail/duplicate_tag.rs:9:30 + | +9 | #[serde(tag = "first", tag = "second")] + | ^^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_type_name.rs b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_type_name.rs new file mode 100644 index 0000000000..6c849db69a --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_type_name.rs @@ -0,0 +1,12 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[reflect(type_name = "One", type_name = "Two")] +struct DuplicateTypeName; + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_type_name.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_type_name.stderr new file mode 100644 index 0000000000..1b52d0ce53 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_type_name.stderr @@ -0,0 +1,5 @@ +error: reflect attribute `type_name` found multiple times + --> tests/ui/fail/duplicate_type_name.rs:9:42 + | +9 | #[reflect(type_name = "One", type_name = "Two")] + | ^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_variant_rename.rs b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_variant_rename.rs new file mode 100644 index 0000000000..d9a98f98e8 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_variant_rename.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +enum DuplicateVariantRename { + #[serde(rename = "first", rename = "second")] + Unit, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_variant_rename.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_variant_rename.stderr new file mode 100644 index 0000000000..d361937e3f --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/duplicate_variant_rename.stderr @@ -0,0 +1,5 @@ +error: serde attribute `rename` found multiple times + --> tests/ui/fail/duplicate_variant_rename.rs:10:40 + | +10 | #[serde(rename = "first", rename = "second")] + | ^^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/generic_type_name.rs b/diskann-benchmark-runner-derive/tests/ui/fail/generic_type_name.rs new file mode 100644 index 0000000000..2fe609a9b4 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/generic_type_name.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[reflect(type_name = "Generic")] +struct Generic { + value: T, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/generic_type_name.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/generic_type_name.stderr new file mode 100644 index 0000000000..d2acd34782 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/generic_type_name.stderr @@ -0,0 +1,5 @@ +error: The `type_name` attribute cannot be applied to types with generics + --> tests/ui/fail/generic_type_name.rs:9:23 + | +9 | #[reflect(type_name = "Generic")] + | ^^^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/prefix_and_type_name.rs b/diskann-benchmark-runner-derive/tests/ui/fail/prefix_and_type_name.rs new file mode 100644 index 0000000000..6a97894fca --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/prefix_and_type_name.rs @@ -0,0 +1,12 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[reflect(prefix = "benchmark::", type_name = "Renamed")] +struct PrefixAndTypeName; + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/prefix_and_type_name.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/prefix_and_type_name.stderr new file mode 100644 index 0000000000..2a279957f3 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/prefix_and_type_name.stderr @@ -0,0 +1,5 @@ +error: reflect attributes `prefix` and `type_name` are mutually exclusive + --> tests/ui/fail/prefix_and_type_name.rs:9:47 + | +9 | #[reflect(prefix = "benchmark::", type_name = "Renamed")] + | ^^^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/tag_content_on_struct.rs b/diskann-benchmark-runner-derive/tests/ui/fail/tag_content_on_struct.rs new file mode 100644 index 0000000000..ace75b315b --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/tag_content_on_struct.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[serde(tag = "kind", content = "value")] +struct AdjacentlyTaggedStruct { + value: usize, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/tag_content_on_struct.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/tag_content_on_struct.stderr new file mode 100644 index 0000000000..45fcdc96c1 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/tag_content_on_struct.stderr @@ -0,0 +1,5 @@ +error: serde attributes `tag` and `content` provided on a non-enum + --> tests/ui/fail/tag_content_on_struct.rs:9:15 + | +9 | #[serde(tag = "kind", content = "value")] + | ^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/tag_on_struct.rs b/diskann-benchmark-runner-derive/tests/ui/fail/tag_on_struct.rs new file mode 100644 index 0000000000..bfdd0417e0 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/tag_on_struct.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[serde(tag = "kind")] +struct TaggedStruct { + value: usize, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/tag_on_struct.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/tag_on_struct.stderr new file mode 100644 index 0000000000..3f249dd08c --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/tag_on_struct.stderr @@ -0,0 +1,5 @@ +error: serde attribute `tag` provided on a non-enum + --> tests/ui/fail/tag_on_struct.rs:9:15 + | +9 | #[serde(tag = "kind")] + | ^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/union.rs b/diskann-benchmark-runner-derive/tests/ui/fail/union.rs new file mode 100644 index 0000000000..2b7551dcc7 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/union.rs @@ -0,0 +1,13 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +union Unsupported { + value: usize, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/union.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/union.stderr new file mode 100644 index 0000000000..1ece235356 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/union.stderr @@ -0,0 +1,7 @@ +error: Reflect cannot be derived for unions + --> tests/ui/fail/union.rs:9:1 + | + 9 | / union Unsupported { +10 | | value: usize, +11 | | } + | |_^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_container_attribute.rs b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_container_attribute.rs new file mode 100644 index 0000000000..4bf89946e1 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_container_attribute.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[serde(default)] +struct UnsupportedContainerAttribute { + value: usize, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_container_attribute.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_container_attribute.stderr new file mode 100644 index 0000000000..a0b0d55b49 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_container_attribute.stderr @@ -0,0 +1,5 @@ +error: unsupported Serde attribute for Reflect + --> tests/ui/fail/unsupported_container_attribute.rs:9:9 + | +9 | #[serde(default)] + | ^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_field_attribute.rs b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_field_attribute.rs new file mode 100644 index 0000000000..80139d3a0e --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_field_attribute.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +struct UnsupportedFieldAttribute { + #[serde(default)] + value: usize, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_field_attribute.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_field_attribute.stderr new file mode 100644 index 0000000000..a285c1422d --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_field_attribute.stderr @@ -0,0 +1,5 @@ +error: unsupported Serde attribute for Reflect + --> tests/ui/fail/unsupported_field_attribute.rs:10:13 + | +10 | #[serde(default)] + | ^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_reflect_attribute.rs b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_reflect_attribute.rs new file mode 100644 index 0000000000..86fdc2405a --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_reflect_attribute.rs @@ -0,0 +1,12 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[reflect(rename = "Unsupported")] +struct UnsupportedReflectAttribute; + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_reflect_attribute.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_reflect_attribute.stderr new file mode 100644 index 0000000000..57e203891f --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_reflect_attribute.stderr @@ -0,0 +1,5 @@ +error: unsupported attribute for Reflect + --> tests/ui/fail/unsupported_reflect_attribute.rs:9:11 + | +9 | #[reflect(rename = "Unsupported")] + | ^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_rename_all.rs b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_rename_all.rs new file mode 100644 index 0000000000..0b0373789e --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_rename_all.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[serde(rename_all = "camelCase")] +struct UnsupportedRenameAll { + field_name: usize, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_rename_all.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_rename_all.stderr new file mode 100644 index 0000000000..6c797221f2 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_rename_all.stderr @@ -0,0 +1,5 @@ +error: unsupported serde `rename_all` rule "camelCase" - expected one of "lowercase", "snake_case", or "kebab-case" + --> tests/ui/fail/unsupported_rename_all.rs:9:22 + | +9 | #[serde(rename_all = "camelCase")] + | ^^^^^^^^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_variant_attribute.rs b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_variant_attribute.rs new file mode 100644 index 0000000000..ed9ac68c61 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_variant_attribute.rs @@ -0,0 +1,14 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +enum UnsupportedVariantAttribute { + #[serde(skip)] + Unit, +} + +fn main() {} diff --git a/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_variant_attribute.stderr b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_variant_attribute.stderr new file mode 100644 index 0000000000..76a6ce81cc --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/fail/unsupported_variant_attribute.stderr @@ -0,0 +1,5 @@ +error: unsupported Serde attribute for Reflect + --> tests/ui/fail/unsupported_variant_attribute.rs:10:13 + | +10 | #[serde(skip)] + | ^^^^ diff --git a/diskann-benchmark-runner-derive/tests/ui/pass/supported.rs b/diskann-benchmark-runner-derive/tests/ui/pass/supported.rs new file mode 100644 index 0000000000..b668392561 --- /dev/null +++ b/diskann-benchmark-runner-derive/tests/ui/pass/supported.rs @@ -0,0 +1,31 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use diskann_benchmark_runner::Reflect; + +#[derive(Reflect)] +#[reflect(prefix = "benchmark::")] +#[serde(rename_all = "kebab-case", tag = "kind", content = "value")] +enum Generic { + Unit, + #[serde(rename = "tuple")] + Tuple(T, Vec), + #[serde(rename_all = "snake_case")] + Struct { + #[serde(rename = "renamed")] + field_name: T, + }, +} + +#[derive(Reflect)] +#[reflect(type_name = "CompleteOverride")] +struct Renamed { + value: usize, +} + +fn main() { + let _ = as Reflect>::ty(); + let _ = ::ty(); +} diff --git a/diskann-benchmark-runner/src/reflect/mod.rs b/diskann-benchmark-runner/src/reflect/mod.rs index aca1f5872a..a322295705 100644 --- a/diskann-benchmark-runner/src/reflect/mod.rs +++ b/diskann-benchmark-runner/src/reflect/mod.rs @@ -3,6 +3,8 @@ * Licensed under the MIT license. */ +//! Run time inspection of types. + use std::{ any::TypeId, fmt::{self, Write}, @@ -17,17 +19,123 @@ pub use tree::Type; #[cfg(test)] mod test; +/// Provide run-time information about compile-time types, including documentation. +/// +/// This trait is derivable and supports a subset of `serde` attributes. +/// +/// ``` +/// use diskann_benchmark_runner::{Reflect, Reflection}; +/// +/// /// An example struct. +/// #[derive(Reflect)] +/// struct Foo { +/// /// An awesome field. +/// #[serde(rename = "bar")] +/// foo: usize, +/// baz: usize, +/// } +/// +/// // Get information about `Foo`. +/// let ty = Foo::ty(); +/// +/// // The documentation of `Foo` will automatically be extracted from the docstrings. +/// assert_eq!(ty.doc().unwrap(), "An example struct."); +/// +/// // The type name can be extracted as well. +/// let mut name = String::new(); +/// Foo::format_type_name(&mut name); +/// assert_eq!(name, "Foo"); +/// +/// // To make typenames easier to extract, a `Reflection` can be use. +/// let reflection = Reflection::new::(); +/// assert_eq!(reflection.type_name().to_string(), "Foo"); +/// ``` +/// +/// # Attributes +/// +/// For derivation, a combination of `serde` attributes and custom `reflect` attributes +/// are supported. +/// +/// ## Serde Attributes +/// +/// Like the [serde crate](https://serde.rs/attributes.html), attributes are categorized by +/// container, variant, or field. +/// +/// ### Container Attributes +/// +/// * `#[serde(rename_all = "...")]`: Rename all fields (if a struct) or variants (if enum) +/// arrocding to the given case. +/// +/// Possible values are "lowercase", "snake_case", and "kebab-case". +/// +/// * `#[serde(tag = "type")]`: Used for internally tagged enums. +/// +/// * `#[serde(tag = "t", content = "c")]`: Used for adjacently tagged enums. +/// +/// ### Variant Attributes +/// +/// * `#[serde(rename = "name")]`: Describe with the given name instead of its Rust name. +/// +/// * `#[serde(rename_all = "...)]`: Rename all fiels of this struct variant with the given +/// case convention. +/// +/// Possible values are "lowercase", "snake_case", and "kebab-case". +/// +/// ### Field Attributes +/// +/// * `#[serde(rename = "name")]`: Describe with the given name instead of its Rust name. +/// +/// ## Reflect Attributes +/// +/// ### Container Attributes +/// +/// * `#[reflect(type_name = "name")]`: Use the given name to describe a struct instead of +/// its Rust name. This is used to avoid naming conflicts within a [`Registry`], which +/// enforces that type names are unique. +/// +/// Because of this, this attribute cannot be used on generic structs. +/// +/// This is mutually exclusive with the `prefix` attribute. +/// +/// ``` +/// use diskann_benchmark_runner::{Reflect, Reflection}; +/// +/// #[derive(Reflect)] +/// #[reflect(type_name = "Bar")] +/// struct Foo; +/// +/// assert_eq!(Reflection::new::().type_name().to_string(), "Bar"); +/// ``` +/// +/// * `#[reflect(prefix = "...")]`: Prefix the Rust name with the provided prefix. Like the +/// `type_name` attribute, this can be used to create name spaces to help generate unique +/// type names. +/// +/// ``` +/// use diskann_benchmark_runner::{Reflect, Reflection}; +/// +/// #[derive(Reflect)] +/// #[reflect(prefix = "mod::")] +/// struct Foo; +/// +/// assert_eq!(Reflection::new::().type_name().to_string(), "mod::Foo"); +/// ``` pub trait Reflect: 'static { + /// Return the [`Type`] containing the information about `self`. fn ty() -> Type; + + /// Write the type-name for `self` into the buffer. fn format_type_name(f: &mut dyn Write) -> fmt::Result; } +/// A [`Reflect`]ed type. #[derive(Clone, Copy)] pub struct Reflection { reflection: &'static internal::VTable, } impl Reflection { + /// Construct a new [`Reflection`] for `T`. pub const fn new() -> Self where T: Reflect, @@ -37,22 +145,36 @@ impl Reflection { } } + /// Return the [`Type`] for the type being reflected. pub fn ty(&self) -> Type { (self.reflection.ty)() } + /// Return a [`std::fmt::Display`] compatible struct for rendering the name of the type + /// being reflected. pub fn type_name(&self) -> TypeName { TypeName(*self) } + /// Return the [`TypeId`] of the pub fn type_id(&self) -> TypeId { (self.reflection.type_id)() } - pub fn render(&self) -> Render { + /// Return a [`std::fmt::Display`] comaptible struct for rendering the reflected type. + pub(crate) fn render(&self) -> Render { Render(*self) } + /// Visit all types reachable from the reflected type. + /// + /// This will traverse through all structs, enum variants, container types etc. + /// + /// The closure `f` can be used to direct the exploration by returning the following values: + /// + /// * `Ok(true)`: Continue exploring through the argument [`Reflection`]. + /// * `Ok(false)`: Do not continue exploring through the argument [`Reflection`]. + /// * `Err(E)`: Immediately stop exploring and return the error `E`. pub(crate) fn visit_with(&self, f: F) -> Result<(), E> where F: FnMut(Reflection) -> Result, @@ -69,6 +191,9 @@ impl fmt::Debug for Reflection { } } +/// A [`std::fmt::Display`] compatible type for [`Reflection`]. +/// +/// See: [`Reflection::type_name`]. pub struct TypeName(Reflection); impl TypeName { @@ -89,7 +214,10 @@ impl std::fmt::Display for TypeName { } } -pub struct Render(Reflection); +/// A [`std::fmt::Display`] compatible type for rendering a [`Reflection`]. +/// +/// See: [`Reflection::render`]. +pub(crate) struct Render(Reflection); impl std::fmt::Display for Render { fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { diff --git a/diskann-benchmark-runner/src/reflect/tree.rs b/diskann-benchmark-runner/src/reflect/tree.rs index d92df1ae25..989cd09061 100644 --- a/diskann-benchmark-runner/src/reflect/tree.rs +++ b/diskann-benchmark-runner/src/reflect/tree.rs @@ -3,10 +3,17 @@ * Licensed under the MIT license. */ +//! Model of the Rust type system. +//! +//! This is largely based on the structure of types in [`syn`](https://docs.rs/syn/latest/syn/) +//! with some extra entries (e.g. [`Type::Optional`]) for a closer representation with how +//! types are serialized by [`serde`]. + use super::{Reflect, Reflection}; pub type Doc = std::borrow::Cow<'static, str>; +/// Classification of types. #[derive(Debug)] pub enum Type { Primitive(Primitive), @@ -17,27 +24,31 @@ pub enum Type { } impl Type { + /// Construc a new [`Type::Primitive`]. pub fn primitive(kind: PrimitiveKind, doc: Option) -> Self { - Self::from(Primitive::new(kind, doc)) + Self::Primitive(Primitive::new(kind, doc)) } + /// Construc a new [`Type::Aggregate`]. pub fn aggregate(fields: Fields, doc: Option) -> Self { - Self::from(Aggregate::new(fields, doc)) + Self::Aggregate(Aggregate::new(fields, doc)) } + /// Construc a new [`Type::Enum`]. pub fn enum_( repr: EnumRepr, variants: impl IntoIterator, doc: Option, ) -> Self { - Self::from(Enum::new(repr, variants, doc)) + Self::Enum(Enum::new(repr, variants, doc)) } + /// Construc a new [`Type::Sequence`]. pub fn sequence(doc: Option) -> Self where T: Reflect, { - Self::from(Sequence::new::(doc)) + Self::Sequence(Sequence::new::(doc)) } /// Keep the constructor private since we don't want users constructing the very special @@ -46,9 +57,10 @@ impl Type { where T: Reflect, { - Self::from(Optional::new::(doc)) + Self::Optional(Optional::new::(doc)) } + /// Return the struct level documentation if available. pub fn doc(&self) -> Option<&str> { match self { Self::Primitive(p) => p.doc(), @@ -59,6 +71,7 @@ impl Type { } } + /// Return `true` if there is field level information of some kind to render. pub(super) fn has_body(&self) -> bool { match self { Self::Primitive(_) => false, @@ -88,40 +101,11 @@ impl Type { } } -impl From for Type { - fn from(primitive: Primitive) -> Self { - Self::Primitive(primitive) - } -} - -impl From for Type { - fn from(aggergate: Aggregate) -> Self { - Self::Aggregate(aggergate) - } -} - -impl From for Type { - fn from(e: Enum) -> Self { - Self::Enum(e) - } -} - -impl From for Type { - fn from(s: Sequence) -> Self { - Self::Sequence(s) - } -} - -impl From for Type { - fn from(s: Optional) -> Self { - Self::Optional(s) - } -} - //-----------// // Primitive // //-----------// +/// The native JSON representation for a primitive. #[derive(Debug, Clone, Copy)] pub enum PrimitiveKind { Null, @@ -130,6 +114,7 @@ pub enum PrimitiveKind { String, } +/// A primitive type that maps closely to a native JSON type. #[derive(Debug)] pub struct Primitive { kind: PrimitiveKind, @@ -137,7 +122,7 @@ pub struct Primitive { } impl Primitive { - pub fn new(kind: PrimitiveKind, doc: Option) -> Self { + pub(crate) fn new(kind: PrimitiveKind, doc: Option) -> Self { Self { kind, doc } } @@ -154,6 +139,7 @@ impl Primitive { // Aggregate // //-----------// +/// A representation of aggregates like normal structs, unit structs, and tuple-like structs. #[derive(Debug)] pub struct Aggregate { fields: Fields, @@ -161,7 +147,7 @@ pub struct Aggregate { } impl Aggregate { - pub fn new(fields: Fields, doc: Option) -> Self { + pub(crate) fn new(fields: Fields, doc: Option) -> Self { Self { fields, doc } } @@ -178,11 +164,21 @@ impl Aggregate { } } +/// Represent the fields of a struct. #[derive(Debug)] pub enum Fields { + /// Standard Rust structs. Named(Vec), + + /// Tuple-like structs. Unnamed(Vec), + + /// Tuple-like structs with a single named field. + /// + /// These are treated specially by `serde` and thus get their own variant. NewType(UnnamedField), + + /// Unit structs. Unit, } @@ -196,14 +192,17 @@ impl Fields { } } + /// Construct [`Fields::Named`] from the iterator. pub fn named(itr: impl IntoIterator) -> Self { Self::Named(itr.into_iter().collect()) } + /// Construct [`Fields::Unamed`] from the iterator. pub fn unnamed(itr: impl IntoIterator) -> Self { Self::Unnamed(itr.into_iter().collect()) } + /// Construct [`Fields::NewType`] from the iterator. pub fn newtype(field: UnnamedField) -> Self { Self::NewType(field) } @@ -235,11 +234,13 @@ impl Fields { } } + /// Return `true` if `self` is [`Self::Unit`]. pub(super) fn is_unit(&self) -> bool { matches!(self, Self::Unit) } } +/// A struct field with a name. #[derive(Debug)] pub struct NamedField { name: &'static str, @@ -248,6 +249,7 @@ pub struct NamedField { } impl NamedField { + /// Construct a new [`NamedField`] for `T`. pub fn new(name: &'static str, doc: Option) -> Self where T: Reflect, @@ -272,6 +274,7 @@ impl NamedField { } } +/// An unnamed field. #[derive(Debug)] pub struct UnnamedField { field: Reflection, @@ -279,6 +282,7 @@ pub struct UnnamedField { } impl UnnamedField { + /// Construct a new [`UnnamedField`] for `T`. pub fn new(doc: Option) -> Self where T: Reflect, @@ -302,6 +306,7 @@ impl UnnamedField { // Enum // //------// +/// A representation for enums. #[derive(Debug)] pub struct Enum { repr: EnumRepr, @@ -310,6 +315,7 @@ pub struct Enum { } impl Enum { + /// Construct a new [`Enum`]. pub fn new( repr: EnumRepr, variants: impl IntoIterator, @@ -339,6 +345,7 @@ impl Enum { } } +/// Describe how an enum is being represented by `serde`. #[derive(Debug)] #[non_exhaustive] pub enum EnumRepr { @@ -381,6 +388,7 @@ pub enum EnumRepr { }, } +/// A variant of an [`Enum`]. #[derive(Debug)] pub struct Variant { name: &'static str, @@ -389,6 +397,7 @@ pub struct Variant { } impl Variant { + /// Construct a new [`Variant`]. pub fn new(name: &'static str, fields: Fields, doc: Option) -> Self { Self { name, fields, doc } } @@ -410,6 +419,7 @@ impl Variant { // Sequence // //----------// +/// A homogeneous sequence of values. #[derive(Debug)] pub struct Sequence { element: Reflection, @@ -417,6 +427,7 @@ pub struct Sequence { } impl Sequence { + /// Create a new [`Sequence`] containing `T`. pub fn new(doc: Option) -> Self where T: Reflect, @@ -440,6 +451,7 @@ impl Sequence { // Optional // //----------// +/// An [`Option`]. #[derive(Debug)] pub struct Optional { element: Reflection, diff --git a/diskann-benchmark-simd/src/lib.rs b/diskann-benchmark-simd/src/lib.rs index ba03f41455..57e74705e5 100644 --- a/diskann-benchmark-simd/src/lib.rs +++ b/diskann-benchmark-simd/src/lib.rs @@ -26,7 +26,7 @@ use diskann_benchmark_runner::{ num::{relative_change, NonNegativeFinite}, percentiles, MicroSeconds, }, - Benchmark, Checker, Input, Registry, Reflect, + Benchmark, Checker, Input, Reflect, Registry, }; //////////////// From 89f5acec7713dc2773c7536267001f7b8a6a5935 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Tue, 29 Sep 2026 17:54:26 -0700 Subject: [PATCH 17/21] minimal CLI hookup. --- diskann-benchmark-runner/dev/main.rs | 8 ++- diskann-benchmark-runner/src/app.rs | 49 +++++++++++++++++-- diskann-benchmark-runner/src/registry.rs | 5 ++ diskann-benchmark-runner/src/test/dim.rs | 2 + diskann-benchmark-runner/src/test/gated.rs | 2 + diskann-benchmark-runner/src/test/typed.rs | 2 + .../src/utils/datatype.rs | 1 + .../tests/benchmark/test-2/stdout.txt | 11 ++++- .../tests/benchmark/test-3/stdout.txt | 32 +++++++++++- 9 files changed, 106 insertions(+), 6 deletions(-) diff --git a/diskann-benchmark-runner/dev/main.rs b/diskann-benchmark-runner/dev/main.rs index 0d627c421d..9a499b408f 100644 --- a/diskann-benchmark-runner/dev/main.rs +++ b/diskann-benchmark-runner/dev/main.rs @@ -25,7 +25,13 @@ fn main() -> anyhow::Result<()> { #[derive(Debug, clap::Parser)] struct Cli { /// Emulate enabling various features for feature gated functionality. - #[arg(long, value_delimiter = ',', num_args = 0..)] + /// + /// Currently available features: + /// * `gated-feature-0` + /// * `gated-feature-1` + /// * `gated-feature-2` + /// * `gated-feature-3` + #[arg(long, value_delimiter = ',', num_args = 1)] features: Vec, /// The actual application. diff --git a/diskann-benchmark-runner/src/app.rs b/diskann-benchmark-runner/src/app.rs index d355d988fd..1d499e980b 100644 --- a/diskann-benchmark-runner/src/app.rs +++ b/diskann-benchmark-runner/src/app.rs @@ -121,6 +121,11 @@ pub enum Commands { }, #[command(subcommand)] Check(Check), + /// Provide information about all registered types. + TypeInfo { + /// Provide information for the given type. + describe: Option, + }, } /// Subcommands for regression check operations. @@ -207,15 +212,27 @@ impl App { if let Some(describe) = describe { if let Some(input) = registry.input(describe) { let repr = jobs::Unprocessed::format_input(input)?; + + // Render JSON. writeln!( output, - "The example JSON representation for \"{}\" is:", + "The example JSON representation for \"{}\" is:\n", describe )?; writeln!(output, "{}", serde_json::to_string_pretty(&repr)?)?; - // // TODO: Make a little nicer. - // writeln!(output, "{}", input.raw_reflection().unwrap().render())?; + // Render Type Info. + match input.raw_reflection() { + Some(reflection) => { + writeln!(output, "\nType Information:\n\n{}", reflection.render())?; + + writeln!( + output, + "More type information available using `type-info`" + )?; + } + None => writeln!(output, "\n\nNo Type Information Available")?, + } return Ok(()); } else { @@ -390,6 +407,9 @@ impl App { } // Extensions Commands::Check(check) => return self.check(check, registry, output), + + // Types + Commands::TypeInfo { describe } => self.type_info(describe.as_deref(), registry, output)?, }; Ok(()) } @@ -486,6 +506,29 @@ impl App { } } } + + fn type_info( + &self, + describe: Option<&str>, + registry: ®istry::Registry, + mut output: &mut dyn Output, + ) -> anyhow::Result<()> { + match describe { + Some(type_name) => match registry.type_info(type_name) { + Some(reflection) => writeln!(output, "{}", reflection.render())?, + None => anyhow::bail!("No type information for \"{}\"", type_name), + }, + None => { + let mut all_types: Vec<_> = registry.type_names().collect(); + all_types.sort_unstable(); + writeln!(output, "All registered types:")?; + for type_name in all_types { + writeln!(output, " {}", type_name)?; + } + } + } + Ok(()) + } } /////////// diff --git a/diskann-benchmark-runner/src/registry.rs b/diskann-benchmark-runner/src/registry.rs index fb1e0c07d2..948204b2e6 100644 --- a/diskann-benchmark-runner/src/registry.rs +++ b/diskann-benchmark-runner/src/registry.rs @@ -73,6 +73,11 @@ impl Registry { self.name_map.get(type_name).copied() } + /// Return an iterator over all the registered type-names. + pub fn type_names(&self) -> impl ExactSizeIterator { + self.name_map.keys().map(|v| &**v) + } + //--------------// // Registration // //--------------// diff --git a/diskann-benchmark-runner/src/test/dim.rs b/diskann-benchmark-runner/src/test/dim.rs index 6d99c15d05..faa7b26b97 100644 --- a/diskann-benchmark-runner/src/test/dim.rs +++ b/diskann-benchmark-runner/src/test/dim.rs @@ -17,6 +17,7 @@ use crate::{ /////////// #[derive(Debug, Clone, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "benchmark::test::")] pub(super) struct DimInput { dim: Option, } @@ -56,6 +57,7 @@ impl Input for DimInput { /////////////// #[derive(Debug, Clone, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "benchmark::test::")] pub(super) struct Tolerance { succeed: bool, error_in_check: bool, diff --git a/diskann-benchmark-runner/src/test/gated.rs b/diskann-benchmark-runner/src/test/gated.rs index 0d264d65eb..b62e2279dd 100644 --- a/diskann-benchmark-runner/src/test/gated.rs +++ b/diskann-benchmark-runner/src/test/gated.rs @@ -87,6 +87,7 @@ impl Benchmark for AnotherGatedBench { ////////////////////////////////////////////////// #[derive(Debug, Clone, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "benchmark::test::")] pub(super) struct SampleInput { value: String, } @@ -147,6 +148,7 @@ impl Benchmark for GatedWithIndependentInput { // controlling feature set is enabled, standing in for an input whose validation would otherwise // pull in a heavy optional dependency. #[derive(Debug, Clone, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "benchmark::test::")] pub(super) struct PhantomInput { value: usize, } diff --git a/diskann-benchmark-runner/src/test/typed.rs b/diskann-benchmark-runner/src/test/typed.rs index c863f3be44..e620ac45cc 100644 --- a/diskann-benchmark-runner/src/test/typed.rs +++ b/diskann-benchmark-runner/src/test/typed.rs @@ -26,6 +26,7 @@ pub(crate) struct TypeInput { /// This is a test input for testing corner cases in the benchmark runner. #[derive(Serialize, Deserialize, Reflect)] +#[reflect(prefix = "benchmark::test::")] pub(crate) struct TypeInputRaw { data_type: DataType, dim: usize, @@ -80,6 +81,7 @@ impl Input for TypeInput { /////////////// #[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "benchmark::test::")] pub(super) struct Tolerance { /// Should we return an error when `from_raw` is called? pub(super) error_when_checked: bool, diff --git a/diskann-benchmark-runner/src/utils/datatype.rs b/diskann-benchmark-runner/src/utils/datatype.rs index 5bf13d6d8c..bce0d70288 100644 --- a/diskann-benchmark-runner/src/utils/datatype.rs +++ b/diskann-benchmark-runner/src/utils/datatype.rs @@ -13,6 +13,7 @@ use crate::Reflect; /// See also: [`AsDataType`]. #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "lowercase")] +#[reflect(prefix = "benchmark::")] pub enum DataType { Float64, Float32, diff --git a/diskann-benchmark-runner/tests/benchmark/test-2/stdout.txt b/diskann-benchmark-runner/tests/benchmark/test-2/stdout.txt index 0c2c55cb3e..0440dc40af 100644 --- a/diskann-benchmark-runner/tests/benchmark/test-2/stdout.txt +++ b/diskann-benchmark-runner/tests/benchmark/test-2/stdout.txt @@ -1,7 +1,16 @@ The example JSON representation for "test-input-dim" is: + { "type": "test-input-dim", "content": { "dim": 128 } -} \ No newline at end of file +} + +Type Information: + +benchmark::test::DimInput + "dim": Option + May be `null`. + +More type information available using `type-info` \ No newline at end of file diff --git a/diskann-benchmark-runner/tests/benchmark/test-3/stdout.txt b/diskann-benchmark-runner/tests/benchmark/test-3/stdout.txt index 06d12d7086..dcf65baba1 100644 --- a/diskann-benchmark-runner/tests/benchmark/test-3/stdout.txt +++ b/diskann-benchmark-runner/tests/benchmark/test-3/stdout.txt @@ -1,4 +1,5 @@ The example JSON representation for "test-input-types" is: + { "type": "test-input-types", "content": { @@ -6,4 +7,33 @@ The example JSON representation for "test-input-types" is: "dim": 128, "error_when_checked": false } -} \ No newline at end of file +} + +Type Information: + +benchmark::test::TypeInputRaw + This is a test input for testing corner cases in the benchmark runner. + + "data_type": benchmark::DataType + Representation: externally tagged + + Options: + "float64" + "float32" + "float16" + "uint8" + "uint16" + "uint32" + "uint64" + "int8" + "int16" + "int32" + "int64" + "bool" + + "dim": usize + + "error_when_checked": bool + Should we return an error when deserializing? + +More type information available using `type-info` \ No newline at end of file From a550a079d94ece4efacb956834fa3deab169f9bd Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Tue, 29 Sep 2026 18:03:39 -0700 Subject: [PATCH 18/21] CLI baselines. --- .../benchmark/type-info-describe/stdin.txt | 1 + .../benchmark/type-info-describe/stdout.txt | 24 +++++++++++++++++++ .../type-info-list-all-features/features.txt | 4 ++++ .../type-info-list-all-features/stdin.txt | 1 + .../type-info-list-all-features/stdout.txt | 10 ++++++++ .../tests/benchmark/type-info-list/stdin.txt | 1 + .../tests/benchmark/type-info-list/stdout.txt | 9 +++++++ .../benchmark/type-info-missing/stdin.txt | 1 + .../benchmark/type-info-missing/stdout.txt | 1 + diskann-benchmark-simd/src/lib.rs | 5 ++++ 10 files changed, 57 insertions(+) create mode 100644 diskann-benchmark-runner/tests/benchmark/type-info-describe/stdin.txt create mode 100644 diskann-benchmark-runner/tests/benchmark/type-info-describe/stdout.txt create mode 100644 diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/features.txt create mode 100644 diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/stdin.txt create mode 100644 diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/stdout.txt create mode 100644 diskann-benchmark-runner/tests/benchmark/type-info-list/stdin.txt create mode 100644 diskann-benchmark-runner/tests/benchmark/type-info-list/stdout.txt create mode 100644 diskann-benchmark-runner/tests/benchmark/type-info-missing/stdin.txt create mode 100644 diskann-benchmark-runner/tests/benchmark/type-info-missing/stdout.txt diff --git a/diskann-benchmark-runner/tests/benchmark/type-info-describe/stdin.txt b/diskann-benchmark-runner/tests/benchmark/type-info-describe/stdin.txt new file mode 100644 index 0000000000..8e5adac4a3 --- /dev/null +++ b/diskann-benchmark-runner/tests/benchmark/type-info-describe/stdin.txt @@ -0,0 +1 @@ +type-info benchmark::test::TypeInputRaw diff --git a/diskann-benchmark-runner/tests/benchmark/type-info-describe/stdout.txt b/diskann-benchmark-runner/tests/benchmark/type-info-describe/stdout.txt new file mode 100644 index 0000000000..34b5423302 --- /dev/null +++ b/diskann-benchmark-runner/tests/benchmark/type-info-describe/stdout.txt @@ -0,0 +1,24 @@ +benchmark::test::TypeInputRaw + This is a test input for testing corner cases in the benchmark runner. + + "data_type": benchmark::DataType + Representation: externally tagged + + Options: + "float64" + "float32" + "float16" + "uint8" + "uint16" + "uint32" + "uint64" + "int8" + "int16" + "int32" + "int64" + "bool" + + "dim": usize + + "error_when_checked": bool + Should we return an error when deserializing? \ No newline at end of file diff --git a/diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/features.txt b/diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/features.txt new file mode 100644 index 0000000000..da13a750a7 --- /dev/null +++ b/diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/features.txt @@ -0,0 +1,4 @@ +gated-feature-0 +gated-feature-1 +gated-feature-2 +gated-feature-3 diff --git a/diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/stdin.txt b/diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/stdin.txt new file mode 100644 index 0000000000..3ae7b5f6ff --- /dev/null +++ b/diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/stdin.txt @@ -0,0 +1 @@ +type-info diff --git a/diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/stdout.txt b/diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/stdout.txt new file mode 100644 index 0000000000..007a1c6460 --- /dev/null +++ b/diskann-benchmark-runner/tests/benchmark/type-info-list-all-features/stdout.txt @@ -0,0 +1,10 @@ +All registered types: + Option + benchmark::DataType + benchmark::test::DimInput + benchmark::test::PhantomInput + benchmark::test::SampleInput + benchmark::test::TypeInputRaw + bool + string + usize \ No newline at end of file diff --git a/diskann-benchmark-runner/tests/benchmark/type-info-list/stdin.txt b/diskann-benchmark-runner/tests/benchmark/type-info-list/stdin.txt new file mode 100644 index 0000000000..3ae7b5f6ff --- /dev/null +++ b/diskann-benchmark-runner/tests/benchmark/type-info-list/stdin.txt @@ -0,0 +1 @@ +type-info diff --git a/diskann-benchmark-runner/tests/benchmark/type-info-list/stdout.txt b/diskann-benchmark-runner/tests/benchmark/type-info-list/stdout.txt new file mode 100644 index 0000000000..1b1f397e47 --- /dev/null +++ b/diskann-benchmark-runner/tests/benchmark/type-info-list/stdout.txt @@ -0,0 +1,9 @@ +All registered types: + Option + benchmark::DataType + benchmark::test::DimInput + benchmark::test::SampleInput + benchmark::test::TypeInputRaw + bool + string + usize \ No newline at end of file diff --git a/diskann-benchmark-runner/tests/benchmark/type-info-missing/stdin.txt b/diskann-benchmark-runner/tests/benchmark/type-info-missing/stdin.txt new file mode 100644 index 0000000000..a17e298506 --- /dev/null +++ b/diskann-benchmark-runner/tests/benchmark/type-info-missing/stdin.txt @@ -0,0 +1 @@ +type-info benchmark::test::Missing diff --git a/diskann-benchmark-runner/tests/benchmark/type-info-missing/stdout.txt b/diskann-benchmark-runner/tests/benchmark/type-info-missing/stdout.txt new file mode 100644 index 0000000000..3ae52a773c --- /dev/null +++ b/diskann-benchmark-runner/tests/benchmark/type-info-missing/stdout.txt @@ -0,0 +1 @@ +No type information for "benchmark::test::Missing" \ No newline at end of file diff --git a/diskann-benchmark-simd/src/lib.rs b/diskann-benchmark-simd/src/lib.rs index 57e74705e5..a003879a2f 100644 --- a/diskann-benchmark-simd/src/lib.rs +++ b/diskann-benchmark-simd/src/lib.rs @@ -57,6 +57,7 @@ impl std::ops::Deref for DisplayWrapper<'_, T> { #[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "snake_case")] +#[reflect(prefix = "simd::")] pub enum SimilarityMeasure { SquaredL2, InnerProduct, @@ -81,6 +82,7 @@ impl std::fmt::Display for SimilarityMeasure { /// This enum selects the target micro-architecture's implementation to benchmark. #[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] +#[reflect(prefix = "simd::")] enum Arch { /// Target AVX-512 with additional VNNI and population-count instructions. /// @@ -130,6 +132,7 @@ impl std::fmt::Display for Arch { /// Internal timers measure how long it takes to compute all distances from `query` to each /// data vector. #[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "simd::")] struct Run { /// The kernel type to test. distance: SimilarityMeasure, @@ -145,6 +148,7 @@ struct Run { /// A full-precision SIMD Accelerated kernel. #[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "simd::")] pub struct SimdOp { /// The data type for the left-hand side. /// @@ -246,6 +250,7 @@ impl Input for SimdOp { /// Each field specifies the maximum allowed relative increase in the corresponding metric. /// For example, a value of `0.10` means a 10% increase is tolerated. #[derive(Debug, Clone, Copy, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "simd::")] struct SimdTolerance { min_time_regression: NonNegativeFinite, } From 8a29df99fab9df2c56b3d6a30640f3a790296686 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Wed, 30 Sep 2026 18:00:34 -0700 Subject: [PATCH 19/21] CLoser! --- diskann-benchmark-runner-derive/src/lib.rs | 6 +- diskann-benchmark-runner/src/app.rs | 4 +- diskann-benchmark-runner/src/files.rs | 12 +- diskann-benchmark-runner/src/reflect/mod.rs | 7 + diskann-benchmark-runner/src/utils/mod.rs | 2 + .../src/utils/required.rs | 83 +++++++ diskann-benchmark-simd/src/lib.rs | 2 +- diskann-benchmark/src/index/benchmarks.rs | 7 +- diskann-benchmark/src/index/inmem/product.rs | 2 +- diskann-benchmark/src/index/inmem/scalar.rs | 2 +- diskann-benchmark/src/index/inmem2.rs | 26 +- diskann-benchmark/src/inputs/bftree.rs | 222 +++++++++--------- diskann-benchmark/src/inputs/disk.rs | 4 +- diskann-benchmark/src/inputs/exhaustive.rs | 32 ++- diskann-benchmark/src/inputs/filters.rs | 12 +- diskann-benchmark/src/inputs/flat.rs | 6 +- diskann-benchmark/src/inputs/graph_index.rs | 138 +++++++---- diskann-benchmark/src/inputs/mod.rs | 2 +- diskann-benchmark/src/inputs/multi_vector.rs | 8 +- diskann-benchmark/src/main.rs | 8 +- diskann-benchmark/src/multi_vector/driver.rs | 4 +- diskann-benchmark/src/utils/mod.rs | 4 +- 22 files changed, 381 insertions(+), 212 deletions(-) create mode 100644 diskann-benchmark-runner/src/utils/required.rs diff --git a/diskann-benchmark-runner-derive/src/lib.rs b/diskann-benchmark-runner-derive/src/lib.rs index 6da9748d21..c1ead2eea7 100644 --- a/diskann-benchmark-runner-derive/src/lib.rs +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -421,7 +421,11 @@ fn process_enum( rename_variant_fields, } = attributes::Variant::parse(&v.attrs)?; - let fields = build_fields(&v.fields, &mut generics, rename_variant_fields)?; + let fields = build_fields( + &v.fields, + &mut generics, + rename_variant_fields, + )?; // Rename the variant as needed. let name = rename_variant.apply_to_variant(name, rename_all); diff --git a/diskann-benchmark-runner/src/app.rs b/diskann-benchmark-runner/src/app.rs index 1d499e980b..63355e5b10 100644 --- a/diskann-benchmark-runner/src/app.rs +++ b/diskann-benchmark-runner/src/app.rs @@ -409,7 +409,9 @@ impl App { Commands::Check(check) => return self.check(check, registry, output), // Types - Commands::TypeInfo { describe } => self.type_info(describe.as_deref(), registry, output)?, + Commands::TypeInfo { describe } => { + self.type_info(describe.as_deref(), registry, output)? + } }; Ok(()) } diff --git a/diskann-benchmark-runner/src/files.rs b/diskann-benchmark-runner/src/files.rs index eb161d52ed..c7e037a9b0 100644 --- a/diskann-benchmark-runner/src/files.rs +++ b/diskann-benchmark-runner/src/files.rs @@ -7,7 +7,7 @@ use std::path::{Path, PathBuf}; use serde::{Deserialize, Serialize}; -use super::Checker; +use crate::{Checker, Reflect}; /// A file that is used as an input to for a benchmark. /// @@ -27,6 +27,16 @@ pub struct InputFile { path: PathBuf, } +impl Reflect for InputFile { + fn ty() -> crate::reflect::Type { + ::ty() + } + + fn format_type_name(f: &mut dyn std::fmt::Write) -> std::fmt::Result { + f.write_str("benchmark::InputFile") + } +} + impl InputFile { /// Create a new input file from the path-like `path`. pub fn new

(path: P) -> Self diff --git a/diskann-benchmark-runner/src/reflect/mod.rs b/diskann-benchmark-runner/src/reflect/mod.rs index a322295705..d535ab290b 100644 --- a/diskann-benchmark-runner/src/reflect/mod.rs +++ b/diskann-benchmark-runner/src/reflect/mod.rs @@ -324,6 +324,12 @@ primitive!(i64, Number, "A 64-bit signed integer", "i64"); primitive!(f32, Number, "An 32-bit floating-point number", "f32"); primitive!(f64, Number, "An 64-bit floating-point number", "f64"); +primitive!( + std::num::NonZeroU32, + Number, + "A system dependent, 32-bit unsigned integer", + "NonZero" +); primitive!( std::num::NonZeroUsize, Number, @@ -333,6 +339,7 @@ primitive!( primitive!(bool, Boolean, "A value of \"true\" or \"false\"", "bool"); primitive!(String, String, "A string", "string"); +primitive!(std::path::PathBuf, String, "A file path", "PathBuf"); impl Reflect for Option where diff --git a/diskann-benchmark-runner/src/utils/mod.rs b/diskann-benchmark-runner/src/utils/mod.rs index 9f48f9413e..24caccbd92 100644 --- a/diskann-benchmark-runner/src/utils/mod.rs +++ b/diskann-benchmark-runner/src/utils/mod.rs @@ -8,5 +8,7 @@ pub mod fmt; pub mod microseconds; pub mod num; pub mod percentiles; +mod required; pub use microseconds::MicroSeconds; +pub use required::RequiredOption; diff --git a/diskann-benchmark-runner/src/utils/required.rs b/diskann-benchmark-runner/src/utils/required.rs new file mode 100644 index 0000000000..fb31c18731 --- /dev/null +++ b/diskann-benchmark-runner/src/utils/required.rs @@ -0,0 +1,83 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use serde::{Deserialize, Serialize}; + +use crate::{reflect::tree, Reflect}; + +/// Like `Option`, but requires the containing field to be present in the input JSON. +/// +/// To represent [`None`], the field must be explicitly set to `null`. +#[derive(Default, Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Serialize, Deserialize)] +#[serde(transparent)] +pub struct RequiredOption(Option); + +impl RequiredOption { + pub fn new(opt: Option) -> Self { + Self(opt) + } + + pub fn some(value: T) -> Self { + Self::new(Some(value)) + } + + pub fn none() -> Self { + Self::new(None) + } + + pub fn unwrap_or(self, default: T) -> T { + self.0.unwrap_or(default) + } + + pub fn get_or_insert(&mut self, v: T) -> &mut T { + self.0.get_or_insert(v) + } + + pub fn into_inner(self) -> Option { + self.0 + } + + pub fn as_ref(&self) -> Option<&T> { + self.0.as_ref() + } + + pub fn as_mut(&mut self) -> Option<&mut T> { + self.0.as_mut() + } + + pub fn as_deref(&self) -> Option<&T::Target> + where + T: std::ops::Deref, + { + self.0.as_deref() + } + + pub fn as_deref_mut(&mut self) -> Option<&mut T::Target> + where + T: std::ops::DerefMut, + { + self.0.as_deref_mut() + } +} + +impl Reflect for RequiredOption +where + T: Reflect, +{ + fn ty() -> tree::Type { + let doc = "A required optional type.\n\n\ + Unlike `Option`, where an omitted field implies `None`, the field must be \ + present and explicitly set to `null` to represent `None`."; + + // We're special - we get to instantiate the `Optional` type. + tree::Type::optional::(Some(doc.into())) + } + + fn format_type_name(f: &mut dyn std::fmt::Write) -> std::fmt::Result { + f.write_str("RequiredOption<")?; + T::format_type_name(f)?; + f.write_str(">") + } +} diff --git a/diskann-benchmark-simd/src/lib.rs b/diskann-benchmark-simd/src/lib.rs index a003879a2f..e02af18589 100644 --- a/diskann-benchmark-simd/src/lib.rs +++ b/diskann-benchmark-simd/src/lib.rs @@ -56,7 +56,7 @@ impl std::ops::Deref for DisplayWrapper<'_, T> { //////////// #[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] -#[serde(rename_all = "snake_case")] +#[serde(rename_all = "kebab-case")] #[reflect(prefix = "simd::")] pub enum SimilarityMeasure { SquaredL2, diff --git a/diskann-benchmark/src/index/benchmarks.rs b/diskann-benchmark/src/index/benchmarks.rs index 34ae5ff5f5..aa231ae934 100644 --- a/diskann-benchmark/src/index/benchmarks.rs +++ b/diskann-benchmark/src/index/benchmarks.rs @@ -225,7 +225,7 @@ where build::set_start_points( index.provider(), data.as_view(), - *build.start_point_strategy(), + build.start_point_strategy(), )?; Ok(index) }, @@ -842,9 +842,10 @@ where { let topk = input.search_phase.as_topk()?; - let consolidate_threshold: f32 = input + let consolidate_threshold: f32 = *input .runbook_params .consolidate_threshold + .as_ref() .ok_or_else(|| anyhow::anyhow!("consolidate_threshold is required for inmem streaming"))?; let data = datafiles::load_dataset::(datafiles::BinFile(input.build.data()))?; @@ -864,7 +865,7 @@ where build::set_start_points( index.provider(), data.as_view(), - *input.build.start_point_strategy(), + input.build.start_point_strategy(), )?; let num_threads_and_tasks = NonZeroUsize::new(input.build.num_threads()).unwrap(); diff --git a/diskann-benchmark/src/index/inmem/product.rs b/diskann-benchmark/src/index/inmem/product.rs index 931ca725ce..0b2ab8f108 100644 --- a/diskann-benchmark/src/index/inmem/product.rs +++ b/diskann-benchmark/src/index/inmem/product.rs @@ -202,7 +202,7 @@ mod imp { build::set_start_points( index.provider(), data_view, - *build.start_point_strategy(), + build.start_point_strategy(), )?; Ok(index) }; diff --git a/diskann-benchmark/src/index/inmem/scalar.rs b/diskann-benchmark/src/index/inmem/scalar.rs index b10fd874be..34ca4de20f 100644 --- a/diskann-benchmark/src/index/inmem/scalar.rs +++ b/diskann-benchmark/src/index/inmem/scalar.rs @@ -254,7 +254,7 @@ mod imp { inmem::WithBits::<$N>::new(quantizer), common::NoDeletes, )?; - build::set_start_points(index.provider(), data_view, *build.start_point_strategy())?; + build::set_start_points(index.provider(), data_view, build.start_point_strategy())?; Ok(index) }; diff --git a/diskann-benchmark/src/index/inmem2.rs b/diskann-benchmark/src/index/inmem2.rs index 60b776d635..3e09e50368 100644 --- a/diskann-benchmark/src/index/inmem2.rs +++ b/diskann-benchmark/src/index/inmem2.rs @@ -62,14 +62,18 @@ pub(crate) fn register_benchmarks(registry: &mut Registry) -> anyhow::Result<()> mod dto { use super::*; - #[derive(Debug, Serialize, Deserialize)] + use diskann_benchmark_runner::Reflect; + + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "inmem::")] pub(super) struct KnnSweep { pub(super) search_n: usize, pub(super) search_l: Vec, pub(super) recall_k: usize, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "inmem::")] pub(super) struct KnnSearch { pub(super) queries: InputFile, pub(super) groundtruth: InputFile, @@ -78,14 +82,16 @@ mod dto { pub(super) runs: Vec, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "inmem::")] pub(super) struct Data { pub(super) data_type: DataType, pub(super) data: InputFile, pub(super) distance: SimilarityMeasure, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "inmem::")] pub(super) struct BuildParams { pub(super) pruned_degree: usize, pub(super) max_degree: usize, @@ -98,7 +104,8 @@ mod dto { // Streaming // //-----------// - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "inmem::")] pub(super) struct StreamingKnnSearch { pub(super) queries: InputFile, pub(super) reps: NonZeroUsize, @@ -106,7 +113,8 @@ mod dto { pub(super) runs: Vec, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "inmem::")] pub(super) struct RunBook { pub(super) path: InputFile, pub(super) dataset: String, @@ -119,14 +127,16 @@ mod dto { // Top Level Inputs // //------------------// - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "inmem::")] pub(super) struct StaticBuild { pub(super) data: Data, pub(super) build: BuildParams, pub(super) search: KnnSearch, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "inmem::")] pub(super) struct BigANNStreaming { pub(super) data: Data, pub(super) build: BuildParams, diff --git a/diskann-benchmark/src/inputs/bftree.rs b/diskann-benchmark/src/inputs/bftree.rs index 32cc7adb16..22dfc870e3 100644 --- a/diskann-benchmark/src/inputs/bftree.rs +++ b/diskann-benchmark/src/inputs/bftree.rs @@ -10,7 +10,7 @@ use crate::inputs::{ write_field, Example, PRINT_WIDTH, }; use diskann::graph::config; -use diskann_benchmark_runner::{utils::datatype::DataType, Checker}; +use diskann_benchmark_runner::{utils::{RequiredOption, datatype::DataType}, Checker, Reflect}; use diskann_bftree::BfTreeProviderParameters; use serde::{Deserialize, Serialize}; @@ -20,7 +20,8 @@ use serde::{Deserialize, Serialize}; /// /// Required fields control memory sizing before data spills to disk. /// Optional fields tune internal behavior and default to bf_tree's defaults. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "bftree::")] pub(crate) struct BfTreeStoreConfig { /// Size of the circular buffer (in-memory write cache) in bytes. pub(crate) cb_size_byte: usize, @@ -29,32 +30,25 @@ pub(crate) struct BfTreeStoreConfig { pub(crate) leaf_page_size: usize, /// Maximum record size that can be stored in the circular buffer. - #[serde(deserialize_with = "Deserialize::deserialize")] - pub(crate) cb_max_record_size: Option, + pub(crate) cb_max_record_size: RequiredOption, /// Minimum record size for the circular buffer. - #[serde(deserialize_with = "Deserialize::deserialize")] - pub(crate) cb_min_record_size: Option, + pub(crate) cb_min_record_size: RequiredOption, /// Probability (0-100) of promoting a read record to the front of the buffer. - #[serde(deserialize_with = "Deserialize::deserialize")] - pub(crate) read_promotion_rate: Option, + pub(crate) read_promotion_rate: RequiredOption, /// Probability (0-100) of promoting a scanned record to the front of the buffer. - #[serde(deserialize_with = "Deserialize::deserialize")] - pub(crate) scan_promotion_rate: Option, + pub(crate) scan_promotion_rate: RequiredOption, /// Ratio of buffer used before copy-on-access kicks in. - #[serde(deserialize_with = "Deserialize::deserialize")] - pub(crate) cb_copy_on_access_ratio: Option, + pub(crate) cb_copy_on_access_ratio: RequiredOption, /// Whether to cache full pages on read. - #[serde(deserialize_with = "Deserialize::deserialize")] - pub(crate) read_record_cache: Option, + pub(crate) read_record_cache: RequiredOption, /// If true, only use the in-memory circular buffer (no disk pages). - #[serde(deserialize_with = "Deserialize::deserialize")] - pub(crate) cache_only: Option, + pub(crate) cache_only: RequiredOption, /// Whether to enable CPR snapshot support for this store. pub(crate) use_snapshot: bool, @@ -86,26 +80,26 @@ impl BfTreeStoreConfig { let mut c = bf_tree::Config::default(); c.cb_size_byte(self.cb_size_byte); c.leaf_page_size(self.leaf_page_size); - if let Some(v) = self.cb_max_record_size { - c.cb_max_record_size(v); + if let Some(v) = self.cb_max_record_size.as_ref() { + c.cb_max_record_size(*v); } - if let Some(v) = self.cb_min_record_size { - c.cb_min_record_size(v); + if let Some(v) = self.cb_min_record_size.as_ref() { + c.cb_min_record_size(*v); } - if let Some(v) = self.read_promotion_rate { - c.read_promotion_rate(v); + if let Some(v) = self.read_promotion_rate.as_ref() { + c.read_promotion_rate(*v); } - if let Some(v) = self.scan_promotion_rate { - c.scan_promotion_rate(v); + if let Some(v) = self.scan_promotion_rate.as_ref() { + c.scan_promotion_rate(*v); } - if let Some(v) = self.cb_copy_on_access_ratio { - c.cb_copy_on_access_ratio(v); + if let Some(v) = self.cb_copy_on_access_ratio.as_ref() { + c.cb_copy_on_access_ratio(*v); } - if let Some(v) = self.read_record_cache { - c.read_record_cache(v); + if let Some(v) = self.read_record_cache.as_ref() { + c.read_record_cache(*v); } - if let Some(v) = self.cache_only { - c.cache_only(v); + if let Some(v) = self.cache_only.as_ref() { + c.cache_only(*v); } c.use_snapshot(self.use_snapshot); c @@ -139,13 +133,13 @@ impl Default for BfTreeStoreConfig { Self { cb_size_byte: 32 * 1024 * 1024, // 32MB leaf_page_size: 4096, - cb_max_record_size: None, - cb_min_record_size: None, - read_promotion_rate: None, - scan_promotion_rate: None, - cb_copy_on_access_ratio: None, - read_record_cache: None, - cache_only: None, + cb_max_record_size: RequiredOption::none(), + cb_min_record_size: RequiredOption::none(), + read_promotion_rate: RequiredOption::none(), + scan_promotion_rate: RequiredOption::none(), + cb_copy_on_access_ratio: RequiredOption::none(), + read_record_cache: RequiredOption::none(), + cache_only: RequiredOption::none(), use_snapshot: false, } } @@ -161,25 +155,25 @@ impl std::fmt::Display for BfTreeStoreConfig { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { write_field!(f, "cb_size_byte", self.cb_size_byte)?; write_field!(f, "leaf_page_size", self.leaf_page_size)?; - if let Some(v) = self.cb_max_record_size { + if let Some(v) = self.cb_max_record_size.as_ref() { write_field!(f, "cb_max_record_size", v)?; } - if let Some(v) = self.cb_min_record_size { + if let Some(v) = self.cb_min_record_size.as_ref() { write_field!(f, "cb_min_record_size", v)?; } - if let Some(v) = self.read_promotion_rate { + if let Some(v) = self.read_promotion_rate.as_ref() { write_field!(f, "read_promotion_rate", v)?; } - if let Some(v) = self.scan_promotion_rate { + if let Some(v) = self.scan_promotion_rate.as_ref() { write_field!(f, "scan_promotion_rate", v)?; } - if let Some(v) = self.cb_copy_on_access_ratio { + if let Some(v) = self.cb_copy_on_access_ratio.as_ref() { write_field!(f, "cb_copy_on_access_ratio", v)?; } - if let Some(v) = self.read_record_cache { + if let Some(v) = self.read_record_cache.as_ref() { write_field!(f, "read_record_cache", v)?; } - if let Some(v) = self.cache_only { + if let Some(v) = self.cache_only.as_ref() { write_field!(f, "cache_only", v)?; } Ok(()) @@ -190,7 +184,7 @@ impl std::fmt::Display for BfTreeStoreConfig { /// /// Returns the agreed-upon value. If no configs are present or none set /// `use_snapshot`, returns `false`. If configs disagree, returns an error. -fn reconcile_use_snapshot(configs: &[(&str, &Option)]) -> anyhow::Result { +fn reconcile_use_snapshot(configs: &[(&str, Option<&BfTreeStoreConfig>)]) -> anyhow::Result { let mut resolved: Option = None; for (name, config) in configs { @@ -218,9 +212,9 @@ fn bftree_parameters_from( build: &IndexBuild, num_points: usize, dim: usize, - vector_store_config: &Option, - neighbor_store_config: &Option, - quant_store_config: &Option, + vector_store_config: Option<&BfTreeStoreConfig>, + neighbor_store_config: Option<&BfTreeStoreConfig>, + quant_store_config: Option<&BfTreeStoreConfig>, ) -> anyhow::Result { let use_snapshot = reconcile_use_snapshot(&[ ("vector_store_config", vector_store_config), @@ -236,14 +230,14 @@ fn bftree_parameters_from( dim, metric: build.distance().into(), vector_provider_config: vector_store_config - .clone() + .cloned() .unwrap_or_default() .into_config(), neighbor_list_provider_config: neighbor_store_config - .clone() + .cloned() .unwrap_or_default() .into_config(), - quant_vector_provider_config: quant_store_config.clone().unwrap_or_default().into_config(), + quant_vector_provider_config: quant_store_config.cloned().unwrap_or_default().into_config(), graph_params: None, use_snapshot, }) @@ -251,14 +245,13 @@ fn bftree_parameters_from( as_input!(BfTreeFullPrecisionBuild); -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "bftree::")] pub(crate) struct BfTreeFullPrecisionBuild { build: IndexBuild, search_phase: SearchPhase, - #[serde(deserialize_with = "Deserialize::deserialize")] - vector_store_config: Option, - #[serde(deserialize_with = "Deserialize::deserialize")] - neighbor_store_config: Option, + vector_store_config: RequiredOption, + neighbor_store_config: RequiredOption, } impl BfTreeFullPrecisionBuild { @@ -291,19 +284,19 @@ impl BfTreeFullPrecisionBuild { &self.build, num_points, dim, - &self.vector_store_config, - &self.neighbor_store_config, - &None, + self.vector_store_config.as_ref(), + self.neighbor_store_config.as_ref(), + None, ) } pub(crate) fn validate(&mut self, checker: &mut Checker) -> Result<(), anyhow::Error> { self.build.validate(checker)?; self.search_phase.validate(checker)?; - if let Some(cfg) = &mut self.neighbor_store_config { + if let Some(cfg) = self.neighbor_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } - if let Some(cfg) = &mut self.vector_store_config { + if let Some(cfg) = self.vector_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } @@ -318,8 +311,8 @@ impl Example for BfTreeFullPrecisionBuild { Self { build, search_phase: SearchPhase::Topk(TopkSearchPhase::example()), - vector_store_config: None, - neighbor_store_config: None, + vector_store_config: RequiredOption::none(), + neighbor_store_config: RequiredOption::none(), } } } @@ -335,11 +328,11 @@ impl std::fmt::Display for BfTreeFullPrecisionBuild { writeln!(f)?; self.build.summarize_fields(f)?; - if let Some(ref cfg) = self.vector_store_config { + if let Some(cfg) = self.vector_store_config.as_ref() { writeln!(f, "\n Vector Store:")?; write!(f, "{}", cfg)?; } - if let Some(ref cfg) = self.neighbor_store_config { + if let Some(cfg) = self.neighbor_store_config.as_ref() { writeln!(f, "\n Neighbor Store:")?; write!(f, "{}", cfg)?; } @@ -350,15 +343,14 @@ impl std::fmt::Display for BfTreeFullPrecisionBuild { as_input!(BfTreeDynamicRun); -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "bftree::")] pub(crate) struct BfTreeDynamicRun { build: IndexBuild, search_phase: SearchPhase, runbook_params: DynamicRunbookParams, - #[serde(deserialize_with = "Deserialize::deserialize")] - vector_store_config: Option, - #[serde(deserialize_with = "Deserialize::deserialize")] - neighbor_store_config: Option, + vector_store_config: RequiredOption, + neighbor_store_config: RequiredOption, } impl BfTreeDynamicRun { @@ -395,20 +387,20 @@ impl BfTreeDynamicRun { &self.build, num_points, dim, - &self.vector_store_config, - &self.neighbor_store_config, - &None, + self.vector_store_config.as_ref(), + self.neighbor_store_config.as_ref(), + None, ) } pub(crate) fn validate(&mut self, checker: &mut Checker) -> Result<(), anyhow::Error> { self.build.validate(checker)?; self.search_phase.validate(checker)?; self.runbook_params.validate(checker)?; - if let Some(cfg) = &mut self.vector_store_config { + if let Some(cfg) = self.vector_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } - if let Some(cfg) = &mut self.neighbor_store_config { + if let Some(cfg) = self.neighbor_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } @@ -424,8 +416,8 @@ impl Example for BfTreeDynamicRun { build, search_phase: SearchPhase::Topk(TopkSearchPhase::example()), runbook_params: DynamicRunbookParams::example_immediate(), - vector_store_config: None, - neighbor_store_config: None, + vector_store_config: RequiredOption::none(), + neighbor_store_config: RequiredOption::none(), } } } @@ -441,11 +433,11 @@ impl std::fmt::Display for BfTreeDynamicRun { writeln!(f)?; self.build.summarize_fields(f)?; - if let Some(ref cfg) = self.vector_store_config { + if let Some(cfg) = self.vector_store_config.as_ref() { writeln!(f, "\n Vector Store:")?; write!(f, "{}", cfg)?; } - if let Some(ref cfg) = self.neighbor_store_config { + if let Some(cfg) = self.neighbor_store_config.as_ref() { writeln!(f, "\n Neighbor Store:")?; write!(f, "{}", cfg)?; } @@ -458,7 +450,8 @@ impl std::fmt::Display for BfTreeDynamicRun { as_input!(BfTreeSphericalBuild); -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "bftree::")] pub(crate) struct BfTreeSphericalBuild { build: IndexBuild, search_phase: SearchPhase, @@ -466,12 +459,9 @@ pub(crate) struct BfTreeSphericalBuild { transform_kind: exhaustive::TransformKind, num_bits: NonZeroUsize, pre_scale: Option, - #[serde(deserialize_with = "Deserialize::deserialize")] - vector_store_config: Option, - #[serde(deserialize_with = "Deserialize::deserialize")] - neighbor_store_config: Option, - #[serde(deserialize_with = "Deserialize::deserialize")] - quant_store_config: Option, + vector_store_config: RequiredOption, + neighbor_store_config: RequiredOption, + quant_store_config: RequiredOption, } impl BfTreeSphericalBuild { @@ -492,9 +482,9 @@ impl BfTreeSphericalBuild { &self.build, num_points, dim, - &self.vector_store_config, - &self.neighbor_store_config, - &self.quant_store_config, + self.vector_store_config.as_ref(), + self.neighbor_store_config.as_ref(), + self.quant_store_config.as_ref(), ) } @@ -529,15 +519,15 @@ impl BfTreeSphericalBuild { pub(crate) fn validate(&mut self, checker: &mut Checker) -> anyhow::Result<()> { self.build.validate(checker)?; self.search_phase.validate(checker)?; - if let Some(cfg) = &mut self.vector_store_config { + if let Some(cfg) = self.vector_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } - if let Some(cfg) = &mut self.neighbor_store_config { + if let Some(cfg) = self.neighbor_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } - if let Some(cfg) = &mut self.quant_store_config { + if let Some(cfg) = self.quant_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } @@ -560,9 +550,9 @@ impl Example for BfTreeSphericalBuild { transform_kind: exhaustive::TransformKind::Null, num_bits: NonZeroUsize::new(1).unwrap(), pre_scale: None, - vector_store_config: None, - neighbor_store_config: None, - quant_store_config: None, + vector_store_config: RequiredOption::none(), + neighbor_store_config: RequiredOption::none(), + quant_store_config: RequiredOption::none(), } } } @@ -581,15 +571,15 @@ impl std::fmt::Display for BfTreeSphericalBuild { writeln!(f)?; self.build.summarize_fields(f)?; - if let Some(ref cfg) = self.vector_store_config { + if let Some(cfg) = self.vector_store_config.as_ref() { writeln!(f, "\n Vector Store:")?; write!(f, "{}", cfg)?; } - if let Some(ref cfg) = self.neighbor_store_config { + if let Some(cfg) = self.neighbor_store_config.as_ref() { writeln!(f, "\n Neighbor Store:")?; write!(f, "{}", cfg)?; } - if let Some(ref cfg) = self.quant_store_config { + if let Some(cfg) = self.quant_store_config.as_ref() { writeln!(f, "\n Quant Store:")?; write!(f, "{}", cfg)?; } @@ -602,7 +592,8 @@ impl std::fmt::Display for BfTreeSphericalBuild { as_input!(BfTreeSphericalDynamicRun); -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "bftree::")] pub(crate) struct BfTreeSphericalDynamicRun { build: IndexBuild, search_phase: SearchPhase, @@ -611,12 +602,9 @@ pub(crate) struct BfTreeSphericalDynamicRun { transform_kind: exhaustive::TransformKind, num_bits: NonZeroUsize, pre_scale: Option, - #[serde(deserialize_with = "Deserialize::deserialize")] - vector_store_config: Option, - #[serde(deserialize_with = "Deserialize::deserialize")] - neighbor_store_config: Option, - #[serde(deserialize_with = "Deserialize::deserialize")] - quant_store_config: Option, + vector_store_config: RequiredOption, + neighbor_store_config: RequiredOption, + quant_store_config: RequiredOption, } impl BfTreeSphericalDynamicRun { @@ -669,9 +657,9 @@ impl BfTreeSphericalDynamicRun { &self.build, num_points, dim, - &self.vector_store_config, - &self.neighbor_store_config, - &self.quant_store_config, + self.vector_store_config.as_ref(), + self.neighbor_store_config.as_ref(), + self.quant_store_config.as_ref(), ) } @@ -679,15 +667,15 @@ impl BfTreeSphericalDynamicRun { self.build.validate(checker)?; self.search_phase.validate(checker)?; self.runbook_params.validate(checker)?; - if let Some(cfg) = &mut self.vector_store_config { + if let Some(cfg) = self.vector_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } - if let Some(cfg) = &mut self.neighbor_store_config { + if let Some(cfg) = self.neighbor_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } - if let Some(cfg) = &mut self.quant_store_config { + if let Some(cfg) = self.quant_store_config.as_mut() { cfg.fill_defaults(); cfg.validate()?; } @@ -711,9 +699,9 @@ impl Example for BfTreeSphericalDynamicRun { transform_kind: exhaustive::TransformKind::Null, num_bits: NonZeroUsize::new(1).unwrap(), pre_scale: None, - vector_store_config: None, - neighbor_store_config: None, - quant_store_config: None, + vector_store_config: RequiredOption::none(), + neighbor_store_config: RequiredOption::none(), + quant_store_config: RequiredOption::none(), } } } @@ -732,15 +720,15 @@ impl std::fmt::Display for BfTreeSphericalDynamicRun { writeln!(f)?; self.build.summarize_fields(f)?; - if let Some(ref cfg) = self.vector_store_config { + if let Some(cfg) = self.vector_store_config.as_ref() { writeln!(f, "\n Vector Store:")?; write!(f, "{}", cfg)?; } - if let Some(ref cfg) = self.neighbor_store_config { + if let Some(cfg) = self.neighbor_store_config.as_ref() { writeln!(f, "\n Neighbor Store:")?; write!(f, "{}", cfg)?; } - if let Some(ref cfg) = self.quant_store_config { + if let Some(cfg) = self.quant_store_config.as_ref() { writeln!(f, "\n Quant Store:")?; write!(f, "{}", cfg)?; } diff --git a/diskann-benchmark/src/inputs/disk.rs b/diskann-benchmark/src/inputs/disk.rs index 032665cd7a..8166bb592a 100644 --- a/diskann-benchmark/src/inputs/disk.rs +++ b/diskann-benchmark/src/inputs/disk.rs @@ -37,13 +37,13 @@ as_input!(DiskIndexOperation); // Input // /////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] pub(crate) struct DiskIndexOperation { pub(crate) source: DiskIndexSource, // either load or build pub(crate) search_phase: DiskSearchPhase, } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] #[serde(tag = "disk-index-source")] // Use tagged enums for JSON pub(crate) enum DiskIndexSource { Load(DiskIndexLoad), diff --git a/diskann-benchmark/src/inputs/exhaustive.rs b/diskann-benchmark/src/inputs/exhaustive.rs index 20583de85c..a2e30c3641 100644 --- a/diskann-benchmark/src/inputs/exhaustive.rs +++ b/diskann-benchmark/src/inputs/exhaustive.rs @@ -6,7 +6,7 @@ use std::num::NonZeroUsize; use anyhow::{anyhow, Context}; -use diskann_benchmark_runner::{files::InputFile, utils::datatype::DataType, Checker}; +use diskann_benchmark_runner::{files::InputFile, utils::datatype::DataType, Checker, Reflect}; use serde::{Deserialize, Serialize}; use crate::{ @@ -33,7 +33,8 @@ as_input!(MinMax); // Search // //////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "exhaustive::")] pub(crate) struct SearchValues { pub(crate) recall_k: Vec, pub(crate) recall_n: Vec, @@ -85,7 +86,8 @@ impl SearchValues { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "exhaustive::")] pub(crate) struct SearchPhase { pub(crate) queries: InputFile, pub(crate) groundtruth: InputFile, @@ -124,8 +126,9 @@ impl Example for SearchPhase { //////////////////////////////// // Transforms related methods // /////////////////////////////// -#[derive(Debug, Clone, Copy, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, Serialize, Deserialize, Reflect)] #[serde(rename_all = "snake_case")] +#[reflect(prefix = "transform::")] pub(crate) enum TargetDim { Same, Natural, @@ -152,8 +155,9 @@ impl From for diskann_quantization::algorithms::transforms::TargetDim } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] #[serde(rename_all = "snake_case")] +#[reflect(prefix = "transform::")] pub(crate) enum TransformKind { PaddingHadamard(TargetDim), RandomRotation(TargetDim), @@ -200,7 +204,8 @@ impl From<&TransformKind> for diskann_quantization::algorithms::transforms::Tran // Product Quantization Methods // ////////////////////////////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "product::")] pub(crate) struct Product { pub(crate) data: InputFile, pub(crate) data_type: DataType, @@ -272,8 +277,9 @@ impl std::fmt::Display for Product { // Spherical-quantization-based methods // ////////////////////////////////////////// -#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "snake_case")] +#[reflect(prefix = "spherical::")] pub(crate) enum SphericalQuery { SameAsData, FourBitTransposed, @@ -333,8 +339,9 @@ pub(super) fn check_compatibility(num_bits: usize, query: SphericalQuery) -> any } } -#[derive(Debug, Clone, Copy, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, Serialize, Deserialize, Reflect)] #[serde(rename_all = "snake_case")] +#[reflect(prefix = "spherical::")] pub(crate) enum PreScale { None, Some(f32), @@ -378,7 +385,8 @@ impl PreScale { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "exhaustive::spherical::")] pub(crate) struct Spherical { pub(crate) data: InputFile, pub(crate) data_type: DataType, @@ -461,8 +469,9 @@ impl std::fmt::Display for Spherical { // MinMax-quantization-based methods // /////////////////////////////////////// -#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "snake_case")] +#[reflect(prefix = "minmax::")] pub(crate) enum MinMaxQuery { SameAsData, FullPrecision, @@ -480,7 +489,8 @@ impl std::fmt::Display for MinMaxQuery { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "minmax::exhaustive::")] pub(crate) struct MinMax { pub(crate) data: InputFile, pub(crate) data_type: DataType, diff --git a/diskann-benchmark/src/inputs/filters.rs b/diskann-benchmark/src/inputs/filters.rs index 942c6da12b..40b9685139 100644 --- a/diskann-benchmark/src/inputs/filters.rs +++ b/diskann-benchmark/src/inputs/filters.rs @@ -3,7 +3,7 @@ * Licensed under the MIT license. */ -use diskann_benchmark_runner::{files::InputFile, Checker}; +use diskann_benchmark_runner::{files::InputFile, Checker, Reflect}; use serde::{Deserialize, Serialize}; use crate::inputs::{as_input, Example}; @@ -18,7 +18,7 @@ as_input!(MetadataIndexBuild); // Metadata-only Index Build // /////////////////////////////// -#[derive(Default, Debug, Serialize, Deserialize, Clone, Copy)] +#[derive(Default, Debug, Serialize, Deserialize, Reflect, Clone, Copy)] pub(crate) enum InvertedIndexKind { #[serde(rename = "bftree")] #[default] @@ -32,22 +32,20 @@ impl std::fmt::Display for InvertedIndexKind { } } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] pub(crate) struct FilterParams { pub(crate) query_predicates: InputFile, pub(crate) data_labels: InputFile, } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] pub(crate) struct MetadataIndexBuild { /// Filter parameters describing predicate and label file locations. The /// actual label file used to build the inverted index is taken from /// `filter_params.data_labels`. pub(crate) filter_params: FilterParams, - /// Which inverted-index implementation to use when building/evaluating - /// bitmap filters. If omitted in input files, defaults to `fast`. - #[serde(default)] + /// Which inverted-index implementation to use when building/evaluating bitmap filters. pub(crate) inverted_index_type: InvertedIndexKind, } diff --git a/diskann-benchmark/src/inputs/flat.rs b/diskann-benchmark/src/inputs/flat.rs index d25e77cf76..b998aceae8 100644 --- a/diskann-benchmark/src/inputs/flat.rs +++ b/diskann-benchmark/src/inputs/flat.rs @@ -6,7 +6,7 @@ use std::num::NonZeroUsize; use anyhow::Context; -use diskann_benchmark_runner::{files::InputFile, utils::datatype::DataType, Checker}; +use diskann_benchmark_runner::{files::InputFile, utils::datatype::DataType, Checker, Reflect}; use serde::{Deserialize, Serialize}; use crate::{ @@ -25,7 +25,7 @@ as_input!(FlatSearch); /////////// /// Input specification for a flat-index (brute-force kNN) benchmark. -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] pub(crate) struct FlatSearch { /// Path to the dataset vectors (`.bin` format). pub(crate) data: InputFile, @@ -91,7 +91,7 @@ impl Example for FlatSearch { /////////////////// /// Parameters controlling the search phase of a flat benchmark. -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] pub(crate) struct SearchPhase { /// Path to the query vectors (`.bin` format). pub(crate) queries: InputFile, diff --git a/diskann-benchmark/src/inputs/graph_index.rs b/diskann-benchmark/src/inputs/graph_index.rs index c7ed6c1410..d9196caaab 100644 --- a/diskann-benchmark/src/inputs/graph_index.rs +++ b/diskann-benchmark/src/inputs/graph_index.rs @@ -7,11 +7,13 @@ use std::num::{NonZero, NonZeroU32, NonZeroUsize}; use anyhow::{anyhow, Context}; use diskann::{ - graph::{self, config, search::Range, RangeSearchError, StartPointStrategy}, + graph::{self, config, search::Range, RangeSearchError}, utils::IntoUsize, }; use diskann_benchmark_core::streaming::executors::bigann; -use diskann_benchmark_runner::{files::InputFile, utils::datatype::DataType, Checker}; +use diskann_benchmark_runner::{ + files::InputFile, utils::datatype::DataType, utils::RequiredOption, Checker, Reflect, +}; use diskann_providers::{ model::{ configuration::IndexConfiguration, @@ -41,7 +43,8 @@ as_input!(DynamicIndexRun); // Search // //////////// -#[derive(Debug, Serialize, Deserialize, Clone)] +#[derive(Debug, Serialize, Deserialize, Reflect, Clone)] +#[reflect(prefix = "graph::")] pub(crate) struct GraphSearch { pub(crate) search_n: usize, pub(crate) search_l: Vec, @@ -65,7 +68,8 @@ impl GraphSearch { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct GraphRangeSearch { pub(crate) initial_search_l: Vec, pub(crate) radius: f32, @@ -102,7 +106,8 @@ impl GraphRangeSearch { } } -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct TopkSearchPhase { pub(crate) queries: InputFile, pub(crate) groundtruth: InputFile, @@ -156,7 +161,8 @@ impl Example for TopkSearchPhase { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct RangeSearchPhase { pub(crate) queries: InputFile, pub(crate) groundtruth: InputFile, @@ -166,7 +172,8 @@ pub(crate) struct RangeSearchPhase { pub(crate) runs: Vec, } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct FilteredRangeSearchPhase { pub(crate) queries: InputFile, pub(crate) query_predicates: InputFile, @@ -206,7 +213,8 @@ impl RangeSearchPhase { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct BetaSearchPhase { pub(crate) queries: InputFile, pub(crate) query_predicates: InputFile, @@ -242,7 +250,8 @@ impl BetaSearchPhase { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct MultihopFilterSearchPhase { pub(crate) queries: InputFile, pub(crate) query_predicates: InputFile, @@ -269,7 +278,8 @@ impl MultihopFilterSearchPhase { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct AdaptiveL { pub(crate) sample_count: NonZeroUsize, pub(crate) scale_factor: f64, @@ -283,7 +293,8 @@ impl AdaptiveL { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct InlineFilterSearchPhase { pub(crate) queries: InputFile, pub(crate) query_predicates: InputFile, @@ -292,8 +303,7 @@ pub(crate) struct InlineFilterSearchPhase { pub(crate) data_labels: InputFile, pub(crate) num_threads: Vec, pub(crate) runs: Vec, - #[serde(deserialize_with = "Deserialize::deserialize")] - pub(crate) adaptive_l: Option, + pub(crate) adaptive_l: RequiredOption, } impl InlineFilterSearchPhase { @@ -306,7 +316,7 @@ impl InlineFilterSearchPhase { run.validate(checker) .with_context(|| format!("search run {}", i))?; } - if let Some(ref adaptive_l) = self.adaptive_l { + if let Some(adaptive_l) = self.adaptive_l.as_ref() { adaptive_l.validate(checker)?; } @@ -314,7 +324,7 @@ impl InlineFilterSearchPhase { } pub(crate) fn adaptive_l(&self) -> Result, anyhow::Error> { - if let Some(ref adaptive_l) = self.adaptive_l { + if let Some(adaptive_l) = self.adaptive_l.as_ref() { let adaptive_l = graph::search::AdaptiveL::new( adaptive_l.sample_count.into(), adaptive_l.scale_factor, @@ -327,8 +337,9 @@ impl InlineFilterSearchPhase { } /// A one-to-one correspondence with [`diskann::graph::config::IntraBatchCandidates`]. -#[derive(Debug, Clone, Copy, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] +#[reflect(prefix = "graph::")] pub(crate) enum IntraBatchCandidates { /// No intra-batch candidates will be considered. None, @@ -359,7 +370,8 @@ impl From for config::IntraBatchCandidates { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct MultiInsert { pub(crate) batch_size: NonZeroUsize, pub(crate) batch_parallelism: NonZeroUsize, @@ -379,7 +391,8 @@ impl Example for MultiInsert { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct TopkDeterminantDiversityPhase { pub(crate) queries: InputFile, pub(crate) groundtruth: InputFile, @@ -426,8 +439,9 @@ impl Example for TopkDeterminantDiversityPhase { } } -#[derive(Debug, Deserialize, Serialize)] +#[derive(Debug, Deserialize, Serialize, Reflect)] #[serde(tag = "search-type", rename_all = "kebab-case")] +#[reflect(prefix = "graph::")] pub(crate) enum SearchPhase { Topk(TopkSearchPhase), Range(RangeSearchPhase), @@ -596,7 +610,8 @@ impl std::fmt::Display for SearchPhaseKind { // Build - Full Precision // //////////////////////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct IndexLoad { pub(crate) data_type: DataType, pub(crate) distance: SimilarityMeasure, @@ -673,17 +688,18 @@ impl std::fmt::Display for IndexLoad { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct InsertRetry { num_insert_attempts: NonZeroU32, retry_threshold: f32, saturate_inserts: bool, } -#[derive(Debug, Serialize, Deserialize, Clone, Copy, PartialEq)] -#[serde(remote = "StartPointStrategy")] +#[derive(Debug, Serialize, Deserialize, Reflect, Clone, Copy, PartialEq)] #[serde(rename_all = "snake_case")] -pub enum StartPointStrategyRef { +#[reflect(prefix = "graph::")] +pub enum StartPointStrategy { /// Randomly select vector(s) with given norm as starting points with seed provided. /// Requires the norm (f32), number of samples (usize), and random seed (u64) to be provided. RandomVectors { @@ -707,7 +723,32 @@ pub enum StartPointStrategyRef { FirstVector, } -#[derive(Debug, Serialize, Deserialize)] +impl StartPointStrategy { + fn to_diskann(&self) -> graph::StartPointStrategy { + match *self { + Self::RandomVectors { + norm, + nsamples, + seed, + } => graph::StartPointStrategy::RandomVectors { + norm, + nsamples, + seed, + }, + Self::RandomSamples { nsamples, seed } => { + graph::StartPointStrategy::RandomSamples { nsamples, seed } + } + Self::Medoid => graph::StartPointStrategy::Medoid, + Self::LatinHyperCube { nsamples, seed } => { + graph::StartPointStrategy::LatinHyperCube { nsamples, seed } + } + Self::FirstVector => graph::StartPointStrategy::FirstVector, + } + } +} + +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct IndexBuild { data_type: DataType, data: InputFile, @@ -715,7 +756,6 @@ pub(crate) struct IndexBuild { max_degree: usize, l_build: usize, insert_retry: Option, - #[serde(with = "StartPointStrategyRef")] start_point_strategy: StartPointStrategy, alpha: f32, backedge_ratio: f32, @@ -789,7 +829,7 @@ impl IndexBuild { ) -> DefaultProviderParameters { DefaultProviderParameters { max_points: num_points, - frozen_points: NonZero::new(self.start_point_strategy.count()).unwrap(), + frozen_points: NonZero::new(self.start_point_strategy().count()).unwrap(), metric: self.distance.into(), dim, max_degree: self.exact_max_degree() as u32, @@ -804,7 +844,7 @@ impl IndexBuild { write_field!(f, "max degree", self.max_degree)?; write_field!(f, "L-build", self.l_build)?; write_field!(f, "alpha", self.alpha)?; - write_field!(f, "start point strategy", self.start_point_strategy)?; + write_field!(f, "start point strategy", self.start_point_strategy())?; write_field!(f, "backedge ratio", self.backedge_ratio)?; match &self.multi_insert { None => write_field!(f, "Using Multi Insert", "NO")?, @@ -814,7 +854,7 @@ impl IndexBuild { write_field!(f, "Intra Batch Candidates", mi.intra_batch_candidates)?; } } - write_field!(f, "start_point_strategy", self.start_point_strategy)?; + write_field!(f, "start_point_strategy", self.start_point_strategy())?; write_field!(f, "build threads", self.num_threads)?; match &self.save_path { None => write_field!(f, "Save Path", "None")?, @@ -864,8 +904,8 @@ impl IndexBuild { &self.data } - pub(crate) fn start_point_strategy(&self) -> &StartPointStrategy { - &self.start_point_strategy + pub(crate) fn start_point_strategy(&self) -> graph::StartPointStrategy { + self.start_point_strategy.to_diskann() } pub(crate) fn multi_insert(&self) -> Option<&MultiInsert> { @@ -906,8 +946,9 @@ impl std::fmt::Display for IndexBuild { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] #[serde(tag = "index-source")] // Use tagged enums for JSON +#[reflect(prefix = "graph::")] pub enum IndexSource { Load(IndexLoad), Build(IndexBuild), @@ -936,7 +977,8 @@ impl IndexSource { } } -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct IndexOperation { pub(crate) source: IndexSource, // either load or build pub(crate) search_phase: SearchPhase, @@ -978,7 +1020,8 @@ impl std::fmt::Display for IndexOperation { // Graph Index Build PQ // ////////////////////////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct IndexPQOperation { pub(crate) index_operation: IndexOperation, // either load or build pub(crate) num_pq_chunks: usize, @@ -1068,7 +1111,8 @@ impl std::fmt::Display for IndexPQOperation { // Graph Index Build SQ // ////////////////////////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct IndexSQOperation { pub(crate) index_operation: IndexOperation, pub(crate) num_bits: usize, @@ -1158,7 +1202,8 @@ impl std::fmt::Display for IndexSQOperation { // Graph Index Build Spherical // ///////////////////////////////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct SphericalQuantBuild { pub(crate) build: IndexBuild, // spherical does not support saving and loading pub(crate) search_phase: SearchPhase, @@ -1276,8 +1321,9 @@ impl std::fmt::Display for SphericalQuantBuild { // Dynamic Runbook Params // //////////////////////////// -#[derive(Copy, Clone, Debug, serde::Serialize, serde::Deserialize)] +#[derive(Copy, Clone, Debug, serde::Serialize, serde::Deserialize, Reflect)] #[serde(tag = "method", content = "params")] +#[reflect(prefix = "graph::")] pub enum InplaceDeleteMethod { #[serde(rename = "visited_and_top_k")] VisitedAndTopK { k_value: usize, l_value: usize }, @@ -1300,7 +1346,8 @@ impl From for graph::InplaceDeleteMethod { } /// Runbook loading and phase type definitions are in utils.datafiles -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct DynamicRunbookParams { pub(crate) runbook_path: InputFile, pub(crate) dataset_name: String, @@ -1309,9 +1356,7 @@ pub(crate) struct DynamicRunbookParams { pub(crate) ip_delete_num_to_replace: usize, /// Threshold for deferred consolidation. Required for soft-delete providers (inmem). /// Hard-delete providers (bf-tree) ignore this field. - #[serde(default, skip_serializing_if = "Option::is_none")] - pub(crate) consolidate_threshold: Option, - #[serde(skip)] + pub(crate) consolidate_threshold: RequiredOption, pub(crate) resolved_gt_directory: Option, } @@ -1324,7 +1369,7 @@ impl DynamicRunbookParams { self.runbook_path.resolve(checker)?; // Validate consolidate_threshold if provided - if let Some(threshold) = self.consolidate_threshold { + if let Some(&threshold) = self.consolidate_threshold.as_ref() { if threshold <= 0.0 { return Err(anyhow::anyhow!( "consolidate_threshold must be greater than 0, but got {}", @@ -1366,7 +1411,7 @@ impl Example for DynamicRunbookParams { l_value: 64, }, ip_delete_num_to_replace: 3, - consolidate_threshold: Some(0.2), + consolidate_threshold: RequiredOption::some(0.2), resolved_gt_directory: None, } } @@ -1377,7 +1422,7 @@ impl DynamicRunbookParams { #[cfg(feature = "bftree")] pub(crate) fn example_immediate() -> Self { Self { - consolidate_threshold: None, + consolidate_threshold: RequiredOption::none(), ..Self::example() } } @@ -1410,7 +1455,7 @@ impl std::fmt::Display for DynamicRunbookParams { } } write_field!(f, "IP Delete Num to Replace", self.ip_delete_num_to_replace)?; - if let Some(threshold) = self.consolidate_threshold { + if let Some(threshold) = self.consolidate_threshold.as_ref() { write_field!(f, "Consolidate Threshold", threshold)?; } @@ -1422,7 +1467,8 @@ impl std::fmt::Display for DynamicRunbookParams { // Graph Index Dynamic // /////////////////////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "graph::")] pub(crate) struct DynamicIndexRun { pub(crate) build: IndexBuild, pub(crate) search_phase: SearchPhase, diff --git a/diskann-benchmark/src/inputs/mod.rs b/diskann-benchmark/src/inputs/mod.rs index 978d93cb31..75f5512ad6 100644 --- a/diskann-benchmark/src/inputs/mod.rs +++ b/diskann-benchmark/src/inputs/mod.rs @@ -3,7 +3,7 @@ * Licensed under the MIT license. */ -pub(crate) mod disk; +// pub(crate) mod disk; pub(crate) mod exhaustive; pub(crate) mod filters; pub(crate) mod flat; diff --git a/diskann-benchmark/src/inputs/multi_vector.rs b/diskann-benchmark/src/inputs/multi_vector.rs index c74f9d232f..d362d9ce3b 100644 --- a/diskann-benchmark/src/inputs/multi_vector.rs +++ b/diskann-benchmark/src/inputs/multi_vector.rs @@ -5,7 +5,7 @@ use std::num::NonZeroUsize; -use diskann_benchmark_runner::{utils::datatype::DataType, Checker, Input}; +use diskann_benchmark_runner::{utils::datatype::DataType, Checker, Input, Reflect}; use diskann_quantization::multi_vector::MaxSimIsa; use serde::{Deserialize, Serialize}; @@ -16,7 +16,7 @@ use serde::{Deserialize, Serialize}; /// JSON-facing shadow of [`MaxSimIsa`]. The library's enum is deliberately /// serde-free; this owns the kebab-case JSON shape and converts via `From`. /// Stays variant-for-variant in sync with `MaxSimIsa` manually. -#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] #[non_exhaustive] pub(crate) enum BenchIsa { @@ -60,7 +60,7 @@ impl From for MaxSimIsa { } /// One benchmark configuration: a single shape measurement. -#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Reflect)] pub(crate) struct Run { pub(crate) num_query_vectors: NonZeroUsize, pub(crate) num_doc_vectors: NonZeroUsize, @@ -74,7 +74,7 @@ pub(crate) struct Run { /////////////////////// /// A complete multi-vector benchmark job. -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] pub(crate) struct MultiVectorOp { pub(crate) element_type: DataType, pub(crate) isa: BenchIsa, diff --git a/diskann-benchmark/src/main.rs b/diskann-benchmark/src/main.rs index 13cf1ec31f..331dc6e93f 100644 --- a/diskann-benchmark/src/main.rs +++ b/diskann-benchmark/src/main.rs @@ -5,7 +5,7 @@ //! Command-line benchmarks for DiskANN. -mod disk_index; +// mod disk_index; mod exhaustive; mod filters; mod flat; @@ -49,15 +49,19 @@ impl Cli { fn run(&self, output: &mut dyn runner::Output) -> anyhow::Result<()> { self.check_target(output)?; + let now = std::time::Instant::now(); + // Collect benchmarks. let mut registry = runner::Registry::new(); exhaustive::register_benchmarks(&mut registry)?; - disk_index::register_benchmarks(&mut registry)?; + // disk_index::register_benchmarks(&mut registry)?; flat::register_benchmarks(&mut registry)?; index::register_benchmarks(&mut registry)?; filters::register_benchmarks(&mut registry)?; multi_vector::register_benchmarks(&mut registry)?; + println!("registration took {}us", now.elapsed().as_micros()); + self.app.run(®istry, output) } diff --git a/diskann-benchmark/src/multi_vector/driver.rs b/diskann-benchmark/src/multi_vector/driver.rs index e69c708451..7befcf738b 100644 --- a/diskann-benchmark/src/multi_vector/driver.rs +++ b/diskann-benchmark/src/multi_vector/driver.rs @@ -13,6 +13,7 @@ use diskann_benchmark_runner::{ percentiles, MicroSeconds, }, Checker, Input, + Reflect, }; use diskann_quantization::multi_vector::{Mat, MatRef, MaxSimKernel, Overflow, Standard}; use rand::{ @@ -30,7 +31,8 @@ use crate::utils::DisplayWrapper; ////////////////////// /// Tolerance thresholds for multi-vector benchmark regression detection. -#[derive(Debug, Clone, Copy, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "multi_vector::")] pub(super) struct MultiVectorTolerance { pub(super) min_time_regression: NonNegativeFinite, } diff --git a/diskann-benchmark/src/utils/mod.rs b/diskann-benchmark/src/utils/mod.rs index cd8e6510cf..71d6ec671d 100644 --- a/diskann-benchmark/src/utils/mod.rs +++ b/diskann-benchmark/src/utils/mod.rs @@ -6,6 +6,7 @@ use diskann_benchmark_runner::{ benchmark::Score, utils::datatype::{AsDataType, DataType}, + Reflect, }; use serde::{Deserialize, Serialize}; @@ -26,8 +27,9 @@ where } } -#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "snake_case")] +#[reflect(prefix = "benchmark::")] pub(crate) enum SimilarityMeasure { SquaredL2, InnerProduct, From 78a2698f37f71a5216022e5db2cc7208c6d5fc82 Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Wed, 30 Sep 2026 18:08:20 -0700 Subject: [PATCH 20/21] Finish porting. --- diskann-benchmark/src/inputs/bftree.rs | 10 ++++- diskann-benchmark/src/multi_vector/driver.rs | 3 +- diskann-inmem/integration/index/runner.rs | 40 +++++++++++-------- diskann-inmem/integration/store/checked.rs | 3 +- diskann-inmem/integration/store/intrusive.rs | 3 +- diskann-inmem/integration/store/mod.rs | 5 ++- diskann-inmem/integration/support/datatype.rs | 3 +- .../integration/support/tolerance.rs | 4 +- 8 files changed, 44 insertions(+), 27 deletions(-) diff --git a/diskann-benchmark/src/inputs/bftree.rs b/diskann-benchmark/src/inputs/bftree.rs index 22dfc870e3..b214a6b68d 100644 --- a/diskann-benchmark/src/inputs/bftree.rs +++ b/diskann-benchmark/src/inputs/bftree.rs @@ -10,7 +10,10 @@ use crate::inputs::{ write_field, Example, PRINT_WIDTH, }; use diskann::graph::config; -use diskann_benchmark_runner::{utils::{RequiredOption, datatype::DataType}, Checker, Reflect}; +use diskann_benchmark_runner::{ + utils::{datatype::DataType, RequiredOption}, + Checker, Reflect, +}; use diskann_bftree::BfTreeProviderParameters; use serde::{Deserialize, Serialize}; @@ -237,7 +240,10 @@ fn bftree_parameters_from( .cloned() .unwrap_or_default() .into_config(), - quant_vector_provider_config: quant_store_config.cloned().unwrap_or_default().into_config(), + quant_vector_provider_config: quant_store_config + .cloned() + .unwrap_or_default() + .into_config(), graph_params: None, use_snapshot, }) diff --git a/diskann-benchmark/src/multi_vector/driver.rs b/diskann-benchmark/src/multi_vector/driver.rs index 7befcf738b..e5deb79d31 100644 --- a/diskann-benchmark/src/multi_vector/driver.rs +++ b/diskann-benchmark/src/multi_vector/driver.rs @@ -12,8 +12,7 @@ use diskann_benchmark_runner::{ num::{relative_change, NonNegativeFinite}, percentiles, MicroSeconds, }, - Checker, Input, - Reflect, + Checker, Input, Reflect, }; use diskann_quantization::multi_vector::{Mat, MatRef, MaxSimKernel, Overflow, Standard}; use rand::{ diff --git a/diskann-inmem/integration/index/runner.rs b/diskann-inmem/integration/index/runner.rs index 97e711a5d1..6d593c5c6b 100644 --- a/diskann-inmem/integration/index/runner.rs +++ b/diskann-inmem/integration/index/runner.rs @@ -11,7 +11,8 @@ use diskann_benchmark_runner::{ Checker, Checkpoint, Output, Registry, RegistryError, benchmark::{MatchContext, PassFail, Regression, Score}, files::InputFile, - utils::fmt::Indent, + utils::{RequiredOption, fmt::Indent}, + Reflect, }; use diskann_utils::views::Matrix; use diskann_vector::distance::Metric; @@ -43,8 +44,9 @@ mod dto { use serde::{Deserialize, Serialize}; - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] + #[reflect(prefix = "graph::")] pub(super) enum SerdeMetric { L2, InnerProduct, @@ -73,8 +75,9 @@ mod dto { } } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] + #[reflect(prefix = "graph::")] pub(super) enum Preprocess { Halve, Floor, @@ -98,7 +101,8 @@ mod dto { } } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "graph::")] pub(super) struct Data { pub(super) data: InputFile, pub(super) queries: InputFile, @@ -108,12 +112,14 @@ mod dto { pub(super) preprocess: Vec, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "graph::")] pub(super) enum Representation { FullPrecision { data_type: DataType }, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "graph::")] pub(super) struct Build { pub(super) pruned_degree: usize, pub(super) max_degree: usize, @@ -121,20 +127,22 @@ mod dto { pub(super) alpha: f32, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "graph::")] pub(super) struct KnnSearch { pub(super) knn: usize, pub(super) search_l: usize, - #[serde(deserialize_with = "Deserialize::deserialize")] - pub(super) beam_width: Option, + pub(super) beam_width: RequiredOption, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "graph::")] pub(super) struct Search { pub(super) knn: Vec, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] + #[reflect(prefix = "graph::")] pub(super) struct Test { pub(super) data: Data, pub(super) representation: Representation, @@ -296,7 +304,7 @@ struct Search { impl Search { fn from_raw(raw: dto::Search) -> anyhow::Result { fn make_knn(raw: &dto::KnnSearch) -> anyhow::Result<(usize, Knn)> { - Ok((raw.knn, Knn::new(raw.search_l, raw.beam_width)?)) + Ok((raw.knn, Knn::new(raw.search_l, raw.beam_width.as_ref().copied())?)) } Ok(Self { @@ -313,7 +321,7 @@ impl Search { dto::KnnSearch { knn: *k, search_l: knn.l_value().get(), - beam_width: Some(knn.beam_width().get()), + beam_width: RequiredOption::some(knn.beam_width().get()), } } @@ -471,17 +479,17 @@ impl diskann_benchmark_runner::Input for Test { dto::KnnSearch { knn: 10, search_l: 50, - beam_width: None, + beam_width: RequiredOption::none(), }, dto::KnnSearch { knn: 10, search_l: 50, - beam_width: Some(3), + beam_width: RequiredOption::some(3), }, dto::KnnSearch { knn: 20, search_l: 100, - beam_width: Some(3), + beam_width: RequiredOption::some(3), }, ], }, diff --git a/diskann-inmem/integration/store/checked.rs b/diskann-inmem/integration/store/checked.rs index 9a139349ad..1ece858429 100644 --- a/diskann-inmem/integration/store/checked.rs +++ b/diskann-inmem/integration/store/checked.rs @@ -14,7 +14,8 @@ pub(super) fn register(registry: &mut dbr::Registry) -> Result<(), dbr::Registry } /// Configuration for a [`Stress`] run. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, dbr::Reflect)] +#[reflect(prefix = "store::checked::")] struct Input { /// Shared stress test setup. setup: super::Setup, diff --git a/diskann-inmem/integration/store/intrusive.rs b/diskann-inmem/integration/store/intrusive.rs index 8b9c4b8433..ae42b5f2d5 100644 --- a/diskann-inmem/integration/store/intrusive.rs +++ b/diskann-inmem/integration/store/intrusive.rs @@ -14,7 +14,8 @@ pub(super) fn register(registry: &mut dbr::Registry) -> Result<(), dbr::Registry } /// Configuration for a [`Stress`] run. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, dbr::Reflect)] +#[reflect(prefix = "store::intrusive::")] struct Input { /// Shared stress test setup. setup: super::Setup, diff --git a/diskann-inmem/integration/store/mod.rs b/diskann-inmem/integration/store/mod.rs index ade8bd370b..991b524977 100644 --- a/diskann-inmem/integration/store/mod.rs +++ b/diskann-inmem/integration/store/mod.rs @@ -27,7 +27,7 @@ use std::{ time::{Duration, Instant}, }; -use diskann_benchmark_runner::{Registry, RegistryError, utils::fmt::KeyValue}; +use diskann_benchmark_runner::{Registry, RegistryError, Reflect, utils::fmt::KeyValue}; use rand::{Rng, SeedableRng, distr::Uniform, rngs::StdRng}; use serde::{Deserialize, Serialize}; @@ -56,7 +56,8 @@ pub(super) fn register(registry: &mut Registry) -> Result<(), RegistryError> { // Input // /////////// -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "store::")] pub struct Setup { /// Number of reader threads. Must be below `epoch_guard_slots`. readers: usize, diff --git a/diskann-inmem/integration/support/datatype.rs b/diskann-inmem/integration/support/datatype.rs index 34a60190aa..77eb57966f 100644 --- a/diskann-inmem/integration/support/datatype.rs +++ b/diskann-inmem/integration/support/datatype.rs @@ -8,6 +8,7 @@ use diskann_utils::{ views::{Matrix, MatrixView, MutMatrixView}, }; use diskann_wide::{cast_f16_to_f32, cast_f32_to_f16}; +use diskann_benchmark_runner::Reflect; use half::f16; use serde::{Deserialize, Serialize}; use thiserror::Error; @@ -16,7 +17,7 @@ use thiserror::Error; // DataType // ////////////// -#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] pub(crate) enum DataType { F32, diff --git a/diskann-inmem/integration/support/tolerance.rs b/diskann-inmem/integration/support/tolerance.rs index 7f8efa8709..2f6555325d 100644 --- a/diskann-inmem/integration/support/tolerance.rs +++ b/diskann-inmem/integration/support/tolerance.rs @@ -3,12 +3,12 @@ * Licensed under the MIT license. */ -use diskann_benchmark_runner::{Checker, Input}; +use diskann_benchmark_runner::{Checker, Input, Reflect}; use serde::{Deserialize, Serialize}; /// A tolerance [`Input`] for [`diskann_benchmark_runner::benchmark::Regression`]s that /// do not need any external tolerances. -#[derive(Debug, Clone, Copy, Serialize, Deserialize)] +#[derive(Debug, Clone, Copy, Serialize, Deserialize, Reflect)] pub(crate) struct Empty; impl Input for Empty { From 3036872aa0a15a7e40551d581b37a5d84eb1cecf Mon Sep 17 00:00:00 2001 From: Mark Hildebrand Date: Wed, 30 Sep 2026 18:12:17 -0700 Subject: [PATCH 21/21] Remove temp files. --- diskann-benchmark-runner/Cargo.toml | 4 ---- diskann-benchmark-runner/temp/main.rs | 18 ------------------ 2 files changed, 22 deletions(-) delete mode 100644 diskann-benchmark-runner/temp/main.rs diff --git a/diskann-benchmark-runner/Cargo.toml b/diskann-benchmark-runner/Cargo.toml index 331cfc6e4b..19d97a8805 100644 --- a/diskann-benchmark-runner/Cargo.toml +++ b/diskann-benchmark-runner/Cargo.toml @@ -41,7 +41,3 @@ ux-tools = [] name = "dev" path = "dev/main.rs" required-features = ["test-app"] - -[[bin]] -name = "temp" -path = "temp/main.rs" diff --git a/diskann-benchmark-runner/temp/main.rs b/diskann-benchmark-runner/temp/main.rs deleted file mode 100644 index d0aea2a292..0000000000 --- a/diskann-benchmark-runner/temp/main.rs +++ /dev/null @@ -1,18 +0,0 @@ -/* - * Copyright (c) Microsoft Corporation. - * Licensed under the MIT license. - */ - -//! Development CLI for exercising the benchmark runner with its test registry. - -use diskann_benchmark_runner::{reflect, Reflect, Reflection}; - -fn main() -> anyhow::Result<()> { - // println!("{}", Reflection::new::().render()); - - // println!("{}", Reflection::new::().render()); - - // println!("{}", Reflection::new::().render()); - - Ok(()) -}