use std::collections::HashSet;
use std::fmt::Write;
use crate::ir::{ByteOrder, Ir, Presence, PrimitiveType, Signal, Token};
use crate::structured_ir::*;
use crate::{GenerationConfig, Schema};
pub(crate) mod conversion_helpers;
pub(crate) use conversion_helpers::*;
pub(crate) mod conversion_traits;
pub(crate) use conversion_traits::*;
pub(crate) mod converter_impls;
pub(crate) use converter_impls::generate_converter_impls;
pub(crate) mod decoder_display;
pub(crate) use decoder_display::generate_decoder_display;
pub(crate) mod domain_cluster;
pub(crate) use domain_cluster::*;
pub(crate) mod encoded_length;
pub(crate) mod field_type;
pub(crate) use field_type::field_type_ident;
pub(crate) mod message_header_template;
pub(crate) use message_header_template::*;
pub(crate) mod nullification;
pub(crate) use nullification::*;
pub(crate) mod runtime;
use quote::format_ident;
pub(crate) use runtime::*;
pub(crate) mod group_encoder;
pub(crate) use group_encoder::generate_group_encoder;
pub(crate) mod group_decoder;
pub(crate) use group_decoder::generate_group_decoder;
pub(crate) mod tail_stages;
pub(crate) use tail_stages::*;
pub(crate) mod message_decoder;
pub(crate) use message_decoder::generate_message_decoder;
pub(crate) mod message_encoder;
pub(crate) use message_encoder::generate_message_encoder;
use sha2::{Digest, Sha256};
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct GeneratedModule {
pub path: String,
pub source: String,
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct GeneratedModuleSet {
modules: Vec<GeneratedModule>,
warnings: Vec<String>,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum GenerateError {
HeaderValueOutOfRange {
field: String,
value: u64,
maximum: u64,
context: String,
},
InvalidConversion {
selector: String,
reason: String,
},
ConversionCollision {
method: String,
selector_a: String,
selector_b: String,
},
}
impl core::fmt::Display for GenerateError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::HeaderValueOutOfRange {
field,
value,
maximum,
context,
} => {
write!(
f,
"message header field '{field}' value {value} for {context} exceeds declared maximum {maximum}"
)
}
Self::InvalidConversion { selector, reason } => {
write!(f, "invalid conversion '{selector}': {reason}")
}
Self::ConversionCollision {
method,
selector_a,
selector_b,
} => {
write!(
f,
"conversion method collision: '{method}' from '{selector_a}' and '{selector_b}'"
)
}
}
}
}
impl core::error::Error for GenerateError {}
impl GeneratedModuleSet {
pub(crate) fn push(&mut self, module: GeneratedModule) {
self.modules.push(module);
}
#[must_use]
pub fn modules(&self) -> impl ExactSizeIterator<Item = &GeneratedModule> {
self.modules.iter()
}
#[must_use]
pub fn warnings(&self) -> &[String] {
&self.warnings
}
}
#[allow(missing_docs)]
pub(crate) struct GenerationContext {
pub elements: SchemaElements,
pub byte_order: ByteOrder,
pub schema_id: u16,
pub schema_version: u16,
pub header_type: String,
pub header_size: usize,
pub schema_name: String,
pub multi_message: bool,
pub conversions: Vec<crate::ConversionSelector>,
pub domain_types: Vec<(crate::ConversionSelector, String)>,
pub unchecked_companions: bool,
pub domain_objects: bool,
pub domain_var_data: crate::config::DomainVarData,
}
impl GenerationContext {
fn from_schema(schema: &Schema, config: &GenerationConfig, multi_message: bool) -> Self {
let elements = partition_tokens(&schema.ir.tokens);
let header_size = elements
.composites
.iter()
.find(|c| c[0].name == schema.ir.header_type)
.and_then(|c| c[0].encoding.offset)
.unwrap_or(8);
Self {
elements,
byte_order: schema.ir.byte_order,
schema_id: schema.ir.id,
schema_version: schema.ir.version,
header_type: schema.ir.header_type.clone(),
header_size,
schema_name: schema.ir.package.clone(),
multi_message,
conversions: config.conversions.clone(),
domain_types: config.domain_types.clone(),
unchecked_companions: config.unchecked_companions,
domain_objects: config.domain_objects,
domain_var_data: config.domain_var_data,
}
}
}
#[derive(Debug)]
pub struct Generator {
config: GenerationConfig,
}
impl Generator {
#[must_use]
pub const fn new(config: GenerationConfig) -> Self {
Self { config }
}
fn validate_header_values(&self, schema: &Schema) -> Result<(), GenerateError> {
let elements = partition_tokens(&schema.ir.tokens);
let Some(header) = elements
.composites
.iter()
.find(|tokens| tokens[0].name == schema.ir.header_type)
else {
return Ok(());
};
let check = |field_name: &str, value: u64, context: String| -> Result<(), GenerateError> {
let Some(field) = header
.iter()
.find(|token| token.signal == Signal::BeginField && token.name == field_name)
else {
return Ok(());
};
if field.encoding.presence == Presence::Constant {
return Ok(());
}
let maximum = field.encoding.max_value.unwrap_or(u64::MAX);
if value > maximum {
return Err(GenerateError::HeaderValueOutOfRange {
field: field_name.to_string(),
value,
maximum,
context,
});
}
Ok(())
};
let schema_context = format!("schema '{}'", schema.package);
check("schemaId", u64::from(schema.id), schema_context.clone())?;
check("version", u64::from(schema.version), schema_context)?;
for message_tokens in &elements.messages {
let message = parse_message_structure(message_tokens, &elements);
let context = format!("message '{}'", message.name);
check("templateId", u64::from(message.id), context.clone())?;
let block_length = u64::try_from(message.block_length).unwrap_or(u64::MAX);
check("blockLength", block_length, context)?;
}
Ok(())
}
fn validate_conversions(&self, schema: &Schema) -> Result<(), GenerateError> {
if !self.config.has_conversions() {
return Ok(());
}
let elements = partition_tokens(&schema.ir.tokens);
for sel in &self.config.conversions {
let matched = match sel {
crate::ConversionSelector::NamedType(name) => {
elements.composites.iter().any(|c| c[0].name == *name)
|| elements.enums.iter().any(|e| e[0].name == *name)
|| elements.sets.iter().any(|s| s[0].name == *name)
}
crate::ConversionSelector::SemanticType(_) => {
true
}
crate::ConversionSelector::FieldPath(_) => {
true
}
};
if !matched {
return Err(GenerateError::InvalidConversion {
selector: format!("{sel:?}"),
reason: "no matching type found in schema".into(),
});
}
}
for (sel, rust_type) in &self.config.domain_types {
if rust_type.is_empty() {
return Err(GenerateError::InvalidConversion {
selector: format!("{sel:?}"),
reason: "domain type path must not be empty".into(),
});
}
}
Ok(())
}
fn field_has_conversion(
field: &MessageField,
conversions: &[crate::ConversionSelector],
) -> bool {
field_has_conversion_free(field, conversions)
}
fn has_conversion_for(
&self,
type_name: &str,
semantic_type: Option<&str>,
owner_name: Option<&str>,
field_name: &str,
) -> bool {
for sel in &self.config.conversions {
match sel {
crate::ConversionSelector::NamedType(name) if name == type_name => return true,
crate::ConversionSelector::SemanticType(st)
if semantic_type == Some(st.as_str()) =>
{
return true;
}
crate::ConversionSelector::FieldPath(path) => {
let expected = format!("{}.{}", owner_name.unwrap_or(""), field_name);
if path == &expected || path == field_name {
return true;
}
}
_ => {}
}
}
false
}
#[allow(missing_docs)]
fn effective_domain_types(
&self,
schemas: &[(&Schema, &str)],
) -> Vec<(crate::ConversionSelector, String)> {
let mut types = self.config.domain_types.clone();
if self.config.auto_bool_domain {
for (schema, _) in schemas {
let elements = partition_tokens(&schema.ir.tokens);
for e in &elements.enums {
let name = &e[0].name;
if crate::structured_ir::is_bool_value_enum(&elements, name) {
let sel = crate::ConversionSelector::named_type(name);
if !types.iter().any(|(s, _)| s == &sel) {
types.push((sel, "bool".into()));
}
}
}
}
}
types
}
pub fn generate(&self, schema: &Schema) -> Result<GeneratedModuleSet, GenerateError> {
let effective = self.effective_domain_types(&[(schema, "")]);
with_keyword_append(&self.config.keyword_append_token, || {
with_deprecated_attrs(self.config.deprecated_attrs, || {
self.validate_header_values(schema)?;
self.validate_conversions(schema)?;
let mut modules = GeneratedModuleSet::default();
let src = self.gen_schema(schema, &HashSet::new(), false, true, &effective);
modules.push(GeneratedModule {
path: format!("{}.rs", self.config.module_name),
source: src,
});
Ok(modules)
})
})
}
pub fn generate_multi(
&self,
schemas: &[(&Schema, &str)],
) -> Result<GeneratedModuleSet, GenerateError> {
let effective = self.effective_domain_types(schemas);
with_keyword_append(&self.config.keyword_append_token, || {
with_deprecated_attrs(self.config.deprecated_attrs, || {
self.generate_multi_inner(schemas, &effective)
})
})
}
fn generate_multi_inner(
&self,
schemas: &[(&Schema, &str)],
domain_types: &[(crate::ConversionSelector, String)],
) -> Result<GeneratedModuleSet, GenerateError> {
let mut modules = GeneratedModuleSet::default();
let mut shared_types: HashSet<String> = HashSet::new();
let empty_set: HashSet<String> = HashSet::new();
for (i, (schema, module_name)) in schemas.iter().enumerate() {
self.validate_header_values(schema)?;
if i == 0 {
let elements = partition_tokens(&schema.ir.tokens);
for et in &elements.enums {
let name = to_pascal_case(&et[0].name);
shared_types.insert(name.clone());
if let Some(warn) = warn_version_gated(&name, et, schema) {
modules.warnings.push(warn);
}
}
for st in &elements.sets {
shared_types.insert(to_pascal_case(&st[0].name));
}
for ct in &elements.composites {
let name = to_pascal_case(&ct[0].name);
shared_types.insert(name.clone());
if let Some(warn) = warn_version_gated(&name, ct, schema) {
modules.warnings.push(warn);
}
}
}
let is_importing = i > 0 && self.config.shared_module.is_some();
let emit_sbe_rt = i == 0 || self.config.shared_module.is_none();
let skip_set: &HashSet<String> = if self.config.shared_module.is_some() && i > 0 {
&shared_types
} else {
&empty_set
};
let src = self.gen_schema(schema, skip_set, is_importing, emit_sbe_rt, domain_types);
modules.push(GeneratedModule {
path: format!("{}.rs", module_name),
source: src,
});
}
Ok(modules)
}
fn build_enum_ctx<'s>(
tokens: &[crate::ir::Token],
schema: &'s crate::Schema,
) -> crate::ItemContext<'s> {
let name = to_pascal_case(&tokens[0].name);
let encoding_type = tokens[0]
.encoding
.primitive_type
.unwrap_or(PrimitiveType::UInt8);
let et_str = rust_type(encoding_type).to_string();
let variants: Vec<_> = tokens
.iter()
.filter(|t| t.signal == crate::ir::Signal::Encoding)
.filter_map(|t| {
let val = t.encoding.constant_value.as_ref()?;
let value: i128 = if encoding_type == PrimitiveType::Char {
i128::from(val.as_bytes().first().copied().unwrap_or(0))
} else {
val.parse::<i128>().ok()?
};
Some(crate::EnumVariantInfo {
name: to_pascal_case(&t.name),
snake_name: to_snake_case(&t.name),
label: t.name.clone(),
value,
description: t.encoding.description.clone(),
})
})
.collect();
crate::ItemContext::Enum {
schema,
name,
encoding_type: et_str,
variants,
}
}
fn build_message_ctx<'s>(
msg: &MessageStructure,
kind: crate::ItemKind,
schema: &'s crate::Schema,
) -> crate::ItemContext<'s> {
let name = to_pascal_case(&msg.name);
let name_with = |suffix: &str| format!("{name}{suffix}");
let fields = message_field_infos(&msg.fields, &[], None);
let name = match kind {
crate::ItemKind::MessageDecoder => name_with("Decoder"),
crate::ItemKind::MessageEncoder => name_with("Encoder"),
_ => name,
};
match kind {
crate::ItemKind::MessageDecoder => crate::ItemContext::MessageDecoder {
schema,
name,
template_id: msg.id,
block_length: msg.block_length,
fields,
},
crate::ItemKind::MessageEncoder => crate::ItemContext::MessageEncoder {
schema,
name,
template_id: msg.id,
block_length: msg.block_length,
fields,
},
_ => unreachable!("build_message_ctx only for MessageDecoder/MessageEncoder"),
}
}
fn build_composite_ctx<'s>(
tokens: &[crate::ir::Token],
schema: &'s crate::Schema,
) -> crate::ItemContext<'s> {
use crate::structured_ir::MemberType;
let name = to_pascal_case(&tokens[0].name);
let member_field_token = |member_name: &str| {
tokens
.iter()
.find(|t| t.signal == crate::ir::Signal::BeginField && t.name == member_name)
};
let inner_type_token = |field_name: &str| {
let mut it = tokens.iter().skip_while(|t| {
!(t.signal == crate::ir::Signal::BeginField && t.name == field_name)
});
let _ = it.next(); it.find(|t| {
matches!(
t.signal,
crate::ir::Signal::Encoding
| crate::ir::Signal::BeginComposite
| crate::ir::Signal::BeginEnum
| crate::ir::Signal::BeginSet
)
})
};
let fields: Vec<_> = crate::structured_ir::parse_composite_members(tokens)
.into_iter()
.map(|m| {
let field_tok = member_field_token(&m.name);
let inner_tok = inner_type_token(&m.name);
let enc = match &m.member_type {
MemberType::Primitive { .. } => field_tok.map(|t| &t.encoding),
MemberType::Composite { .. }
| MemberType::Enum { .. }
| MemberType::Set { .. } => inner_tok.map(|t| &t.encoding),
};
let (rust_type, presence) = match &m.member_type {
MemberType::Primitive {
prim,
length,
presence,
..
} => {
let base = crate::structured_ir::rust_type(*prim);
let rt = match length {
Some(len) => format!("[{base}; {len}]"),
None => base.to_string(),
};
let ps = match presence {
crate::ir::Presence::Optional => "optional",
crate::ir::Presence::Constant => "constant",
crate::ir::Presence::Required => "required",
};
(rt, ps)
}
MemberType::Composite { name, .. } => (to_pascal_case(name), "required"),
MemberType::Enum { name, .. } => (to_pascal_case(name), "required"),
MemberType::Set { name, .. } => (to_pascal_case(name), "required"),
};
crate::FieldInfo {
name: to_snake_case(&m.name),
rust_type,
offset: Some(m.offset),
since_version: m.since_version,
semantic_type: enc.and_then(|e| e.semantic_type.clone()),
presence,
null_value: enc.and_then(|e| e.null_value),
deprecated: enc.is_some_and(|e| e.deprecated),
description: enc.and_then(|e| e.description.clone()),
}
})
.collect();
crate::ItemContext::Composite {
schema,
name,
fields,
}
}
fn build_set_ctx<'s>(
tokens: &[crate::ir::Token],
schema: &'s crate::Schema,
) -> crate::ItemContext<'s> {
let name = to_pascal_case(&tokens[0].name);
let encoding_type = tokens[0]
.encoding
.primitive_type
.unwrap_or(PrimitiveType::UInt8);
let et_str = rust_type(encoding_type).to_string();
let choices: Vec<_> = tokens
.iter()
.filter(|t| t.signal == crate::ir::Signal::Encoding)
.map(|t| crate::SetChoiceInfo {
name: to_pascal_case(&t.name),
snake_name: to_snake_case(&t.name),
label: t.name.clone(),
bit_position: t
.encoding
.constant_value
.as_ref()
.and_then(|v| v.parse::<u8>().ok())
.unwrap_or(0),
description: t.encoding.description.clone(),
})
.collect();
crate::ItemContext::Set {
schema,
name,
encoding_type: et_str,
choices,
}
}
fn run_hooks(&self, ctx: &crate::ItemContext, src: &mut String) {
if !self.config.has_hooks() {
return;
}
self.config.run_hooks(ctx, src);
}
fn gen_schema(
&self,
schema: &Schema,
shared: &HashSet<String>,
is_importing: bool,
emit_sbe_rt: bool,
domain_types: &[(crate::ConversionSelector, String)],
) -> String {
let ir = &schema.ir;
let mut src = String::new();
writeln!(
src,
"/// Generated from SBE schema package `{}` id {} version {}.",
schema.package, schema.id, schema.version
)
.unwrap();
src.push_str(
"#[allow(clippy::absurd_extreme_comparisons, clippy::double_must_use, \
clippy::erasing_op, clippy::identity_op, clippy::unnecessary_cast, \
unused_assignments, unused_comparisons)]\n",
);
src.push_str("#[allow(non_camel_case_types)]\n");
src.push_str("#[allow(non_snake_case)]\n");
src.push_str("#[allow(clippy::identity_op)]\n");
src.push_str("#[allow(clippy::eq_op)]\n");
src.push_str("#[allow(clippy::needless_borrow)]\n");
src.push_str("#[allow(clippy::manual_range_contains)]\n");
src.push_str("#[allow(unused_imports)]\n");
src.push_str("#[allow(unused_variables)]\n");
src.push_str("#[allow(unused_mut)]\n");
src.push_str("#[allow(dead_code)]\n\n");
if is_importing {
if let Some(ref shared_mod) = self.config.shared_module {
write!(src, "pub use super::{}::*;\n\n", shared_mod).unwrap();
}
}
if let Some(ref ext) = self.config.external_sbe_rt_path {
let _ = writeln!(src, "pub use {ext} as sbe_rt;\n");
if self.config.has_conversions() {
emit_conversion_traits(&mut src);
}
} else if emit_sbe_rt {
src.push_str(&generate_sbe_rt_src());
if self.config.has_conversions() {
emit_conversion_traits(&mut src);
}
}
let elements = partition_tokens(&ir.tokens);
for enum_tokens in &elements.enums {
let type_name = to_pascal_case(&enum_tokens[0].name);
if shared.contains(&type_name) {
continue;
}
generate_enum(&mut src, enum_tokens);
if self.config.has_hooks() {
let ctx = Self::build_enum_ctx(enum_tokens, schema);
self.run_hooks(&ctx, &mut src);
}
}
for set_tokens in &elements.sets {
let type_name = to_pascal_case(&set_tokens[0].name);
if shared.contains(&type_name) {
continue;
}
generate_set(&mut src, set_tokens);
if self.config.has_hooks() {
let ctx = Self::build_set_ctx(set_tokens, schema);
self.run_hooks(&ctx, &mut src);
}
}
for composite_tokens in &elements.composites {
let type_name = to_pascal_case(&composite_tokens[0].name);
if shared.contains(&type_name) {
continue;
}
let comp_byte_order = ir.byte_order;
generate_composite(&mut src, composite_tokens, comp_byte_order);
if self.config.has_hooks() {
let ctx = Self::build_composite_ctx(composite_tokens, schema);
self.run_hooks(&ctx, &mut src);
}
}
let header_pascal = to_pascal_case(&ir.header_type);
if header_pascal != "MessageHeader" && !shared.contains(&header_pascal) {
write!(src, "pub type MessageHeader = {};\n\n", header_pascal).unwrap();
}
let messages: Vec<MessageStructure> = elements
.messages
.iter()
.map(|toks| parse_message_structure(toks, &elements))
.collect();
let mut schema_markers = occupied_type_names(&elements);
let mut message_markers: Vec<(String, String)> = Vec::new();
for msg in &messages {
let multi = messages.len() > 1;
let (decoder_ts, marker) = generate_message_decoder(
msg,
&elements,
&mut schema_markers,
ir.byte_order,
ir.id,
ir.version,
&ir.header_type,
&ir.package,
multi,
self.config.domain_objects,
self.config.domain_var_data,
&self.config.conversions,
domain_types,
self.config.unchecked_companions,
&self.config.hooks,
schema,
);
src.push_str(&decoder_ts.to_string());
src.push('\n');
message_markers.push((to_pascal_case(&msg.name), marker));
if self.config.has_hooks() {
let ctx = Self::build_message_ctx(msg, crate::ItemKind::MessageDecoder, schema);
self.run_hooks(&ctx, &mut src);
}
let encoder_ts = generate_message_encoder(
msg,
&elements,
ir.byte_order,
ir.id,
ir.version,
&ir.header_type,
multi,
&self.config.conversions,
domain_types,
self.config.unchecked_companions,
);
src.push_str(&encoder_ts.to_string());
if self.config.has_hooks() {
let ctx = Self::build_message_ctx(msg, crate::ItemKind::MessageEncoder, schema);
self.run_hooks(&ctx, &mut src);
}
if !&self.config.conversions.is_empty() {
let converter_ts =
generate_converter_impls(msg, &self.config.conversions, domain_types, multi);
src.push_str(&converter_ts);
}
src.push('\n');
generate_message_field_meta(&mut src, msg);
}
if self.config.has_conversions() {
let impl_blocks =
generate_conversion_impl_blocks(&elements, &self.config.conversions, domain_types);
src.push_str(&impl_blocks);
}
{
let has_staged = messages.iter().any(|m| {
matches!(
encoded_length::strategy(m),
encoded_length::LengthStrategy::Staged
)
});
if has_staged {
let support_ts = encoded_length::generate_support();
src.push_str(&support_ts.to_string());
}
}
if let Some(ref sem_ver) = schema.ir.semantic_version {
write!(
src,
"pub const SEMANTIC_VERSION: &str = \"{}\";\n\n",
sem_ver
)
.unwrap();
}
let schema_hash = compute_schema_hash(&schema.package, schema.id, schema.version);
write!(src, "pub const SCHEMA_HASH: u64 = {};\n\n", schema_hash).unwrap();
let sha256_hash = compute_schema_sha256(&schema.ir);
src.push_str("pub const SCHEMA_SHA256: [u8; 32] = [");
for (i, &b) in sha256_hash.iter().enumerate() {
if i > 0 {
src.push_str(", ");
}
write!(src, "0x{:02x}", b).unwrap();
}
src.push_str("];\n\n");
let hex: String = sha256_hash.iter().map(|b| format!("{:02x}", b)).collect();
write!(src, "pub const SCHEMA_SHA256_HEX: &str = \"{}\";\n\n", hex).unwrap();
generate_prelude(&mut src, &elements, &messages, ir.id, ir.version);
if let Some(ref err_path) = self.config.error_from_path {
let err_ty: syn::Type = syn::parse_str(err_path).expect("invalid error_from_path");
let span = proc_macro2::Span::call_site();
let impls = quote::quote! {
impl From<sbe_rt::EncodeError> for #err_ty {
fn from(e: sbe_rt::EncodeError) -> Self {
Self::from(format!("sbe encode: {e}"))
}
}
impl From<sbe_rt::DecodeError> for #err_ty {
fn from(e: sbe_rt::DecodeError) -> Self {
Self::from(format!("sbe decode: {e}"))
}
}
};
src.push_str(&impls.to_string());
src.push('\n');
}
let read_bytes_ts: proc_macro2::TokenStream = quote::quote! {
#[inline]
pub fn read_bytes<const N: usize>(buf: &[u8], offset: usize) -> [u8; N] {
buf[offset..offset + N].try_into().expect("read_bytes: buffer too short")
}
#[inline]
pub fn write_bytes<const N: usize>(buf: &mut [u8], offset: usize, bytes: &[u8; N]) {
buf[offset..offset + N].copy_from_slice(bytes);
}
};
src.push_str(&read_bytes_ts.to_string());
let uc = quote::quote! {
#[inline]
pub fn read_bytes_unchecked<const N: usize>(buf: &[u8], offset: usize) -> [u8; N] {
unsafe {
core::ptr::read_unaligned(buf.as_ptr().add(offset) as *const [u8; N])
}
}
#[inline]
pub fn write_bytes_unchecked<const N: usize>(buf: &mut [u8], offset: usize, bytes: &[u8; N]) {
unsafe {
core::ptr::write_unaligned(buf.as_mut_ptr().add(offset) as *mut [u8; N], *bytes)
}
}
};
src.push_str(&uc.to_string());
src.push('\n');
generate_schema_id_from_header(&mut src, &elements, &ir.header_type, ir.byte_order);
let any_msg_ts = generate_any_message(
&messages,
&elements,
ir.id,
&ir.header_type,
&ir.package,
&message_markers,
);
src.push_str(&any_msg_ts.to_string());
src.push('\n');
let file =
syn::parse_str::<syn::File>(&src).expect("generated code must be valid Rust syntax");
prettyplease::unparse(&file)
}
}
#[cfg(test)]
mod tests {
use super::Generator;
use crate::{GenerationConfig, Schema};
#[test]
fn generator_emits_deterministic_module_name() -> Result<(), Box<dyn std::error::Error>> {
let mut generator = Generator::new(GenerationConfig::new("market_data"));
let schema = Schema::new("fix.sbe", 1, 0);
let modules = generator.generate(&schema)?;
let collected = modules.modules().collect::<Vec<_>>();
assert_eq!(collected.len(), 1);
assert_eq!(collected[0].path, "market_data.rs");
assert!(collected[0].source.contains("fix.sbe"));
Ok(())
}
#[test]
fn generate_multi_creates_separate_modules() -> Result<(), Box<dyn std::error::Error>> {
let mut config = GenerationConfig::new("common");
config.shared_module = Some("common_types".to_string());
let mut generator = Generator::new(config);
let schema_a = Schema::new("common.sbe", 1, 0);
let schema_b = Schema::new("market_data.sbe", 2, 0);
let modules =
generator.generate_multi(&[(&schema_a, "common_types"), (&schema_b, "market_data")])?;
let collected: Vec<_> = modules.modules().collect();
assert_eq!(collected.len(), 2);
assert_eq!(collected[0].path, "common_types.rs");
assert_eq!(collected[1].path, "market_data.rs");
assert!(collected[0].source.contains("pub mod sbe_rt"));
assert!(!collected[1].source.contains("pub mod sbe_rt"));
assert!(
collected[1]
.source
.contains("pub use super::common_types::*;")
);
assert!(collected[0].source.contains("common.sbe"));
assert!(collected[1].source.contains("market_data.sbe"));
Ok(())
}
#[test]
fn generate_multi_without_shared_module_emits_sbe_rt_everywhere()
-> Result<(), Box<dyn std::error::Error>> {
let config = GenerationConfig::new("common");
let mut generator = Generator::new(config);
let schema_a = Schema::new("common.sbe", 1, 0);
let schema_b = Schema::new("market_data.sbe", 2, 0);
let modules = generator.generate_multi(&[(&schema_a, "a_mod"), (&schema_b, "b_mod")])?;
let collected: Vec<_> = modules.modules().collect();
assert_eq!(collected.len(), 2);
assert!(collected[0].source.contains("pub mod sbe_rt"));
assert!(collected[1].source.contains("pub mod sbe_rt"));
assert!(!collected[1].source.contains("\npub use super::"));
Ok(())
}
use super::{
SchemaElements, parse_composite_members, parse_field_structure, parse_group_structure,
parse_message_structure, parse_vardata_structure, to_snake_case,
};
use crate::ir::{Encoding, Signal, Token};
fn make_token(signal: Signal) -> Token {
Token {
id: None,
name: String::new(),
signal,
encoding: Encoding::default(),
span: None,
}
}
fn empty_elements() -> SchemaElements {
SchemaElements {
composites: vec![],
enums: vec![],
sets: vec![],
messages: vec![],
}
}
#[test]
fn message_structure_skips_unexpected_signal() -> Result<(), Box<dyn std::error::Error>> {
let elem = empty_elements();
let _ = parse_message_structure(
&[
make_token(Signal::BeginMessage),
make_token(Signal::BeginEnum), make_token(Signal::EndMessage),
],
&elem,
);
Ok(())
}
#[test]
fn group_structure_skips_unexpected_signal() -> Result<(), Box<dyn std::error::Error>> {
let elem = empty_elements();
let _ = parse_group_structure(
&[
make_token(Signal::BeginGroup),
make_token(Signal::BeginMessage), make_token(Signal::EndGroup),
],
&elem,
);
Ok(())
}
#[test]
fn vardata_structure_skips_non_length_fields() -> Result<(), Box<dyn std::error::Error>> {
let _ = parse_vardata_structure(&[
make_token(Signal::BeginComposite),
make_token(Signal::BeginField),
make_token(Signal::EndField),
make_token(Signal::EndComposite),
]);
Ok(())
}
#[test]
fn composite_members_skips_non_field_signals() -> Result<(), Box<dyn std::error::Error>> {
let _ = parse_composite_members(&[
make_token(Signal::BeginComposite),
make_token(Signal::BeginMessage), make_token(Signal::EndComposite),
]);
Ok(())
}
#[test]
fn field_structure_falls_back_to_uint8_primitive() -> Result<(), Box<dyn std::error::Error>> {
let elem = empty_elements();
let _ = parse_field_structure(
&[
make_token(Signal::BeginField),
make_token(Signal::BeginMessage), make_token(Signal::EndField),
],
&elem,
);
Ok(())
}
#[test]
fn group_array_codegen_uses_the_complete_field_extent_and_element_range()
-> Result<(), Box<dyn std::error::Error>> {
let xml = r#"<?xml version="1.0"?>
<messageSchema package="array.guard" id="305" version="1" byteOrder="littleEndian">
<types>
<composite name="messageHeader">
<type name="blockLength" primitiveType="uint16"/>
<type name="templateId" primitiveType="uint16"/>
<type name="schemaId" primitiveType="uint16"/>
<type name="version" primitiveType="uint16"/>
</composite>
<composite name="groupSizeEncoding">
<type name="blockLength" primitiveType="uint16"/>
<type name="numInGroup" primitiveType="uint16"/>
</composite>
<type name="Values" primitiveType="uint32" length="2"/>
<enum name="State" encodingType="uint8">
<validValue name="Ready">1</validValue>
</enum>
<set name="Flags" encodingType="uint8">
<choice name="Active">0</choice>
</set>
<enum name="BooleanType" encodingType="uint8">
<validValue name="F">0</validValue>
<validValue name="T">1</validValue>
</enum>
</types>
<message name="ArrayBoundaryMessage" id="1">
<group name="entries" id="1">
<field name="base" id="2" type="uint8"/>
<field name="values" id="3" type="Values"/>
<field name="state" id="4" type="State" sinceVersion="1"/>
<field name="flags" id="5" type="Flags" sinceVersion="1"/>
<field name="enabled" id="6" type="BooleanType" sinceVersion="1"/>
</group>
</message>
</messageSchema>"#;
let schema = crate::Schema::from_ir(crate::parse(xml)?);
let mut generator = crate::Generator::new(crate::GenerationConfig::new("array_guard"));
let modules = generator.generate(&schema)?;
let source = &modules
.modules()
.next()
.ok_or("missing generated module")?
.source;
assert!(
source.contains("|| 9 > self.acting_block_length"),
"u32[2] at offset 1 must require all nine entry bytes"
);
assert!(
source.contains("let all: [u8; 8]"),
"u32[2] must bulk-read exactly eight bytes"
);
assert!(
source.contains("all[0usize]") && source.contains("all[7usize]"),
"the unrolled array decode must use the complete byte range"
);
assert!(
source.contains("|| 10 > self.acting_block_length"),
"the versioned enum at offset nine must require its complete tenth byte"
);
assert!(
source.contains("|| 11 > self.acting_block_length"),
"the versioned set at offset ten must require its complete eleventh byte"
);
assert!(
source.contains("pub fn enabled_bool(&self) -> Option<bool>"),
"a versioned BooleanType group field must preserve absence in its bool accessor"
);
Ok(())
}
#[test]
fn snake_case_handles_empty_or_special_input() -> Result<(), Box<dyn std::error::Error>> {
assert_eq!(to_snake_case(""), "");
assert_eq!(to_snake_case("Foo__Bar"), "foo_bar");
Ok(())
}
#[test]
fn partition_skips_unexpected_at_top_level() -> Result<(), Box<dyn std::error::Error>> {
let _ = super::partition_tokens(&[make_token(Signal::BeginField)]);
Ok(())
}
#[test]
fn partition_skips_unexpected_in_message_body() -> Result<(), Box<dyn std::error::Error>> {
let _ = super::partition_tokens(&[
make_token(Signal::BeginMessage),
make_token(Signal::BeginEnum), make_token(Signal::EndMessage),
]);
Ok(())
}
#[test]
fn partition_skips_unexpected_in_group_body() -> Result<(), Box<dyn std::error::Error>> {
let _ = super::partition_tokens(&[
make_token(Signal::BeginGroup),
make_token(Signal::BeginMessage), make_token(Signal::EndGroup),
]);
Ok(())
}
#[test]
fn partition_skips_unexpected_after_top_level_items() -> Result<(), Box<dyn std::error::Error>>
{
let _ = super::partition_tokens(&[
make_token(Signal::BeginMessage),
make_token(Signal::EndMessage),
make_token(Signal::BeginEnum), ]);
Ok(())
}
#[test]
fn semantic_type_matches_primitive_field() -> Result<(), Box<dyn std::error::Error>> {
use crate::ir::{Presence, PrimitiveType};
use crate::structured_ir::{FieldType, MessageField};
let field = MessageField {
name: "exchangeTimestamp".into(),
id: Some(1),
offset: 0,
presence: Presence::Required,
since_version: 0,
null_value: None,
min_value: None,
max_value: None,
description: None,
deprecated: false,
semantic_type: Some("UTCTimestamp".into()),
constant_value: None,
epoch: None,
time_unit: None,
character_encoding: None,
field_type: FieldType::Primitive(PrimitiveType::UInt64, None),
};
let conversions = vec![crate::ConversionSelector::semantic_type("UTCTimestamp")];
assert!(
super::field_has_conversion_free(&field, &conversions),
"SemanticType should match primitive u64 with semanticType=UTCTimestamp"
);
let domain_types = vec![(
crate::ConversionSelector::semantic_type("UTCTimestamp"),
"chrono::DateTime<chrono::Utc>".into(),
)];
let dt = super::find_domain_type(&field, &domain_types);
assert_eq!(
dt,
Some("chrono::DateTime<chrono::Utc>"),
"should find domain type for UTCTimestamp"
);
Ok(())
}
#[test]
fn chrono_converter_generates_accessor() -> Result<(), Box<dyn std::error::Error>> {
let xml = r#"<?xml version="1.0" encoding="UTF-8"?>
<sbe:messageSchema xmlns:sbe="http://fixprotocol.io/2016/sbe"
package="test.chrono" id="1" version="0" byteOrder="littleEndian">
<types>
<composite name="messageHeader">
<type name="blockLength" primitiveType="uint16"/>
<type name="templateId" primitiveType="uint16"/>
<type name="schemaId" primitiveType="uint16"/>
<type name="version" primitiveType="uint16"/>
</composite>
</types>
<sbe:message name="TsMsg" id="1">
<field name="ts" id="1" type="uint64" semanticType="UTCTimestamp"/>
</sbe:message>
</sbe:messageSchema>"#;
let ir = crate::parse(xml)?;
let schema = crate::Schema::from_ir(ir);
let config = crate::GenerationConfig::new("test_chrono").with_domain_type(
crate::ConversionSelector::semantic_type("UTCTimestamp"),
"chrono::DateTime<chrono::Utc>",
);
let mut generator = crate::Generator::new(config);
let modules = generator.generate(&schema)?;
let src = modules.modules().next().unwrap().source.clone();
assert!(
src.contains("fn ts(&self) -> chrono::DateTime"),
"should generate concrete DateTime accessor for UTCTimestamp field"
);
assert!(
src.contains("fn ts_wire"),
"should rename raw u64 getter to _wire"
);
assert!(
src.contains("impl TryFromSbe<u64> for chrono::DateTime<chrono::Utc>"),
"should generate TryFromSbe impl"
);
Ok(())
}
#[test]
fn narrow_message_header_rejects_values_above_declared_field_maximum() {
fn generate(xml: &str) -> Result<super::GeneratedModuleSet, super::GenerateError> {
let ir = crate::parse(xml).expect("schema should parse before codegen validation");
let schema = crate::Schema::from_ir(ir);
crate::Generator::new(crate::GenerationConfig::new("narrow")).generate(&schema)
}
fn schema(schema_id: u16, version: u16, template_id: u16, block_length: u16) -> String {
format!(
r#"<messageSchema package="test" id="{schema_id}" version="{version}" byteOrder="littleEndian">
<types>
<composite name="messageHeader">
<type name="schemaId" primitiveType="uint8"/>
<type name="version" primitiveType="uint8"/>
<type name="templateId" primitiveType="uint8"/>
<type name="blockLength" primitiveType="uint8"/>
</composite>
</types>
<message name="M" id="{template_id}" blockLength="{block_length}"/>
</messageSchema>"#
)
}
for (xml, field) in [
(schema(255, 1, 1, 0), "schemaId"),
(schema(1, 255, 1, 0), "version"),
(schema(1, 1, 255, 0), "templateId"),
(schema(1, 1, 1, 255), "blockLength"),
] {
let error = generate(&xml).expect_err("reserved null/max value must be rejected");
assert!(
error.to_string().contains(field),
"expected {field} error, got: {error}"
);
}
}
#[test]
fn field_named_remaining_is_renamed_to_remaining_field()
-> Result<(), Box<dyn std::error::Error>> {
let xml = r#"<messageSchema package="test" id="1" version="1" byteOrder="littleEndian">
<types>
<composite name="messageHeader">
<type name="blockLength" primitiveType="uint16"/>
<type name="templateId" primitiveType="uint16"/>
<type name="schemaId" primitiveType="uint16"/>
<type name="version" primitiveType="uint16"/>
</composite>
</types>
<message name="Msg" id="1" blockLength="8">
<field name="remaining" id="1" type="int64"/>
</message>
</messageSchema>"#;
let ir = crate::parse(xml).expect("schema should parse");
let schema = crate::Schema::from_ir(ir);
let modules =
crate::Generator::new(crate::GenerationConfig::new("test")).generate(&schema)?;
let src = modules.modules().next().expect("one module").source.clone();
let remaining_count = src.matches("fn remaining(&self)").count();
assert_eq!(
remaining_count, 1,
"expected exactly 2 'remaining' methods (one decoder + one encoder), found {remaining_count}"
);
assert!(
src.contains("fn remaining_field"),
"field accessor 'remaining' must be renamed to 'remaining_field'. src:\n{src}"
);
Ok(())
}
#[test]
fn hook_adds_serde_impls_for_enum_and_set() -> Result<(), Box<dyn std::error::Error>> {
let xml = r#"<messageSchema package="test" id="1" version="0" byteOrder="littleEndian">
<types>
<composite name="messageHeader">
<type name="blockLength" primitiveType="uint16"/>
<type name="templateId" primitiveType="uint16"/>
<type name="schemaId" primitiveType="uint16"/>
<type name="version" primitiveType="uint16"/>
</composite>
<enum name="EventCode" encodingType="uint32">
<validValue name="Ok" description="Success">200</validValue>
<validValue name="Error" description="Failure">400</validValue>
<validValue name="Timeout">408</validValue>
</enum>
<set name="OptionalFields" encodingType="uint8">
<choice name="hasPrice">0</choice>
<choice name="hasQty">1</choice>
<choice name="hasVenue">2</choice>
</set>
</types>
<message name="Msg" id="1" blockLength="0"/>
</messageSchema>"#;
use crate::{EnumVariantInfo, ItemContext, ItemKind, SetChoiceInfo};
use quote::format_ident;
let config = crate::GenerationConfig::new("test")
.with_hook(|ctx: &ItemContext| -> Vec<proc_macro2::TokenStream> {
match ctx {
ItemContext::Enum { name, variants, .. } => {
let ident = format_ident!("{name}");
let var_names: Vec<_> = variants.iter().map(|v| format_ident!("{}", v.name)).collect();
let var_labels: Vec<_> = variants.iter().map(|v| v.name.clone()).collect();
vec![quote::quote! {
impl serde::Serialize for #ident {
fn serialize<S: serde::Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
let label = match self {
#(Self::#var_names => #var_labels,)*
};
s.serialize_str(label)
}
}
impl<'de> serde::Deserialize<'de> for #ident {
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
let s = <&str>::deserialize(d)?;
match s {
#(#var_labels => Ok(Self::#var_names),)*
_ => Err(serde::de::Error::unknown_variant(s, &[#(#var_labels),*])),
}
}
}
}]
}
ItemContext::Set { name, encoding_type, choices, .. } => {
let ident = format_ident!("{name}");
let c_getters: Vec<_> = choices
.iter()
.map(|c| format_ident!("is_{}", c.snake_name))
.collect();
let c_labels: Vec<_> = choices.iter().map(|c| c.label.clone()).collect();
let c_bits: Vec<_> = choices.iter().map(|c| c.bit_position).collect();
let acc_ty: syn::Type = syn::parse_str(encoding_type).unwrap();
vec![quote::quote! {
impl serde::Serialize for #ident {
fn serialize<S: serde::Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
let mut names = Vec::new();
#(if self.#c_getters() { names.push(#c_labels); })*
names.serialize(s)
}
}
impl<'de> serde::Deserialize<'de> for #ident {
fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
let names: Vec<String> = Vec::deserialize(d)?;
let mut value: u64 = 0;
for name in &names {
match name.as_str() {
#(#c_labels => value |= 1u64 << #c_bits,)*
other => return Err(serde::de::Error::unknown_variant(
other, &[#(#c_labels),*])),
}
}
Ok(Self(value as #acc_ty))
}
}
}]
}
_ => vec![],
}
});
let ir = crate::parse(xml).expect("schema should parse");
let schema = crate::Schema::from_ir(ir);
let modules = crate::Generator::new(config).generate(&schema)?;
let src = modules.modules().next().expect("one module").source.clone();
assert!(
src.contains("impl serde::Serialize for EventCode"),
"missing Serialize for enum"
);
assert!(src.contains("\"Ok\""), "missing Ok label");
assert!(src.contains("\"Error\""), "missing Error label");
assert!(
src.contains("impl<'de> serde::Deserialize<'de> for EventCode"),
"missing Deserialize for enum"
);
assert!(
src.contains("unknown_variant"),
"missing error handling in Deserialize"
);
assert!(
src.contains("impl serde::Serialize for OptionalFields"),
"missing Serialize for set"
);
assert!(src.contains("\"hasPrice\""), "missing hasPrice label");
assert!(
src.contains("impl<'de> serde::Deserialize<'de> for OptionalFields"),
"missing Deserialize for set"
);
Ok(())
}
#[test]
fn with_bool_domain_type_works_with_generate_multi() -> Result<(), Box<dyn std::error::Error>> {
let xml_a = r#"<?xml version="1.0"?>
<messageSchema package="a" id="1" version="0" byteOrder="littleEndian">
<types>
<composite name="messageHeader">
<type name="blockLength" primitiveType="uint16"/>
<type name="templateId" primitiveType="uint16"/>
<type name="schemaId" primitiveType="uint16"/>
<type name="version" primitiveType="uint16"/>
</composite>
<enum name="BooleanType" encodingType="uint8">
<validValue name="F">0</validValue>
<validValue name="T">1</validValue>
</enum>
</types>
<message name="MsgA" id="1" blockLength="1">
<field name="flag" id="1" type="BooleanType" offset="0"/>
</message>
</messageSchema>"#;
let xml_b = r#"<?xml version="1.0"?>
<messageSchema package="b" id="2" version="0" byteOrder="littleEndian">
<types>
<composite name="messageHeader">
<type name="blockLength" primitiveType="uint16"/>
<type name="templateId" primitiveType="uint16"/>
<type name="schemaId" primitiveType="uint16"/>
<type name="version" primitiveType="uint16"/>
</composite>
<enum name="BooleanType" encodingType="uint8">
<validValue name="F">0</validValue>
<validValue name="T">1</validValue>
</enum>
</types>
<message name="MsgB" id="1" blockLength="1">
<field name="enabled" id="1" type="BooleanType" offset="0"/>
</message>
</messageSchema>"#;
let schema_a = Schema::from_ir(crate::parse(xml_a)?);
let schema_b = Schema::from_ir(crate::parse(xml_b)?);
let mut generator = Generator::new(
crate::GenerationConfig::new("common_types")
.with_shared_module("common_types")
.with_bool_domain_type(true),
);
let modules =
generator.generate_multi(&[(&schema_a, "common_types"), (&schema_b, "consumer")])?;
let collected: Vec<_> = modules.modules().collect();
assert_eq!(collected.len(), 2);
let consumer_src = &collected[1].source;
assert!(
consumer_src.contains("fn enabled_bool(&self) -> bool"),
"with_bool_domain_type must produce bool getter in multi-schema; got:\n{consumer_src}",
);
Ok(())
}
}