diff --git a/Cargo.lock b/Cargo.lock index bc9903c01c..7eb0cec63d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -519,7 +519,9 @@ version = "0.59.0" dependencies = [ "anyhow", "clap", + "diskann-benchmark-runner-derive", "half", + "hashbrown 0.16.1", "indicatif", "serde", "serde_json", @@ -527,6 +529,17 @@ dependencies = [ "thiserror 2.0.17", ] +[[package]] +name = "diskann-benchmark-runner-derive" +version = "0.59.0" +dependencies = [ + "diskann-benchmark-runner", + "proc-macro2", + "quote", + "syn 2.0.117", + "trybuild", +] + [[package]] name = "diskann-benchmark-simd" version = "0.59.0" diff --git a/Cargo.toml b/Cargo.toml index 357674b55e..22c221748c 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -27,6 +27,7 @@ members = [ "diskann-record", "diskann-tools", "diskann-bftree", + "diskann-benchmark-runner-derive", ] default-members = [ @@ -66,11 +67,11 @@ diskann-inmem = { path = "diskann-inmem", default-features = false, version = "0 diskann-disk = { path = "diskann-disk", version = "0.59.0" } diskann-label-filter = { path = "diskann-label-filter", version = "0.59.0" } # Infra +diskann-benchmark-runner-derive = { path = "diskann-benchmark-runner-derive", version = "0.59.0" } diskann-benchmark-runner = { path = "diskann-benchmark-runner", version = "0.59.0" } 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 new file mode 100644 index 0000000000..8d66decc3a --- /dev/null +++ b/diskann-benchmark-runner-derive/Cargo.toml @@ -0,0 +1,22 @@ +[package] +name = "diskann-benchmark-runner-derive" +version.workspace = true +description.workspace = true +authors.workspace = true +license.workspace = true +edition = "2024" + +[lib] +proc-macro = true + +[dependencies] +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 new file mode 100644 index 0000000000..8b965f78f5 --- /dev/null +++ b/diskann-benchmark-runner-derive/SERDE_TODO.md @@ -0,0 +1,92 @@ +# 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 + +- [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 + `xml_http_request`. + - For fields, `snake_case` is identity and `kebab-case` replaces underscores. + - 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 + +- [X] Update enum compatibility coverage to expect renamed variants such as `"unit"` rather than + `"Unit"`. +- [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: + - [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 + +- [X] Fix the `generate_type_name_body` doctest by returning the final + `f.write_str(">")` result instead of discarding it with a semicolon. +- [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 + +- [ ] 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/TODO.md b/diskann-benchmark-runner-derive/TODO.md new file mode 100644 index 0000000000..1ce50dda66 --- /dev/null +++ b/diskann-benchmark-runner-derive/TODO.md @@ -0,0 +1,128 @@ +# Benchmark Input Discoverability + +The benchmark configuration files are a portable protocol: they support sharing +configurations across machines, batching benchmarks, validating runs, and analyzing results. +The missing piece is a human-facing projection of that protocol. + +Use dedicated deserialization DTOs as the public configuration boundary. Convert DTOs into +validated runtime types after deserialization. The derive macro should document and enforce a +deliberately constrained DTO language rather than attempt general Rust or Serde reflection. + +## Initial contract + +- [ ] Support named structs. +- [ ] Support explicitly tagged enums, including unit, tuple, and struct variants where needed. +- [ ] Require documentation for configuration types, fields, and enum variants. +- [ ] Support known primitive and domain leaf types. +- [ ] Support documented DTO composition through selected containers such as `Option` and + `Vec`. +- [ ] Read representation-changing metadata from Serde attributes so Serde remains the source + of truth. +- [ ] Support `rename` and `rename_all`. +- [ ] Support enum `tag` and `content`. +- [ ] Decide whether `default` and `alias` should be included in the initial metadata model. +- [ ] Reject unsupported types and Serde attributes with actionable `syn::Error` diagnostics. +- [ ] Reject asymmetric serialization and deserialization names unless the metadata model + explicitly represents both. +- [ ] Reject `untagged`, `remote`, `with`, `deserialize_with`, `try_from`, and skipped input + fields initially. +- [ ] Add an explicit escape hatch for intentional opaque or externally implemented leaf + types. +- [ ] Evaluate `flatten` separately; prefer explicit nested DTOs unless flattening provides a + clear authoring benefit. +- [ ] Keep validation and DTO-to-runtime conversion outside the reflection system. +- [ ] Keep complete examples curated rather than synthesizing arbitrary field values or + Cartesian products of enum variants. + +## Runtime metadata + +- [ ] Replace or refine the prototype `Reflect` API around benchmark configuration + documentation rather than general-purpose reflection. +- [ ] Choose names that communicate the constrained public role, such as `BenchmarkInput` for + the derive and `DescribeInput` for the generated runtime trait. +- [ ] Represent type, field, and variant documentation. +- [ ] Represent effective serialized names after applying supported Serde rules. +- [ ] Represent nested DTOs and selected containers. +- [ ] Represent accepted enum variants and their tagging strategy. +- [ ] Decide whether metadata should be statically stored or constructed on demand; favor the + simpler dynamic model unless measurements justify static storage. +- [ ] Provide getters or a renderer-facing API for all metadata. +- [ ] Avoid requiring arbitrary runtime and third-party types to implement the reflection + trait. + +## Examples + +- [ ] Add an API for multiple named, curated examples per registered input. +- [ ] Include a short description with each example. +- [ ] Keep examples on the input/DTO API rather than inferring values in the derive macro. +- [ ] Decide whether the derive accepts an examples function: + + ```rust + #[benchmark(examples = Self::examples)] + ``` + +- [ ] Ensure every example serializes and deserializes successfully. +- [ ] Consider validating examples through the normal DTO-to-runtime conversion path. + +## CLI integration + +- [ ] Expose input descriptions through the dynamic registry. +- [ ] Extend `inputs --describe ` or settle on a clearer equivalent command. +- [ ] Render the input summary, fields, nested objects, enum choices, and important defaults. +- [ ] Render one or more complete examples. +- [ ] Consider `skeleton --input ` for generating a directly editable configuration. +- [ ] Consider emitting a commented JSONC-style template for authoring while retaining strict + JSON as the canonical shared representation. +- [ ] Preserve feature-gated input behavior and diagnostics. +- [ ] Improve deserialization and validation errors with paths such as + `jobs[2].input.search.runs[0].search_l`. + +## Macro implementation + +- [ ] Replace `todo!` branches with structured compile errors. +- [ ] Implement named struct generation. +- [ ] Implement the accepted enum forms. +- [ ] Generate appropriate generic bounds for nested reflected types. +- [ ] Extract and normalize literal rustdoc. +- [ ] Parse the supported Serde subset with `syn::parse_nested_meta`. +- [ ] Implement and test Serde rename rules used by benchmark DTOs. +- [ ] Preserve useful source spans in generated diagnostics. +- [ ] Add explicit diagnostics for every rejected Serde feature. +- [ ] Add documentation-specific helper attributes only for behavior Serde does not control, + such as examples, hiding documentation, or an opaque leaf override. +- [ ] Do not generate serialization, deserialization, validation, or arbitrary example values. + +## Migration + +- [ ] Select one simple DTO and one representative complex DTO as the initial vertical slice. +- [ ] Wire those DTOs through derive, registry, CLI rendering, examples, and validation. +- [ ] Use the vertical slice to confirm the output format before migrating all inputs. +- [ ] Convert remaining benchmark-facing inputs to dedicated DTOs where runtime concerns are + still mixed into deserialization types. +- [ ] Add rustdoc to all exposed DTO fields and variants. +- [ ] Replace custom deserialization patterns where a simpler DTO plus conversion can express + the same behavior. +- [ ] Add explicit overrides only where external or specialized types make them unavoidable. +- [ ] Migrate remaining registered inputs incrementally rather than requiring an atomic + workspace-wide conversion. + +## Tests + +- [ ] Add unit tests for rustdoc extraction and normalization. +- [ ] Add tests for each supported Serde naming and tagging rule. +- [ ] Add compile-fail tests for unsupported item shapes, missing documentation, unsupported + Serde attributes, and invalid combinations. +- [ ] Compare generated names and shapes with actual `serde_json` serialization. +- [ ] Add CLI golden tests for descriptions, examples, nested DTOs, enums, feature-gated + inputs, and error output. +- [ ] Test that curated examples round-trip. +- [ ] Run formatting, clippy with warnings denied, and targeted workspace tests. + +## Suggested rollout + +- [ ] Phase 1: named DTO structs, documentation, known leaves and containers, and CLI rendering. +- [ ] Phase 2: tagged enums, Serde naming, and curated examples. +- [ ] Phase 3: migrate representative inputs and refine diagnostics. +- [ ] Phase 4: add flattening or other Serde behavior only in response to concrete DTO needs. +- [ ] Phase 5: migrate the remaining benchmark inputs and stabilize the public API. + diff --git a/diskann-benchmark-runner-derive/src/attributes.rs b/diskann-benchmark-runner-derive/src/attributes.rs new file mode 100644 index 0000000000..685f89f4fd --- /dev/null +++ b/diskann-benchmark-runner-derive/src/attributes.rs @@ -0,0 +1,506 @@ +/* + * 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") +} + +#[must_use] +fn is_reflect_attr(attr: &syn::Attribute) -> bool { + attr.path().is_ident("reflect") +} + +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(()) + } +} + +/// 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 { + rename_all: RenameAll, + enum_repr: EnumRepr, + type_name: TypeName, +} + +impl Container { + pub(crate) fn parse(attrs: &[syn::Attribute]) -> syn::Result { + 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_once(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")) + })?; + } + + // 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 { + rename_all: self.rename_all, + }) + } + + /// 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 { + tag: syn::LitStr, + }, + Adjacent { + tag: syn::LitStr, + content: syn::LitStr, + }, +} + +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, + ) -> 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`", + )), + } + } + + /// 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(()), + 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", + )), + } + } +} + +/// Supported subset of `serde(rename_all = "...")` +#[derive(Default, Debug, Clone, Copy, PartialEq)] +pub(crate) enum RenameAll { + #[default] + None, + Lower, + Snake, + Kebab, +} + +impl RenameAll { + fn supported() -> &'static str { + "\"lowercase\", \"snake_case\", or \"kebab-case\"" + } + + fn parse(s: &str) -> Option { + match s { + "lowercase" => Some(Self::Lower), + "snake_case" => Some(Self::Snake), + "kebab-case" => Some(Self::Kebab), + _ => None, + } + } + + /// 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, + "serde attribute `rename_all` found multiple times", + )) + } else { + let value = s.value(); + match Self::parse(&value) { + Some(me) => { + *self = me; + Ok(()) + } + None => Err(syn::Error::new_spanned( + s, + format!( + "unsupported serde `rename_all` rule \"{}\" - expected one of {}", + value, + Self::supported() + ), + )), + } + } + } + + /// 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('_', "-"), + } + } + + /// Apply the rename rule to `variant`. + 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()) + } + } + + /// 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('_', "-"), + } + } + + /// Apply the rename rule to `field`. + pub(crate) fn apply_to_field(&self, field: syn::LitStr) -> syn::LitStr { + if *self == Self::None { + field + } else { + syn::LitStr::new(&self.apply_to_field_str(&field.value()), field.span()) + } + } +} + +/// A one-type variant or field renamer. +#[derive(Default)] +pub(crate) struct RenameOnce { + rename: Option, +} + +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 generating 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<()> { + 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") + } + (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(()) + } +} + +//---------// +// Variant // +//---------// + +/// Variant level attributes. +#[derive(Default)] +pub(crate) struct Variant { + pub(crate) rename_variant: RenameOnce, + pub(crate) rename_variant_fields: RenameAll, +} + +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()?; + 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_variant.rename, value, "rename")?; + return Ok(()); + } + + Err(meta.error("unsupported Serde attribute for Reflect")) + })?; + } + + Ok(me) + } +} + +//-------// +// Field // +//-------// + +/// Field level attributes. +#[derive(Default)] +pub(crate) struct Field { + pub(crate) rename_field: RenameOnce, +} + +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_field.rename, value, "rename")?; + return Ok(()); + } + + Err(meta.error("unsupported Serde attribute for Reflect")) + })?; + } + + Ok(me) + } +} + +/////////// +// 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 new file mode 100644 index 0000000000..6da9748d21 --- /dev/null +++ b/diskann-benchmark-runner-derive/src/lib.rs @@ -0,0 +1,509 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use proc_macro2::TokenStream; +use quote::{quote, quote_spanned}; +use syn::{Data, DeriveInput, Fields, parse_macro_input, parse_quote, spanned::Spanned}; + +mod attributes; + +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; +/// +/// /// A test aggregate. +/// #[derive(Reflect)] +/// struct MyInput { +/// /// The number of threads. +/// threads: usize, +/// } +/// ``` +/// +/// # 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(_)) { + return Err(syn::Error::new_spanned( + input, + "Reflect cannot be derived for unions", + )); + } + + 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, container.type_name())?; + + let common = DeriveCommon { + doc, + generics, + format_type_name, + container, + }; + + match &input.data { + Data::Struct(s) => process_struct(&input, s, common), + Data::Enum(e) => process_enum(&input, e, common), + Data::Union(_) => unreachable!("this has already been checked"), + } +} + +/// 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, + /// Serde container-level attributes. + container: attributes::Container, +} + +/// 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. +/// +/// # 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 + .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(); + + // Check that the type-name attributes are compatible with the struct. + let prefix: Option = 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)?; + }; + 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 + f.write_str(#name) + }; + Ok(ts) + } else { + let writes = arguments.iter().enumerate().map(|(index, argument)| { + if index == 0 { + quote! { + #argument + } + } else { + quote! { + f.write_str(", ")?; + #argument + } + } + }); + + 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: attributes::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_all)?; + Ok(quote!(#path::tree::Fields::Named(vec![#(#list),*]))) + } + Fields::Unnamed(fields) => { + add_field_bounds(generics, &fields.unnamed); + let list = unnamed_fields(&fields.unnamed); + + // 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)), + } +} + +fn named_fields<'a, I>( + fields: I, + rename_all: attributes::RenameAll, +) -> 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"); + + 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)?; + let name = rename_field.apply_to_field(name, rename_all); + Ok(quote_spanned! { ty.span()=> #path::tree::NamedField::new::<#ty>(#name, #doc) }) + }) + .collect() +} + +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) } + }) + .collect() +} + +/// Generate the `Reflect` implementation. +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 attributes::Struct { rename_all } = container.try_as_struct()?; + + let type_name = &input.ident; + let path = crate_name(); + + let fields = build_fields(&s.fields, &mut generics, rename_all)?; + + let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); + + let ts = quote! { + impl #impl_generics #path::Reflect for #type_name #ty_generics #where_clause { + fn ty() -> #path::Type { + #path::Type::aggregate( + #fields, + #doc, + ) + } + + fn format_type_name(f: &mut dyn ::std::fmt::Write) -> ::std::fmt::Result { + #format_type_name + } + } + }; + + Ok(ts) +} + +//-------// +// Enums // +//-------// + +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 attributes::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 variants = e + .variants + .iter() + .map(|v| -> syn::Result { + let doc = format_docstrings(&v.attrs); + let name = syn::LitStr::new(strip_raw_prefix(&v.ident.to_string()), v.ident.span()); + let attributes::Variant { + rename_variant, + rename_variant_fields, + } = attributes::Variant::parse(&v.attrs)?; + + 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); + Ok(quote!(#path::tree::Variant::new(#name, #fields, #doc))) + }) + .collect::>>()?; + + // Build the enum representation. + let enum_repr = match enum_repr { + 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 }) + } + }; + + let (impl_generics, ty_generics, where_clause) = generics.split_for_impl(); + + let ts = quote! { + impl #impl_generics #path::Reflect for #type_name #ty_generics #where_clause { + fn ty() -> #path::Type { + #path::Type::enum_( + #enum_repr, + [#(#variants),*], + #doc, + ) + } + + fn format_type_name(f: &mut dyn ::std::fmt::Write) -> ::std::fmt::Result { + #format_type_name + } + } + }; + + Ok(ts) +} + +//-------------// +// 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()) }, + } +} + +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")) + } +} + +//-----// +// raw // +//-----// + +fn strip_raw_prefix(s: &str) -> &str { + s.strip_prefix("r#").unwrap_or(s) +} 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/Cargo.toml b/diskann-benchmark-runner/Cargo.toml index 041f16df33..bc4c3deb6a 100644 --- a/diskann-benchmark-runner/Cargo.toml +++ b/diskann-benchmark-runner/Cargo.toml @@ -10,7 +10,9 @@ edition = "2024" [dependencies] 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/dev/main.rs b/diskann-benchmark-runner/dev/main.rs index 788e9016b3..3e9db47e3f 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 c54baed43a..63355e5b10 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,12 +212,28 @@ 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)?)?; + + // 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 { writeln!(output, "No input found for \"{}\"", describe)?; @@ -386,6 +407,11 @@ 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(()) } @@ -482,6 +508,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(()) + } } /////////// @@ -541,8 +590,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"; @@ -560,42 +607,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, @@ -606,7 +617,7 @@ mod tests { fn new(dir: &Path) -> Self { Self { dir: dir.into(), - overwrite: overwrite(), + overwrite: ux::overwrite(), } } @@ -614,7 +625,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() @@ -723,7 +734,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); } @@ -765,9 +776,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!( @@ -777,7 +788,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/files.rs b/diskann-benchmark-runner/src/files.rs index 02b0c97979..0206dcd31d 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/input.rs b/diskann-benchmark-runner/src/input.rs index 64bb7fac07..a34822c4d9 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::{Checker, internal::visibility::Visibility}; +use crate::{Checker, Reflect, Reflection, internal::visibility::Visibility}; /// 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/lib.rs b/diskann-benchmark-runner/src/lib.rs index 403f15fc4e..c6eb596aa8 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 features; @@ -16,6 +19,7 @@ pub mod app; pub mod files; pub mod input; pub mod output; +pub mod reflect; pub mod registry; pub mod utils; @@ -25,6 +29,7 @@ pub use checker::Checker; pub use features::Features; pub use input::Input; pub use output::Output; +pub use reflect::{Reflect, Reflection}; pub use registry::{Registry, RegistryError}; pub use result::Checkpoint; diff --git a/diskann-benchmark-runner/src/reflect/mod.rs b/diskann-benchmark-runner/src/reflect/mod.rs new file mode 100644 index 0000000000..d535ab290b --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/mod.rs @@ -0,0 +1,416 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +//! Run time inspection of types. + +use std::{ + any::TypeId, + fmt::{self, Write}, +}; + +pub use diskann_benchmark_runner_derive::Reflect; + +mod render; +pub mod tree; +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, + { + Self { + reflection: internal::VTable::new::(), + } + } + + /// 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)() + } + + /// 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, + { + visit_with(*self, f) + } +} + +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() + } +} + +/// A [`std::fmt::Display`] compatible type for [`Reflection`]. +/// +/// See: [`Reflection::type_name`]. +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) + } +} + +/// 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 { + let mut r = render::Renderer::new(f, 2); + r.render_subject(self.0) + } +} + +//------------// +// 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::NewType(newtype) => push(newtype.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()), + Type::Optional(opt) => push(opt.value()), + } + } + + // loop + if let Some(r) = stack.pop() { + reflection = r; + } else { + break Ok(()); + } + } +} + +/////////////// +// Bootstrap // +/////////////// + +impl Reflect for std::marker::PhantomData +where + T: Reflect, +{ + fn ty() -> Type { + Type::primitive(tree::PrimitiveKind::Null, None) + } + + fn format_type_name(f: &mut dyn Write) -> fmt::Result { + write!(f, "PhantomData<{}>", Reflection::new::().type_name()) + } +} + +macro_rules! primitive { + ($T:ty, $kind:ident, $doc:literal, $type_name:literal) => { + impl Reflect for $T { + fn ty() -> Type { + Type::primitive(tree::PrimitiveKind::$kind, Some($doc.into())) + } + + fn format_type_name(f: &mut dyn Write) -> fmt::Result { + f.write_str($type_name) + } + } + }; +} + +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::NonZeroU32, + Number, + "A system dependent, 32-bit unsigned integer", + "NonZero" +); +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"); +primitive!(std::path::PathBuf, String, "A file path", "PathBuf"); + +impl Reflect for Option +where + T: Reflect, +{ + fn ty() -> Type { + Type::optional::(Some("An optional type".into())) + } + + fn format_type_name(f: &mut dyn Write) -> fmt::Result { + f.write_str("Option<")?; + T::format_type_name(f)?; + f.write_str(">") + } +} + +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()) + } +} + +////////////// +// Internal // +////////////// + +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, + } + + 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::, + } + } + } + + fn ty() -> super::Type + where + T: super::Reflect, + { + ::ty() + } + + fn format_type_name(f: &mut dyn std::fmt::Write) -> std::fmt::Result + where + T: super::Reflect, + { + ::format_type_name(f) + } + + 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 new file mode 100644 index 0000000000..0fe650b3b6 --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/render.rs @@ -0,0 +1,610 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::fmt::{self, Write}; + +use crate::utils::fmt::Quote; + +use super::{ + Reflection, + tree::{ + Aggregate, Enum, EnumRepr, Fields, NamedField, Optional, Sequence, Type, 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), + Type::Optional(opt) => self.render_optional(opt), + } + } + + 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 previous: Option<&Variant> = None; + for variant in enum_.variants().iter() { + // 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)?; + previous = Some(variant) + } + + 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(()), + ) + } + + 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 // + //--------// + + 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(Some(i), field)?; + } + } + Fields::NewType(newtype) => self.render_unnamed_field(None, newtype)?, + 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: Option, field: &UnnamedField) -> fmt::Result { + let tagged = Tagged::new(field.field()); + let will_render_body = self.will_render(tagged.ty()); + + 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 { + 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()) + }) + } +} + +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + use std::{ + fs::File, + io::{BufRead, BufReader, Write}, + path::{Path, PathBuf}, + }; + + use serde::{Deserialize, Serialize}; + + use crate::{Reflect, ux}; + + // 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::")] + #[expect(unused, reason = "testing")] + enum SimpleDataType { + Float32, + Float16, + } + + /// Select the data type please. + #[derive(Reflect)] + #[serde(rename_all = "kebab-case")] + #[reflect(prefix = "render::")] + #[expect(unused, reason = "testing")] + enum AnnotatedDataType { + /// Use high-precision. + Float32, + /// Use lower precision. + Float16, + Int8, + } + + #[derive(Reflect)] + #[serde(rename_all = "snake_case")] + #[reflect(prefix = "render::")] + #[expect(unused, reason = "testing")] + 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)] + #[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. + 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/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..d030a6f67a --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/test/serde.rs @@ -0,0 +1,789 @@ +/* + * 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_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}; + +use hashbrown::HashSet; +use serde::Serialize; +use serde_json::Value; + +use crate::reflect::{Reflect, Reflection, Type, tree}; + +/// 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::Arguments`] that describe the sequence of operations that led us into a mess. +/// +/// Use the [`context`] macro for creating nested 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")), + Type::Optional(opt) => check_optional(opt, s, context!(ctx, s, "optional")), + } +} + +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, +/// } +/// ``` +/// Adjacency (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 = m.iter().next().unwrap(); + (&kv.0, Some(Cow::Borrowed(&kv.1))) + } + _ => panic!("invalid representation\n\n{}", ctx), + }, + + // 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 + // 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); + (t, Some(Cow::Owned(Value::Object(map)))) + } + + // 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); + 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))) + } + } + } +} + +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), + }; + + 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()), + ) +} + +//----------// +// 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()), + ); + } +} + +//----------// +// 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 // +//---------// + +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(false); + + 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); + + check_struct(0f32); + check_struct(0f64); + + check_struct(std::marker::PhantomData::); + + check_struct(String::from("hello")); +} + +#[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 newtype 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)] + struct NewTypeEmpty {} + + #[derive(Serialize, Reflect)] + #[serde(rename_all = "kebab-case")] + #[serde(tag = "fizzle")] + enum Enum { + Unit, + EmptyStruct {}, + NewType(NewTypePayload), + NewTypeEmpty(NewTypeEmpty), + 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::NewTypeEmpty(NewTypeEmpty {}), + 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!("adjacently 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)], + }); +} + +#[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/test/unit.rs b/diskann-benchmark-runner/src/reflect/test/unit.rs new file mode 100644 index 0000000000..6dd149082d --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/test/unit.rs @@ -0,0 +1,395 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +//! [`Reflect`] macro unit tests. + +use std::assert_matches; + +use crate::reflect::{Reflect, Reflection, tree::Fields}; + +#[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()); + + // Unit structs should be unit aggregates. + let agg = ty.as_aggregate().unwrap(); + assert!(agg.fields().is_unit()); +} + +#[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()); + + // Unit structs should be unit aggregates. + let agg = ty.as_aggregate().unwrap(); + assert!(agg.fields().is_unit()); +} + +#[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()); + + // Unit structs should be unit aggregates. + let agg = ty.as_aggregate().unwrap(); + assert!(agg.fields().is_unit()); +} + +#[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()); + + // 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] +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()); + + // 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] +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. This should be a "newtype". + #[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. 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_newtype().unwrap(); + assert_eq!(f.doc().unwrap(), "Field 0."); + assert_eq!(f.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_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_newtype().unwrap(); + assert_eq!(f.doc().unwrap(), "Buzz buzz"); + assert_eq!(f.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 new file mode 100644 index 0000000000..989cd09061 --- /dev/null +++ b/diskann-benchmark-runner/src/reflect/tree.rs @@ -0,0 +1,481 @@ +/* + * Copyright (c) Microsoft Corporation. + * 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), + Aggregate(Aggregate), + Enum(Enum), + Sequence(Sequence), + Optional(Optional), +} + +impl Type { + /// Construc a new [`Type::Primitive`]. + pub fn primitive(kind: PrimitiveKind, doc: Option) -> Self { + Self::Primitive(Primitive::new(kind, doc)) + } + + /// Construc a new [`Type::Aggregate`]. + pub fn aggregate(fields: Fields, doc: Option) -> Self { + Self::Aggregate(Aggregate::new(fields, doc)) + } + + /// Construc a new [`Type::Enum`]. + pub fn enum_( + repr: EnumRepr, + variants: impl IntoIterator, + doc: Option, + ) -> Self { + Self::Enum(Enum::new(repr, variants, doc)) + } + + /// Construc a new [`Type::Sequence`]. + pub fn sequence(doc: Option) -> Self + where + T: Reflect, + { + Self::Sequence(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::Optional(Optional::new::(doc)) + } + + /// Return the struct level documentation if available. + 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(), + } + } + + /// 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, + Self::Aggregate(a) => a.has_body(), + Self::Enum(e) => e.has_body(), + Self::Sequence(_) => true, + Self::Optional(_) => true, + } + } + + #[cfg(test)] + pub(super) fn as_aggregate(&self) -> Option<&Aggregate> { + if let Self::Aggregate(aggregate) = self { + Some(aggregate) + } else { + None + } + } + + #[cfg(test)] + pub(super) fn as_enum(&self) -> Option<&Enum> { + if let Self::Enum(enum_) = self { + Some(enum_) + } else { + None + } + } +} + +//-----------// +// Primitive // +//-----------// + +/// The native JSON representation for a primitive. +#[derive(Debug, Clone, Copy)] +pub enum PrimitiveKind { + Null, + Boolean, + Number, + String, +} + +/// A primitive type that maps closely to a native JSON type. +#[derive(Debug)] +pub struct Primitive { + kind: PrimitiveKind, + doc: Option, +} + +impl Primitive { + pub(crate) fn new(kind: PrimitiveKind, doc: Option) -> Self { + Self { kind, doc } + } + + pub(super) fn kind(&self) -> PrimitiveKind { + self.kind + } + + pub(super) fn doc(&self) -> Option<&str> { + self.doc.as_deref() + } +} + +//-----------// +// Aggregate // +//-----------// + +/// A representation of aggregates like normal structs, unit structs, and tuple-like structs. +#[derive(Debug)] +pub struct Aggregate { + fields: Fields, + doc: Option, +} + +impl Aggregate { + pub(crate) 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() + } +} + +/// 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, +} + +impl Fields { + pub(super) fn has_body(&self) -> bool { + match self { + Self::Named(fields) => !fields.is_empty(), + Self::Unnamed(fields) => !fields.is_empty(), + Self::NewType(_) => true, + Self::Unit => false, + } + } + + /// 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) + } + + #[cfg(test)] + pub(super) fn as_named(&self) -> Option<&[NamedField]> { + if let Self::Named(fields) = self { + Some(fields) + } else { + None + } + } + + #[cfg(test)] + pub(super) fn as_unnamed(&self) -> Option<&[UnnamedField]> { + if let Self::Unnamed(fields) = self { + Some(fields) + } else { + None + } + } + + #[cfg(test)] + pub(super) fn as_newtype(&self) -> Option<&UnnamedField> { + if let Self::NewType(field) = self { + Some(field) + } else { + None + } + } + + /// 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, + field: Reflection, + doc: Option, +} + +impl NamedField { + /// Construct a new [`NamedField`] for `T`. + 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() + } +} + +/// An unnamed field. +#[derive(Debug)] +pub struct UnnamedField { + field: Reflection, + doc: Option, +} + +impl UnnamedField { + /// Construct a new [`UnnamedField`] for `T`. + 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 // +//------// + +/// A representation for enums. +#[derive(Debug)] +pub struct Enum { + repr: EnumRepr, + variants: Vec, + doc: Option, +} + +impl Enum { + /// Construct a new [`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() + } +} + +/// Describe how an enum is being represented by `serde`. +#[derive(Debug)] +#[non_exhaustive] +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, + }, +} + +/// A variant of an [`Enum`]. +#[derive(Debug)] +pub struct Variant { + name: &'static str, + fields: Fields, + doc: Option, +} + +impl Variant { + /// Construct a new [`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 // +//----------// + +/// A homogeneous sequence of values. +#[derive(Debug)] +pub struct Sequence { + element: Reflection, + doc: Option, +} + +impl Sequence { + /// Create a new [`Sequence`] containing `T`. + 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() + } +} + +//----------// +// Optional // +//----------// + +/// An [`Option`]. +#[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/registry.rs b/diskann-benchmark-runner/src/registry.rs index b44a845c7f..48c092a24f 100644 --- a/diskann-benchmark-runner/src/registry.rs +++ b/diskann-benchmark-runner/src/registry.rs @@ -3,12 +3,16 @@ * Licensed under the MIT license. */ -use std::collections::{HashMap, hash_map::Entry}; +use std::{ + any::TypeId, + collections::{HashMap, hash_map::Entry}, +}; +use hashbrown::{HashSet, hash_set}; use thiserror::Error; use crate::{ - Checkpoint, Features, Input, Output, + Checkpoint, Features, Input, Output, Reflection, benchmark::{self, Benchmark, MatchContext, Regression, Score, internal::AnnotatedMatch}, input, internal::visibility::Visibility, @@ -16,8 +20,18 @@ use crate::{ /// 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`]. + name_map: HashMap, + + /// The type IDs present in `name_map`. + type_ids: HashSet, + + /// The registered benchmarks. benchmarks: Vec, } @@ -26,10 +40,16 @@ impl Registry { pub fn new() -> Self { Self { inputs: HashMap::new(), + name_map: HashMap::new(), + type_ids: HashSet::new(), benchmarks: Vec::new(), } } + //-------// + // Input // + //-------// + /// Return the input with the registered `tag` if present. Otherwise, return `None`. /// /// Inputs are automatically registered as a side-effect of: @@ -44,6 +64,20 @@ 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() + } + + /// Return an iterator over all the registered type-names. + pub fn type_names(&self) -> impl ExactSizeIterator { + self.name_map.keys().map(|v| &**v) + } + //--------------// // Registration // //--------------// @@ -229,6 +263,14 @@ 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. + Self::register_reflection( + Reflection::new::(), + &mut self.name_map, + &mut self.type_ids, + )?; + v.insert(Box::new(wrapper)); Ok(()) } @@ -242,17 +284,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()), + )) } } } @@ -274,26 +316,121 @@ 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(()) + } + } + //-------------------// // Regression Checks // //-------------------// @@ -479,16 +616,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)] @@ -761,3 +925,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); + } +} diff --git a/diskann-benchmark-runner/src/test/dim.rs b/diskann-benchmark-runner/src/test/dim.rs index 329cb64a34..c017c7fb9d 100644 --- a/diskann-benchmark-runner/src/test/dim.rs +++ b/diskann-benchmark-runner/src/test/dim.rs @@ -8,7 +8,7 @@ use std::io::Write; use serde::{Deserialize, Serialize}; use crate::{ - Benchmark, Checker, Checkpoint, Input, Output, + Benchmark, Checker, Checkpoint, Input, Output, Reflect, benchmark::{MatchContext, PassFail, Regression, Score}, }; @@ -16,7 +16,8 @@ use crate::{ // Input // /////////// -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "benchmark::test::")] pub(super) struct DimInput { dim: Option, } @@ -55,7 +56,8 @@ impl Input for DimInput { // Tolerance // /////////////// -#[derive(Debug, Clone, Serialize, Deserialize)] +#[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 033fab55db..064c051566 100644 --- a/diskann-benchmark-runner/src/test/gated.rs +++ b/diskann-benchmark-runner/src/test/gated.rs @@ -16,7 +16,8 @@ use std::io::Write; use serde::{Deserialize, Serialize}; use crate::{ - Benchmark, Checker, Checkpoint, Input, Output, benchmark::MatchContext, benchmark::Score, + Benchmark, Checker, Checkpoint, Input, Output, Reflect, benchmark::MatchContext, + benchmark::Score, }; use super::{dim::DimInput, typed::TypeInput}; @@ -85,7 +86,8 @@ impl Benchmark for AnotherGatedBench { // Partially Gated with Input Always Registered // ////////////////////////////////////////////////// -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "benchmark::test::")] pub(super) struct SampleInput { value: String, } @@ -145,7 +147,8 @@ 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)] +#[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 3ab0e7a891..031c9a3cdb 100644 --- a/diskann-benchmark-runner/src/test/typed.rs +++ b/diskann-benchmark-runner/src/test/typed.rs @@ -8,7 +8,7 @@ use std::io::Write; use serde::{Deserialize, Serialize}; use crate::{ - Benchmark, Checker, Checkpoint, Input, Output, + Benchmark, Checker, Checkpoint, Input, Output, Reflect, benchmark::{MatchContext, PassFail, Regression, Score}, utils::datatype::{AsDataType, DataType}, }; @@ -24,11 +24,13 @@ pub(crate) struct TypeInput { error_when_checked: bool, } -#[derive(Serialize, Deserialize)] +/// 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, - // Should we return an error when deserializing? + /// Should we return an error when deserializing? error_when_checked: bool, } @@ -78,9 +80,10 @@ impl Input for TypeInput { // Tolerance // /////////////// -#[derive(Debug, Serialize, Deserialize)] +#[derive(Debug, Serialize, Deserialize, Reflect)] +#[reflect(prefix = "benchmark::test::")] 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..bce0d70288 100644 --- a/diskann-benchmark-runner/src/utils/datatype.rs +++ b/diskann-benchmark-runner/src/utils/datatype.rs @@ -6,11 +6,14 @@ 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")] +#[reflect(prefix = "benchmark::")] pub enum DataType { Float64, Float32, 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/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-runner/src/utils/required.rs b/diskann-benchmark-runner/src/utils/required.rs new file mode 100644 index 0000000000..8bbe7a758c --- /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, reflect::tree}; + +/// 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-runner/src/ux.rs b/diskann-benchmark-runner/src/ux.rs index 9f862d5619..3fade01da6 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/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 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-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" +======== diff --git a/diskann-benchmark-simd/src/lib.rs b/diskann-benchmark-simd/src/lib.rs index c2fec58069..b09d4c3eb1 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, Reflect, Registry, }; //////////////// @@ -55,8 +55,9 @@ impl std::ops::Deref for DisplayWrapper<'_, T> { // Inputs // //////////// -#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize)] -#[serde(rename_all = "snake_case")] +#[derive(Debug, Clone, Copy, PartialEq, Serialize, Deserialize, Reflect)] +#[serde(rename_all = "kebab-case")] +#[reflect(prefix = "simd::")] pub enum SimilarityMeasure { SquaredL2, InnerProduct, @@ -74,17 +75,41 @@ 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")] +#[reflect(prefix = "simd::")] 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")] #[expect(non_camel_case_types)] X86_64_V4, + + /// Target AVX2. + /// + /// Only usable when compiling for x86-64. #[serde(rename = "x86-64-v3")] #[expect(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 +126,39 @@ 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)] +#[reflect(prefix = "simd::")] 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)] +#[reflect(prefix = "simd::")] 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 +249,8 @@ 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)] +#[reflect(prefix = "simd::")] struct SimdTolerance { min_time_regression: NonNegativeFinite, } diff --git a/diskann-benchmark/src/index/benchmarks.rs b/diskann-benchmark/src/index/benchmarks.rs index 33f5fca649..b8dc1c2571 100644 --- a/diskann-benchmark/src/index/benchmarks.rs +++ b/diskann-benchmark/src/index/benchmarks.rs @@ -226,7 +226,7 @@ where build::set_start_points( index.provider(), data.as_view(), - *build.start_point_strategy(), + build.start_point_strategy(), )?; Ok(index) }, @@ -843,9 +843,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()))?; @@ -865,7 +866,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 e0287638b1..7ad823bc6b 100644 --- a/diskann-benchmark/src/index/inmem/scalar.rs +++ b/diskann-benchmark/src/index/inmem/scalar.rs @@ -258,7 +258,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 b989917e3b..431a86ce60 100644 --- a/diskann-benchmark/src/index/inmem2.rs +++ b/diskann-benchmark/src/index/inmem2.rs @@ -69,14 +69,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, @@ -85,14 +89,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, @@ -128,7 +134,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, @@ -136,7 +143,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, @@ -149,7 +157,8 @@ 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, @@ -157,7 +166,8 @@ mod dto { pub(super) quantization: Option, } - #[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..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::datatype::DataType, Checker}; +use diskann_benchmark_runner::{ + utils::{datatype::DataType, RequiredOption}, + Checker, Reflect, +}; use diskann_bftree::BfTreeProviderParameters; use serde::{Deserialize, Serialize}; @@ -20,7 +23,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 +33,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 +83,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 +136,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 +158,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 +187,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 +215,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 +233,17 @@ 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 + .cloned() .unwrap_or_default() .into_config(), - quant_vector_provider_config: quant_store_config.clone().unwrap_or_default().into_config(), graph_params: None, use_snapshot, }) @@ -251,14 +251,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 +290,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 +317,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 +334,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 +349,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 +393,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 +422,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 +439,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 +456,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 +465,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 +488,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 +525,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 +556,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 +577,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 +598,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 +608,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 +663,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 +673,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 +705,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 +726,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 6f1ec9dddd..2231c5c752 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 d25a56fc92..1741eb58e7 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..e5deb79d31 100644 --- a/diskann-benchmark/src/multi_vector/driver.rs +++ b/diskann-benchmark/src/multi_vector/driver.rs @@ -12,7 +12,7 @@ use diskann_benchmark_runner::{ num::{relative_change, NonNegativeFinite}, percentiles, MicroSeconds, }, - Checker, Input, + Checker, Input, Reflect, }; use diskann_quantization::multi_vector::{Mat, MatRef, MaxSimKernel, Overflow, Standard}; use rand::{ @@ -30,7 +30,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, diff --git a/diskann-inmem/integration/index/runner.rs b/diskann-inmem/integration/index/runner.rs index f65e80ea4f..e9146cdab6 100644 --- a/diskann-inmem/integration/index/runner.rs +++ b/diskann-inmem/integration/index/runner.rs @@ -8,10 +8,10 @@ use std::{io::Write, sync::Arc}; use anyhow::Context; use diskann::graph::{DiskANNIndex, search::Knn}; use diskann_benchmark_runner::{ - Checker, Checkpoint, Output, Registry, RegistryError, + Checker, Checkpoint, Output, Reflect, Registry, RegistryError, benchmark::{MatchContext, PassFail, Regression, Score}, files::InputFile, - utils::fmt::Indent, + utils::{RequiredOption, fmt::Indent}, }; use diskann_utils::views::Matrix; use diskann_vector::distance::Metric; @@ -43,8 +43,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 +74,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 +100,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, @@ -115,24 +118,27 @@ mod dto { pub(super) mod spherical { use super::*; - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] + #[reflect(prefix = "graph::spherical::")] pub(in crate::index::runner) enum Bits { One, Two, Four, } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] + #[reflect(prefix = "graph::spherical::")] pub(in crate::index::runner) enum Rerank { None, F16, } } - #[derive(Debug, Serialize, Deserialize)] + #[derive(Debug, Serialize, Deserialize, Reflect)] #[serde(rename_all = "kebab-case")] + #[reflect(prefix = "graph::")] pub(super) enum Representation { FullPrecision { data_type: DataType, @@ -143,7 +149,8 @@ mod dto { }, } - #[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, @@ -151,20 +158,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, @@ -400,7 +409,10 @@ 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 { @@ -417,7 +429,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()), } } @@ -635,17 +647,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 4bc5f8a982..63ddd1b4c5 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 9e5882ed1c..bcd4158b22 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 893133b63a..3120a0f5e0 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::{Reflect, Registry, RegistryError, utils::fmt::KeyValue}; use rand::{Rng, SeedableRng, distr::Uniform, rngs::StdRng}; use serde::{Deserialize, Serialize}; @@ -58,7 +58,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/store/simple.rs b/diskann-inmem/integration/store/simple.rs index 11095a67db..aa82486dde 100644 --- a/diskann-inmem/integration/store/simple.rs +++ b/diskann-inmem/integration/store/simple.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::simple::")] struct Input { /// Shared stress test setup. setup: super::Setup, diff --git a/diskann-inmem/integration/support/datatype.rs b/diskann-inmem/integration/support/datatype.rs index 729d91edbd..7ef44dc426 100644 --- a/diskann-inmem/integration/support/datatype.rs +++ b/diskann-inmem/integration/support/datatype.rs @@ -3,6 +3,7 @@ * Licensed under the MIT license. */ +use diskann_benchmark_runner::Reflect; use diskann_utils::{ sampling::medoid::ComputeMedoid, views::{Matrix, MatrixView, MutMatrixView}, @@ -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 { diff --git a/expand.rs b/expand.rs new file mode 100644 index 0000000000..e69de29bb2