From cc0f7190d1254ec89354cea8b2ed0aaecbb3b391 Mon Sep 17 00:00:00 2001 From: Michael Williamson Date: Sun, 9 Aug 2026 17:12:46 +0100 Subject: Generate Encode and Decode types --- examples/15-transient-0/output/rust/src/data.rs | 12 +--- .../output/rust/src/gen/transient_0.rs | 4 ++ .../rusttransient0/RustTransient0Generator.java | 68 ++++++++++++++++++---- .../compiler/output/lang/rust/RustTypes.java | 2 + 4 files changed, 64 insertions(+), 22 deletions(-) diff --git a/examples/15-transient-0/output/rust/src/data.rs b/examples/15-transient-0/output/rust/src/data.rs index cc0eb13..2b890c2 100644 --- a/examples/15-transient-0/output/rust/src/data.rs +++ b/examples/15-transient-0/output/rust/src/data.rs @@ -17,6 +17,7 @@ pub mod transient_0 { use ::std::any::Any; use ::std::sync::Arc; use ::std::vec::Vec; + use crate::transient_0::{Decode, Encode}; use super::{GenericNative, Native}; include!("./gen/data/transient_0.rs"); @@ -40,17 +41,6 @@ pub mod transient_0 { Ok(Native { a, b }) } - type Encode = fn( - &T, - &mut TWrite, - &mut ::std::collections::HashMap::, - ) -> std::io::Result<()>; - - type Decode = fn( - &mut TRead, - &mut Vec>, - ) -> std::io::Result; - pub fn encode_generic_native( value: &GenericNative, writer: &mut TWrite, diff --git a/examples/15-transient-0/output/rust/src/gen/transient_0.rs b/examples/15-transient-0/output/rust/src/gen/transient_0.rs index 6505408..b65ab95 100644 --- a/examples/15-transient-0/output/rust/src/gen/transient_0.rs +++ b/examples/15-transient-0/output/rust/src/gen/transient_0.rs @@ -1,5 +1,9 @@ // Generated by hobgoblin. +pub type Encode = fn(&T, &mut TWrite, &mut ::std::collections::HashMap::<::core::primitive::usize, ::core::primitive::i64>) -> ::std::io::Result::<()>; + +pub type Decode = fn(&mut TRead, &mut ::std::vec::Vec::<::std::sync::Arc::>) -> ::std::io::Result::; + pub fn encode_bool(value: &::core::primitive::bool, writer: &mut impl std::io::Write, shared_values: &mut ::std::collections::HashMap::<::core::primitive::usize, ::core::primitive::i64>) -> ::std::io::Result::<()> { if *value { writer.write_all(&1u8.to_le_bytes())?; diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java index 9d959bd..49a35f6 100644 --- a/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java +++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/generators/rusttransient0/RustTransient0Generator.java @@ -71,6 +71,8 @@ public class RustTransient0Generator implements Generator { var rustModule = new RustModule( rustModuleName, List.of( + generateEncodeType(), + generateDecodeType(), generateEncodeBoolFunction(), generateDecodeBoolFunction(), generateEncodeInt8Function(), @@ -87,6 +89,47 @@ public class RustTransient0Generator implements Generator { this.rustGenerator.write(rustModule); } + private RustItem generateEncodeType() { + var valueType = RustIdentifier.of("T"); + var writeType = RustIdentifier.of("TWrite"); + return new RustTypeAlias( + Optional.of(RustVisibility.PUB), + RustIdentifier.of("Encode"), + List.of( + new RustTypeParam(valueType, List.of()), + new RustTypeParam(writeType, List.of(RustTypes.IO_WRITE)) + ), + new RustFunctionPointerType( + List.of( + new RustSharedReferenceType(RustPath.of(valueType)), + new RustMutableReferenceType(RustPath.of(writeType)), + encoderSharedValuesType() + ), + RustTypes.ioResult(RustTypes.UNIT) + ) + ); + } + + private RustItem generateDecodeType() { + var valueType = RustIdentifier.of("T"); + var readType = RustIdentifier.of("TRead"); + return new RustTypeAlias( + Optional.of(RustVisibility.PUB), + RustIdentifier.of("Decode"), + List.of( + new RustTypeParam(readType, List.of(RustTypes.IO_READ)), + new RustTypeParam(valueType, List.of()) + ), + new RustFunctionPointerType( + List.of( + new RustMutableReferenceType(RustPath.of(readType)), + decoderSharedValuesType() + ), + RustTypes.ioResult(RustPath.of(valueType)) + ) + ); + } + private RustItem generateEncodeBoolFunction() { return generateEncodeFunction( NativeTypes.BOOL, @@ -548,12 +591,13 @@ public class RustTransient0Generator implements Generator { return List.of( new RustFunctionParam(VALUE_NAME, new RustSharedReferenceType(rustType)), new RustFunctionParam(WRITER_NAME, new RustMutableReferenceType(new RustImplTraitType(RustPath.of("std", "io", "Write")))), - new RustFunctionParam( - SHARED_VALUES_NAME, - new RustMutableReferenceType( - RustTypes.hashMap(RustTypes.USIZE, RustPath.primitive("i64")) - ) - ) + new RustFunctionParam(SHARED_VALUES_NAME, encoderSharedValuesType()) + ); + } + + private static RustMutableReferenceType encoderSharedValuesType() { + return new RustMutableReferenceType( + RustTypes.hashMap(RustTypes.USIZE, RustPath.primitive("i64")) ); } @@ -578,11 +622,13 @@ public class RustTransient0Generator implements Generator { private static List generateDecodeParams() { return List.of( new RustFunctionParam(READER_NAME, new RustMutableReferenceType(new RustImplTraitType(RustPath.of("std", "io", "Read")))), - new RustFunctionParam( - SHARED_VALUES_NAME, new RustMutableReferenceType( - RustTypes.vec(RustTypes.arc(RustTypes.dyn(RustTypes.ANY, RustTypes.SYNC, RustTypes.SEND))) - ) - ) + new RustFunctionParam(SHARED_VALUES_NAME, decoderSharedValuesType()) + ); + } + + private static RustMutableReferenceType decoderSharedValuesType() { + return new RustMutableReferenceType( + RustTypes.vec(RustTypes.arc(RustTypes.dyn(RustTypes.ANY, RustTypes.SYNC, RustTypes.SEND))) ); } diff --git a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustTypes.java b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustTypes.java index 3c5ab6c..b3d65d4 100644 --- a/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustTypes.java +++ b/src/main/java/org/zwobble/hobgoblin/compiler/output/lang/rust/RustTypes.java @@ -36,8 +36,10 @@ public class RustTypes { public static RustPath IO_ERROR = RustPath.global("std", "io", "Error"); public static RustPath IO_ERROR_KIND = RustPath.global("std", "io", "ErrorKind"); public static RustPath IO_ERROR_KIND_INVALID_DATA = IO_ERROR_KIND.addSegment(RustPathSegment.of("InvalidData")); + public static RustPath IO_READ = RustPath.global("std", "io", "Read"); public static RustPath IO_RESULT = RustPath.global("std", "io", "Result"); public static RustPath IO_RESULT_OK = IO_RESULT.addSegment(RustPathSegment.of("Ok")); + public static RustPath IO_WRITE = RustPath.global("std", "io", "Write"); public static RustPath ioResult(RustType okType) { return IO_RESULT.withArgs(List.of(okType)); -- cgit v1.2.3