use serde::{Deserialize, Serialize};
use serde_reflection::{
json_converter::{
DeserializationContext, DeserializationEnvironment, SerializationContext,
SerializationEnvironment, SymbolTableEnvironment,
},
Format, Registry,
};
#[cfg(not(target_arch = "wasm32"))]
use serde_reflection::{Samples, Tracer, TracerConfig};
#[derive(Serialize, Deserialize, Debug, Eq, Clone, PartialEq)]
#[serde(rename_all = "UPPERCASE")]
pub struct Formats {
pub registry: Registry,
pub operation: Format,
pub response: Format,
pub message: Format,
pub event_value: Format,
}
pub trait BcsApplication {
type Abi;
fn formats() -> serde_reflection::Result<Formats>;
#[cfg(not(target_arch = "wasm32"))]
fn pruned_formats() -> Result<Formats, PruneError> {
let mut formats = Self::formats()?;
formats.prune_known_primitives()?;
Ok(formats)
}
}
fn bcs_to_json(
bytes: &[u8],
format: &Format,
registry: &Registry,
) -> bcs::Result<serde_json::Value> {
let context = DeserializationContext {
format: format.clone(),
registry,
environment: &LineraEnvironment,
};
bcs::from_bytes_seed(context, bytes)
}
fn json_to_bcs(
value: &serde_json::Value,
format: &Format,
registry: &Registry,
) -> bcs::Result<Vec<u8>> {
let context = SerializationContext {
value,
format,
registry,
environment: &LineraEnvironment,
};
bcs::to_bytes(&context)
}
fn primitive_to_json<'de, T, D>(deserializer: D) -> Result<serde_json::Value, String>
where
T: serde::Deserialize<'de> + serde::Serialize,
D: serde::Deserializer<'de>,
{
let value = T::deserialize(deserializer).map_err(|error| error.to_string())?;
serde_json::to_value(&value).map_err(|error| error.to_string())
}
fn primitive_from_json<T, S>(value: &serde_json::Value, serializer: S) -> Result<S::Ok, S::Error>
where
T: serde::Serialize + serde::de::DeserializeOwned,
S: serde::Serializer,
{
let value: T = T::deserialize(value).map_err(serde::ser::Error::custom)?;
value.serialize(serializer)
}
macro_rules! known_human_readable_primitives {
($($name:literal => $ty:ty),* $(,)?) => {
pub const KNOWN_PRIMITIVE_NAMES: &[&str] = &[$($name),*];
#[derive(Clone, Copy, Debug, Default)]
pub struct LineraEnvironment;
impl SymbolTableEnvironment for LineraEnvironment {}
impl<'de> DeserializationEnvironment<'de> for LineraEnvironment {
fn deserialize<D>(
&self,
name: String,
deserializer: D,
) -> Result<serde_json::Value, String>
where
D: serde::Deserializer<'de>,
{
match name.as_str() {
$( $name => primitive_to_json::<$ty, D>(deserializer), )*
_ => Err(format!("No external definition available for {name}")),
}
}
}
impl SerializationEnvironment for LineraEnvironment {
fn serialize<S>(
&self,
name: &str,
value: &serde_json::Value,
serializer: S,
) -> Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
match name {
$( $name => primitive_from_json::<$ty, S>(value, serializer), )*
_ => Err(serde::ser::Error::custom(format!(
"No external serializer available for {name}"
))),
}
}
}
#[cfg(not(target_arch = "wasm32"))]
fn expected_primitive_registry() -> serde_reflection::Result<Registry> {
let mut tracer = Tracer::new(
TracerConfig::default()
.record_samples_for_newtype_structs(true)
.record_samples_for_tuple_structs(true),
);
let samples = Samples::new();
$( tracer.trace_type::<$ty>(&samples)?; )*
tracer.trace_type::<crate::linera_base_types::VmRuntime>(&samples)?;
tracer.trace_type::<crate::linera_base_types::BlobType>(&samples)?;
tracer.trace_type::<crate::linera_base_types::GenericApplicationId>(&samples)?;
tracer.registry()
}
};
}
known_human_readable_primitives! {
"CryptoHash" => crate::linera_base_types::CryptoHash,
"AccountOwner" => crate::linera_base_types::AccountOwner,
"Amount" => crate::linera_base_types::Amount,
"Epoch" => crate::linera_base_types::Epoch,
"BlobId" => crate::linera_base_types::BlobId,
"StreamId" => crate::linera_base_types::StreamId,
"ModuleId" => crate::linera_base_types::ModuleId,
"ApplicationId" => crate::linera_base_types::ApplicationId,
}
#[cfg(not(target_arch = "wasm32"))]
#[derive(Debug, thiserror::Error)]
pub enum PruneError {
#[error("failed to compute the canonical primitive formats: {0}")]
Reflection(#[from] serde_reflection::Error),
#[error(
"registry entry for `{name}` does not match the canonical linera-base format; \
refusing to prune"
)]
Mismatch {
name: String,
},
}
impl Formats {
pub fn decode_operation(&self, bytes: &[u8]) -> bcs::Result<serde_json::Value> {
bcs_to_json(bytes, &self.operation, &self.registry)
}
pub fn decode_response(&self, bytes: &[u8]) -> bcs::Result<serde_json::Value> {
bcs_to_json(bytes, &self.response, &self.registry)
}
pub fn decode_message(&self, bytes: &[u8]) -> bcs::Result<serde_json::Value> {
bcs_to_json(bytes, &self.message, &self.registry)
}
pub fn decode_event_value(&self, bytes: &[u8]) -> bcs::Result<serde_json::Value> {
bcs_to_json(bytes, &self.event_value, &self.registry)
}
pub fn encode_operation(&self, value: &serde_json::Value) -> bcs::Result<Vec<u8>> {
json_to_bcs(value, &self.operation, &self.registry)
}
pub fn encode_response(&self, value: &serde_json::Value) -> bcs::Result<Vec<u8>> {
json_to_bcs(value, &self.response, &self.registry)
}
pub fn encode_message(&self, value: &serde_json::Value) -> bcs::Result<Vec<u8>> {
json_to_bcs(value, &self.message, &self.registry)
}
pub fn encode_event_value(&self, value: &serde_json::Value) -> bcs::Result<Vec<u8>> {
json_to_bcs(value, &self.event_value, &self.registry)
}
#[cfg(not(target_arch = "wasm32"))]
pub fn prune_known_primitives(&mut self) -> Result<(), PruneError> {
let expected = expected_primitive_registry()?;
for name in KNOWN_PRIMITIVE_NAMES {
let (Some(actual), Some(expected_format)) =
(self.registry.get(*name), expected.get(*name))
else {
continue;
};
if actual != expected_format {
return Err(PruneError::Mismatch {
name: (*name).to_string(),
});
}
}
for name in KNOWN_PRIMITIVE_NAMES {
self.registry.remove(*name);
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use serde::{Deserialize, Serialize};
use serde_json::json;
use serde_reflection::{Samples, Tracer, TracerConfig};
use super::*;
fn trace_format<T>() -> (Format, Registry)
where
T: Serialize + for<'de> Deserialize<'de>,
{
let mut tracer = Tracer::new(
TracerConfig::default()
.record_samples_for_newtype_structs(true)
.record_samples_for_tuple_structs(true),
);
let samples = Samples::new();
let (format, _) = tracer.trace_type::<T>(&samples).unwrap();
let registry = tracer.registry().unwrap();
(format, registry)
}
#[test]
fn primitive_round_trip() {
let (format, registry) = trace_format::<u64>();
let bytes = bcs::to_bytes(&42u64).unwrap();
let value = bcs_to_json(&bytes, &format, ®istry).unwrap();
assert_eq!(value, json!(42));
}
#[test]
fn struct_round_trip() {
#[derive(Serialize, Deserialize)]
struct Point {
x: i32,
y: i32,
}
let (format, registry) = trace_format::<Point>();
let bytes = bcs::to_bytes(&Point { x: 10, y: -7 }).unwrap();
let value = bcs_to_json(&bytes, &format, ®istry).unwrap();
assert_eq!(value, json!({ "x": 10, "y": -7 }));
}
#[test]
fn enum_unit_and_struct_variants() {
#[derive(Serialize, Deserialize)]
enum Op {
Increment,
Set { value: u64 },
Add(i64, i64),
}
let (format, registry) = trace_format::<Op>();
let bytes = bcs::to_bytes(&Op::Increment).unwrap();
let value = bcs_to_json(&bytes, &format, ®istry).unwrap();
assert_eq!(value, json!({ "Increment": null }));
let bytes = bcs::to_bytes(&Op::Set { value: 99 }).unwrap();
let value = bcs_to_json(&bytes, &format, ®istry).unwrap();
assert_eq!(value, json!({ "Set": { "value": 99 } }));
let bytes = bcs::to_bytes(&Op::Add(2, 3)).unwrap();
let value = bcs_to_json(&bytes, &format, ®istry).unwrap();
assert_eq!(value, json!({ "Add": [2, 3] }));
}
#[test]
fn nested_with_option_and_seq() {
#[derive(Serialize, Deserialize)]
struct Outer {
tag: String,
items: Vec<u32>,
note: Option<String>,
}
let (format, registry) = trace_format::<Outer>();
let value = Outer {
tag: "hello".to_string(),
items: vec![1, 2, 3],
note: None,
};
let bytes = bcs::to_bytes(&value).unwrap();
let json_value = bcs_to_json(&bytes, &format, ®istry).unwrap();
assert_eq!(
json_value,
json!({ "tag": "hello", "items": [1, 2, 3], "note": null })
);
}
#[test]
fn formats_decode_helpers() {
#[derive(Serialize, Deserialize)]
enum Operation {
Ping,
Echo(String),
}
#[derive(Serialize, Deserialize)]
struct Response {
ok: bool,
}
let (operation, op_registry) = trace_format::<Operation>();
let (response, resp_registry) = trace_format::<Response>();
let mut registry = op_registry;
registry.extend(resp_registry);
let (message, _) = trace_format::<()>();
let (event_value, _) = trace_format::<()>();
let formats = Formats {
registry,
operation,
response,
message,
event_value,
};
let op_bytes = bcs::to_bytes(&Operation::Echo("hi".to_string())).unwrap();
assert_eq!(
formats.decode_operation(&op_bytes).unwrap(),
json!({ "Echo": "hi" })
);
let resp_bytes = bcs::to_bytes(&Response { ok: true }).unwrap();
assert_eq!(
formats.decode_response(&resp_bytes).unwrap(),
json!({ "ok": true })
);
let unit_bytes = bcs::to_bytes(&()).unwrap();
assert_eq!(formats.decode_message(&unit_bytes).unwrap(), json!(null));
assert_eq!(
formats.decode_event_value(&unit_bytes).unwrap(),
json!(null)
);
assert_eq!(
formats.encode_operation(&json!({ "Echo": "hi" })).unwrap(),
op_bytes
);
assert_eq!(
formats.encode_response(&json!({ "ok": true })).unwrap(),
resp_bytes
);
assert_eq!(formats.encode_message(&json!(null)).unwrap(), unit_bytes);
assert_eq!(
formats.encode_event_value(&json!(null)).unwrap(),
unit_bytes
);
}
#[test]
fn malformed_bytes_return_error() {
let (format, registry) = trace_format::<u64>();
assert!(bcs_to_json(&[1, 2, 3], &format, ®istry).is_err());
}
#[test]
fn expected_registry_builds() {
let registry = expected_primitive_registry().unwrap();
for name in KNOWN_PRIMITIVE_NAMES {
assert!(registry.contains_key(*name), "missing {name}");
}
}
#[test]
fn known_primitives_decode_as_human_readable() {
use std::str::FromStr as _;
use crate::linera_base_types::{AccountOwner, Amount, CryptoHash, ModuleId, VmRuntime};
#[derive(Serialize, Deserialize)]
struct Sample {
owner: AccountOwner,
amount: Amount,
hash: CryptoHash,
module: Option<ModuleId>,
}
let hash = CryptoHash::from_str(&"ab".repeat(32)).unwrap();
let value = Sample {
owner: AccountOwner::Address32(hash),
amount: Amount::from_tokens(5),
hash,
module: Some(ModuleId::new(hash, hash, VmRuntime::Wasm)),
};
let mut tracer = Tracer::new(
TracerConfig::default()
.record_samples_for_newtype_structs(true)
.record_samples_for_tuple_structs(true),
);
let samples = Samples::new();
let (operation, _) = tracer.trace_type::<Sample>(&samples).unwrap();
tracer.trace_type::<AccountOwner>(&samples).unwrap();
tracer.trace_type::<VmRuntime>(&samples).unwrap();
let registry = tracer.registry().unwrap();
let unit = Format::Unit;
let mut formats = Formats {
registry,
operation,
response: unit.clone(),
message: unit.clone(),
event_value: unit,
};
let bytes = bcs::to_bytes(&value).unwrap();
assert!(formats.registry.contains_key("CryptoHash"));
formats.prune_known_primitives().unwrap();
assert!(!formats.registry.contains_key("CryptoHash"));
assert!(!formats.registry.contains_key("AccountOwner"));
let decoded = formats.decode_operation(&bytes).unwrap();
let expected = serde_json::to_value(&value).unwrap();
assert_eq!(decoded, expected);
assert_eq!(decoded["hash"], json!("ab".repeat(32)));
assert_eq!(decoded["amount"], json!(value.amount.to_string()));
let reencoded = formats.encode_operation(&decoded).unwrap();
assert_eq!(reencoded, bytes);
}
#[test]
fn prune_rejects_colliding_format() {
use serde_reflection::ContainerFormat;
let mut registry = Registry::new();
registry.insert(
"CryptoHash".to_string(),
ContainerFormat::NewTypeStruct(Box::new(Format::U64)),
);
let unit = Format::Unit;
let mut formats = Formats {
registry,
operation: Format::TypeName("CryptoHash".to_string()),
response: unit.clone(),
message: unit.clone(),
event_value: unit,
};
let error = formats.prune_known_primitives().unwrap_err();
assert!(matches!(error, PruneError::Mismatch { name } if name == "CryptoHash"));
assert!(formats.registry.contains_key("CryptoHash"));
}
}