From 28efafc89709282dda333e5f6e46eb0b67dbf717 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Tue, 29 Sep 2026 19:12:38 +1000 Subject: [PATCH 01/61] fix bytecode sizes and metadata remapping --- java-linker/src/remap.rs | 54 ++++++++++++++++++++++++---------------- java-linker/src/tests.rs | 9 +++++++ 2 files changed, 41 insertions(+), 22 deletions(-) diff --git a/java-linker/src/remap.rs b/java-linker/src/remap.rs index 48d2c760..f9ca2b43 100644 --- a/java-linker/src/remap.rs +++ b/java-linker/src/remap.rs @@ -116,28 +116,15 @@ pub(crate) fn remap_instruction( } pub(crate) fn instruction_byte_offsets(instructions: &[Instruction]) -> io::Result> { - let mut bytes = Cursor::new(Vec::new()); + // Branch operands are instruction indexes. Measure encoded sizes without treating those indexes as byte offsets. + let mut position = 0; let mut offsets = Vec::with_capacity(instructions.len() + 1); for instruction in instructions { - offsets.push(u16::try_from(bytes.position()).map_err(|_| { - io::Error::new( - io::ErrorKind::InvalidData, - "JVM method exceeds the bytecode offset limit", - ) - })?); - instruction.to_bytes(&mut bytes).map_err(|error| { - io::Error::new( - io::ErrorKind::InvalidData, - format!("could not measure JVM instruction: {error}"), - ) - })?; + offsets + .push(u16::try_from(position).map_err(|e| constant_pool_error("bytecode offset", e))?); + position += jvm_compiler_core::jvm::encoding::instruction_size_at(instruction, position); } - offsets.push(u16::try_from(bytes.position()).map_err(|_| { - io::Error::new( - io::ErrorKind::InvalidData, - "JVM method exceeds the bytecode offset limit", - ) - })?); + offsets.push(u16::try_from(position).map_err(|e| constant_pool_error("bytecode size", e))?); Ok(offsets) } @@ -197,6 +184,23 @@ pub(crate) fn remap_attribute( indexes: &impl ConstantIndexes, ) -> io::Result<()> { match attribute { + Attribute::InnerClasses { + name_index, + classes, + } => { + *name_index = indexes.remap(*name_index)?; + for class in classes { + for index in [ + &mut class.class_info_index, + &mut class.outer_class_info_index, + &mut class.name_index, + ] { + if *index != 0 { + *index = indexes.remap(*index)?; + } + } + } + } Attribute::Code { name_index, code, @@ -205,12 +209,18 @@ pub(crate) fn remap_attribute( .. } => { *name_index = remapped_constant_index(*name_index, indexes)?; - let old_byte_offsets = instruction_byte_offsets(code)?; + let old_byte_offsets = attributes + .iter() + .any(|a| matches!(a, Attribute::LocalVariableTable { .. })) + .then(|| instruction_byte_offsets(code)) + .transpose()?; for instruction in code.iter_mut() { remap_instruction(instruction, indexes)?; } - let new_byte_offsets = instruction_byte_offsets(code)?; - remap_local_variable_ranges(attributes, &old_byte_offsets, &new_byte_offsets)?; + if let Some(old_byte_offsets) = old_byte_offsets { + let new_byte_offsets = instruction_byte_offsets(code)?; + remap_local_variable_ranges(attributes, &old_byte_offsets, &new_byte_offsets)?; + } for exception in exception_table { if exception.catch_type != 0 { exception.catch_type = remapped_constant_index(exception.catch_type, indexes)?; diff --git a/java-linker/src/tests.rs b/java-linker/src/tests.rs index 6e3902d8..2d64cbc3 100644 --- a/java-linker/src/tests.rs +++ b/java-linker/src/tests.rs @@ -267,6 +267,15 @@ fn recognizes_msvc_output_argument_case_insensitively() { assert_eq!(msvc_output_path("/DEBUG"), None); } +#[test] +fn measures_short_branch_near_end_of_large_method_without_encoding_its_target() { + let mut code = vec![Instruction::Iinc_w(0, 1); 7_000]; + // The target is three bytes ahead. Treating its instruction index as a byte offset exceeds the branch limit. + code.extend([Instruction::Goto(7_001), Instruction::Return]); + let offsets = instruction_byte_offsets(&code).unwrap(); + assert_eq!(&offsets[7_000..], &[42_000, 42_003, 42_004]); +} + #[test] fn remaps_local_variable_ranges_when_ldc_widens() { let old_offsets = instruction_byte_offsets(&[Instruction::Ldc(1), Instruction::Return]) From d5da1fdf032d090546eab591088500548736f345 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Tue, 29 Sep 2026 20:42:29 +1000 Subject: [PATCH 02/61] share codec adapters and array codecs --- runtime/src/ArrayMemoryCodec.java | 63 ++++ runtime/src/CodecCalls.java | 74 ++++ runtime/src/MemoryCodec.java | 197 ++++++++++ runtime/src/Pointer.java | 346 +++--------------- .../pointer_provenance/CodecAdapters.java | 66 ++++ .../integration/pointer_provenance/Main.java | 1 + .../pointer_provenance/MemoryViews.java | 33 +- 7 files changed, 478 insertions(+), 302 deletions(-) create mode 100644 runtime/src/ArrayMemoryCodec.java create mode 100644 runtime/src/CodecCalls.java create mode 100644 runtime/src/MemoryCodec.java create mode 100644 tests/integration/pointer_provenance/CodecAdapters.java diff --git a/runtime/src/ArrayMemoryCodec.java b/runtime/src/ArrayMemoryCodec.java new file mode 100644 index 00000000..5d3a34c7 --- /dev/null +++ b/runtime/src/ArrayMemoryCodec.java @@ -0,0 +1,63 @@ +package org.rustlang.runtime; + +/** Shared little-endian codecs for primitive arrays, independent of Rust length. */ +final class ArrayMemoryCodec { + private ArrayMemoryCodec() { } + + static int elementSize(Class type) { + if (type == byte.class || type == boolean.class) return 1; + if (type == short.class || type == char.class) return 2; + if (type == int.class || type == float.class) return 4; + if (type == long.class || type == double.class) return 8; + throw new IllegalArgumentException("not a primitive array element: " + type); + } + + static void write(Object source, byte[] target, int offset) { + if (source instanceof byte[]) { + byte[] values = (byte[]) source; + System.arraycopy(values, 0, target, offset, values.length); + } else if (source instanceof boolean[]) { + for (boolean value : (boolean[]) source) target[offset++] = (byte) (value ? 1 : 0); + } else if (source instanceof short[]) { + for (short value : (short[]) source) { MemoryBytes.write(target, offset, 2, value); offset += 2; } + } else if (source instanceof char[]) { + for (char value : (char[]) source) { MemoryBytes.write(target, offset, 2, value); offset += 2; } + } else if (source instanceof int[]) { + for (int value : (int[]) source) { MemoryBytes.write(target, offset, 4, value); offset += 4; } + } else if (source instanceof float[]) { + for (float value : (float[]) source) { MemoryBytes.write(target, offset, 4, Float.floatToRawIntBits(value)); offset += 4; } + } else if (source instanceof long[]) { + for (long value : (long[]) source) { MemoryBytes.write(target, offset, 8, value); offset += 8; } + } else if (source instanceof double[]) { + for (double value : (double[]) source) { MemoryBytes.write(target, offset, 8, Double.doubleToRawLongBits(value)); offset += 8; } + } else throw new IllegalArgumentException("not a primitive array"); + } + + static void read(byte[] source, int offset, Object target) { + if (target instanceof byte[]) { + byte[] values = (byte[]) target; + System.arraycopy(source, offset, values, 0, values.length); + } else if (target instanceof boolean[]) { + boolean[] values = (boolean[]) target; + for (int i = 0; i < values.length; i++) values[i] = source[offset++] != 0; + } else if (target instanceof short[]) { + short[] values = (short[]) target; + for (int i = 0; i < values.length; i++, offset += 2) values[i] = (short) MemoryBytes.read(source, offset, 2); + } else if (target instanceof char[]) { + char[] values = (char[]) target; + for (int i = 0; i < values.length; i++, offset += 2) values[i] = (char) MemoryBytes.read(source, offset, 2); + } else if (target instanceof int[]) { + int[] values = (int[]) target; + for (int i = 0; i < values.length; i++, offset += 4) values[i] = (int) MemoryBytes.read(source, offset, 4); + } else if (target instanceof float[]) { + float[] values = (float[]) target; + for (int i = 0; i < values.length; i++, offset += 4) values[i] = Float.intBitsToFloat((int) MemoryBytes.read(source, offset, 4)); + } else if (target instanceof long[]) { + long[] values = (long[]) target; + for (int i = 0; i < values.length; i++, offset += 8) values[i] = MemoryBytes.read(source, offset, 8); + } else if (target instanceof double[]) { + double[] values = (double[]) target; + for (int i = 0; i < values.length; i++, offset += 8) values[i] = Double.longBitsToDouble(MemoryBytes.read(source, offset, 8)); + } else throw new IllegalArgumentException("not a primitive array"); + } +} diff --git a/runtime/src/CodecCalls.java b/runtime/src/CodecCalls.java new file mode 100644 index 00000000..cb6e1d60 --- /dev/null +++ b/runtime/src/CodecCalls.java @@ -0,0 +1,74 @@ +package org.rustlang.runtime; + +import java.lang.invoke.MethodHandle; +import java.lang.invoke.MethodType; + +/** One adapter class per operation, shared by every exact Rust layout. */ +final class CodecCalls { + private CodecCalls() {} + + interface Encoder { byte[] encode(Object value); } + interface RangeEncoder { void encode(Object value, byte[] bytes, int offset); } + interface Decoder { Object decode(byte[] bytes); } + interface RangeDecoder { Object decode(byte[] bytes, int offset); } + interface Binder { void bind(Pointer pointer, Object value); } + + static Encoder wholeEncoder(RangeEncoder operation, int byteSize) { + return value -> { + byte[] bytes = new byte[byteSize]; + operation.encode(value, bytes, 0); + return bytes; + }; + } + + static Decoder wholeDecoder(RangeDecoder operation) { + return bytes -> operation.decode(bytes, 0); + } + + static Encoder encoder(MethodHandle target) { + MethodHandle call = target.asType(MethodType.methodType(byte[].class, Object.class)); + return value -> { + try { return (byte[]) call.invokeExact(value); } + catch (Throwable error) { throw failure(error); } + }; + } + + static RangeEncoder rangeEncoder(MethodHandle target) { + MethodHandle call = target.asType(MethodType.methodType( + void.class, Object.class, byte[].class, int.class)); + return (value, bytes, offset) -> { + try { call.invokeExact(value, bytes, offset); } + catch (Throwable error) { throw failure(error); } + }; + } + + static Decoder decoder(MethodHandle target) { + MethodHandle call = target.asType(MethodType.methodType(Object.class, byte[].class)); + return bytes -> { + try { return (Object) call.invokeExact(bytes); } + catch (Throwable error) { throw failure(error); } + }; + } + + static RangeDecoder rangeDecoder(MethodHandle target) { + MethodHandle call = target.asType(MethodType.methodType(Object.class, byte[].class, int.class)); + return (bytes, offset) -> { + try { return (Object) call.invokeExact(bytes, offset); } + catch (Throwable error) { throw failure(error); } + }; + } + + static Binder binder(MethodHandle target) { + MethodHandle call = target.asType(MethodType.methodType(void.class, Pointer.class, Object.class)); + return (pointer, value) -> { + try { call.invokeExact(pointer, value); } + catch (Throwable error) { throw failure(error); } + }; + } + + private static RuntimeException failure(Throwable error) { + if (error instanceof RuntimeException) return (RuntimeException) error; + if (error instanceof Error) throw (Error) error; + return new IllegalStateException("Rust codec threw a checked exception", error); + } +} diff --git a/runtime/src/MemoryCodec.java b/runtime/src/MemoryCodec.java new file mode 100644 index 00000000..fdff0e32 --- /dev/null +++ b/runtime/src/MemoryCodec.java @@ -0,0 +1,197 @@ +package org.rustlang.runtime; + +import java.lang.invoke.MethodHandle; +import java.lang.invoke.MethodHandles; +import java.lang.invoke.MethodType; +import java.lang.reflect.Field; + +/** Resolves codec operations for an exact Rust layout only when needed. */ +final class MemoryCodec { + final Class encodeParameterType; + final int arrayElementSize; + final String arrayElementCodec; + private final Class owner; + private final String key; + private final int byteSize; + private final Object[] operations = new Object[5]; + private volatile int initialized; + + static MemoryCodec load(String recipe, ClassLoader loader) throws ReflectiveOperationException { + if (recipe.startsWith("@zero-sized:")) { + return new MemoryCodec(Class.forName(recipe.substring("@zero-sized:".length()).replace('/', '.'), false, loader)); + } + int first = recipe.indexOf('#'), second = recipe.indexOf('#', first + 1); + if (first < 1 || second <= first + 1) throw new IllegalArgumentException("invalid pointer codec recipe: " + recipe); + Class owner = Class.forName(recipe.substring(0, first).replace('/', '.'), false, loader); + int third = recipe.indexOf('#', second + 1); + String descriptor = third < 0 ? recipe.substring(second + 1) : recipe.substring(second + 1, third); + int size = third < 0 ? -1 : Integer.parseInt(recipe.substring(third + 1)); + if (third >= 0 && size < 0) throw new IllegalArgumentException("invalid codec byte size: " + recipe); + Class value = MethodType.fromMethodDescriptorString("()" + descriptor, loader).returnType(); + return new MemoryCodec(owner, recipe.substring(first + 1, second), value, size); + } + + private MemoryCodec(Class value) throws ReflectiveOperationException { + owner = null; + key = null; + byteSize = 0; + encodeParameterType = value; + arrayElementSize = -1; + arrayElementCodec = null; + MethodHandle constructor = MethodHandles.publicLookup().unreflectConstructor(value.getConstructor()) + .asType(MethodType.methodType(Object.class)); + operations[0] = (CodecCalls.Encoder) ignored -> new byte[0]; + operations[2] = (CodecCalls.Decoder) ignored -> { + try { return (Object) constructor.invokeExact(); } + catch (Throwable error) { throw failure(error); } + }; + initialized = 31; + } + + private MemoryCodec(Class owner, String key, Class value, int layoutSize) throws ReflectiveOperationException { + this.owner = owner; + this.key = key; + byteSize = layoutSize; + encodeParameterType = value; + if (owner == ArrayMemoryCodec.class && key.equals("array")) { + Class element = value.getComponentType(); + arrayElementSize = ArrayMemoryCodec.elementSize(element); + arrayElementCodec = null; + if (layoutSize < 0 || layoutSize % arrayElementSize != 0) + throw new IllegalArgumentException("invalid primitive array codec size"); + int length = layoutSize / arrayElementSize; + CodecCalls.RangeEncoder encoder = (array, bytes, offset) -> { + if (!value.isInstance(array) || java.lang.reflect.Array.getLength(array) != length) + throw new IllegalArgumentException("primitive array codec layout mismatch"); + Pointer.encodeArrayMemory(array, bytes, offset, arrayElementSize, null); + }; + CodecCalls.RangeDecoder decoder = (bytes, offset) -> { + Object array = java.lang.reflect.Array.newInstance(element, length); + Pointer.decodeArrayMemory(bytes, offset, array, arrayElementSize, null); + return array; + }; + operations[0] = CodecCalls.wholeEncoder(encoder, layoutSize); + operations[1] = encoder; + operations[2] = CodecCalls.wholeDecoder(decoder); + operations[3] = decoder; + initialized = 31; + return; + } + // Resolve array metadata only for array layouts. + if (value.isArray()) { + MethodHandle size = optional("s$", MethodType.methodType(int.class)); + MethodHandle codec = optional("c$", MethodType.methodType(String.class)); + try { + arrayElementSize = size == null ? -1 : (int) size.invokeExact(); + arrayElementCodec = codec == null ? null : (String) codec.invokeExact(); + } catch (Throwable error) { throw failure(error); } + } else { + arrayElementSize = -1; + arrayElementCodec = null; + } + } + + CodecCalls.Encoder encode() { return (CodecCalls.Encoder) operation(0); } + CodecCalls.RangeEncoder encodeAt() { return (CodecCalls.RangeEncoder) operation(1); } + CodecCalls.Decoder decode() { return (CodecCalls.Decoder) operation(2); } + CodecCalls.RangeDecoder decodeAt() { return (CodecCalls.RangeDecoder) operation(3); } + CodecCalls.Binder bind() { return (CodecCalls.Binder) operation(4); } + + private Object operation(int index) { + if ((initialized & (1 << index)) == 0) initialize(index); + return operations[index]; + } + + private synchronized void initialize(int index) { + int flag = 1 << index; + if ((initialized & flag) != 0) return; + try { + MethodHandle target; + switch (index) { + case 0: + target = optional("e$", MethodType.methodType(byte[].class, encodeParameterType)); + if (target != null) operations[index] = CodecCalls.encoder(target); + else if (byteSize >= 0 && encodeAt() != null) { + operations[index] = CodecCalls.wholeEncoder(encodeAt(), byteSize); + } else throw new NoSuchMethodException("codec has no encoder"); + break; + case 1: + target = optional("w$", MethodType.methodType(void.class, encodeParameterType, byte[].class, int.class)); + operations[index] = target == null ? null : CodecCalls.rangeEncoder(target); + break; + case 2: + target = optional("d$", MethodType.methodType(encodeParameterType, byte[].class)); + if (target != null) operations[index] = CodecCalls.decoder(target); + else if (decodeAt() != null) operations[index] = CodecCalls.wholeDecoder(decodeAt()); + else throw new NoSuchMethodException("codec has no decoder"); + break; + case 3: + target = optional("a$", MethodType.methodType(encodeParameterType, byte[].class, int.class)); + operations[index] = target == null ? null : CodecCalls.rangeDecoder(target); + break; + case 4: + target = optional("b$", MethodType.methodType(void.class, Pointer.class, encodeParameterType)); + operations[index] = target == null ? null : CodecCalls.binder(target); + break; + default: throw new AssertionError(index); + } + // Release publication also records a missing optional operation. + initialized |= flag; + } catch (ReflectiveOperationException error) { + throw new IllegalStateException("could not resolve Rust codec " + owner.getName() + "#" + key, error); + } + } + + private MethodHandle optional(String operation, MethodType type) throws IllegalAccessException { + try { return MethodHandles.lookup().findStatic(owner, operation + key, type); } + catch (NoSuchMethodException missing) { return null; } + } + + boolean isArrayCodecFor(Object value) { + return arrayElementSize > 0 && value != null && value.getClass().isArray(); + } + + // Resolve union fields only when a union byte operation needs them. + private static final ClassValue UNIONS = new ClassValue() { + protected UnionAccess computeValue(Class type) { return new UnionAccess(type); } + }; + + byte[] directUnionBytes(Object value) { + if (value == null || !encodeParameterType.isInstance(value)) return null; + MethodHandle read = UNIONS.get(encodeParameterType).bytes; + if (read == null) return null; + try { return (byte[]) read.invokeExact(value); } + catch (Throwable error) { throw failure(error); } + } + + Object[] directUnionObjects(Object value) { + if (value == null || !encodeParameterType.isInstance(value)) return null; + MethodHandle read = UNIONS.get(encodeParameterType).objects; + if (read == null) return null; + try { return (Object[]) read.invokeExact(value); } + catch (Throwable error) { throw failure(error); } + } + + private static final class UnionAccess { + final MethodHandle bytes, objects; + UnionAccess(Class type) { + MethodHandle b = null, o = null; + try { + Field bytes = type.getField("_bytes"), objects = type.getField("_objects"); + if (bytes.getType() == byte[].class && objects.getType() == Object[].class) { + b = MethodHandles.lookup().unreflectGetter(bytes).asType(MethodType.methodType(byte[].class, Object.class)); + o = MethodHandles.lookup().unreflectGetter(objects).asType(MethodType.methodType(Object[].class, Object.class)); + } + } catch (NoSuchFieldException ignored) { + } catch (IllegalAccessException error) { throw failure(error); } + this.bytes = b; + this.objects = o; + } + } + + private static RuntimeException failure(Throwable error) { + if (error instanceof RuntimeException) return (RuntimeException) error; + if (error instanceof Error) throw (Error) error; + return new IllegalStateException("Rust memory codec failed", error); + } +} diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index 6ca57649..89b01f77 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -1,7 +1,5 @@ package org.rustlang.runtime; -import java.lang.invoke.CallSite; -import java.lang.invoke.LambdaMetafactory; import java.lang.ref.ReferenceQueue; import java.lang.ref.WeakReference; import java.lang.reflect.Array; @@ -262,7 +260,7 @@ private static void rethrowUnchecked(Throwable failure) { new TreeMap<>(); private static final ReferenceQueue ALLOCATION_RANGE_QUEUE = new ReferenceQueue<>(); - private static final ConcurrentHashMap CODEC_METHODS = + private static final ConcurrentHashMap CODEC_METHODS = new ConcurrentHashMap<>(); private static final ThreadLocal RECENT_CODEC_PLANS = new ThreadLocal() { @@ -1622,176 +1620,20 @@ private Object newInstance(Object... arguments) throws ReflectiveOperationExcept } } - private static final class CodecPlan { - private final Class encodeParameterType; - private final CodecEncoder encode; - private final CodecDecoder decode; - private final CodecRangeDecoder decodeAt; - private final CodecBinder bind; - private final MethodHandle unionBytes; - private final MethodHandle unionObjects; - private final int arrayElementSize; - private final String arrayElementCodec; - - /** A fieldless Rust ZST needs only its concrete JVM constructor. */ - private CodecPlan(Class valueType) throws ReflectiveOperationException { - encodeParameterType = valueType; - MethodHandle constructor = MethodHandles.publicLookup() - .unreflectConstructor(valueType.getConstructor()) - .asType(MethodType.methodType(Object.class)); - encode = value -> new byte[0]; - decode = bytes -> { - try { - return (Object) constructor.invokeExact(); - } catch (RuntimeException | Error error) { - throw error; - } catch (Throwable error) { - throw new IllegalStateException("could not construct zero-sized Rust value", error); - } - }; - bind = null; - decodeAt = null; - unionBytes = null; - unionObjects = null; - arrayElementSize = -1; - arrayElementCodec = null; - } - - private CodecPlan( - Method encode, - Method decode, - Method decodeAt, - Method bind, - Method arrayElementSize, - Method arrayElementCodec) - throws IllegalAccessException { - encodeParameterType = encode.getParameterTypes()[0]; - MethodHandles.Lookup lookup = MethodHandles.lookup(); - try { - this.encode = (CodecEncoder) - lambdaAdapter( - lookup, - "encode", - CodecEncoder.class, - MethodType.methodType(byte[].class, Object.class), - lookup.unreflect(encode), - MethodType.methodType( - byte[].class, encode.getParameterTypes()[0])); - this.decode = (CodecDecoder) - lambdaAdapter( - lookup, - "decode", - CodecDecoder.class, - MethodType.methodType(Object.class, byte[].class), - lookup.unreflect(decode), - MethodType.methodType( - decode.getReturnType(), byte[].class)); - this.decodeAt = decodeAt == null ? null : (CodecRangeDecoder) - lambdaAdapter( - lookup, - "decode", - CodecRangeDecoder.class, - MethodType.methodType(Object.class, byte[].class, int.class), - lookup.unreflect(decodeAt), - MethodType.methodType( - decodeAt.getReturnType(), byte[].class, int.class)); - this.bind = bind == null - ? null - : (CodecBinder) - lambdaAdapter( - lookup, - "bind", - CodecBinder.class, - MethodType.methodType( - void.class, Pointer.class, Object.class), - lookup.unreflect(bind), - MethodType.methodType( - void.class, - Pointer.class, - bind.getParameterTypes()[1])); - MethodHandle directUnionBytes = null; - MethodHandle directUnionObjects = null; - try { - Field bytes = encodeParameterType.getField("_bytes"); - Field objects = encodeParameterType.getField("_objects"); - if (bytes.getType() == byte[].class - && objects.getType() == Object[].class) { - directUnionBytes = lookup.unreflectGetter(bytes).asType( - MethodType.methodType(byte[].class, Object.class)); - directUnionObjects = lookup.unreflectGetter(objects).asType( - MethodType.methodType(Object[].class, Object.class)); - } - } catch (NoSuchFieldException ignored) { - // Ordinary generated aggregates use their encode/decode methods. - } - unionBytes = directUnionBytes; - unionObjects = directUnionObjects; - if (arrayElementSize == null || arrayElementCodec == null) { - this.arrayElementSize = -1; - this.arrayElementCodec = null; - } else { - this.arrayElementSize = - (int) lookup.unreflect(arrayElementSize).invokeExact(); - this.arrayElementCodec = - (String) lookup.unreflect(arrayElementCodec).invokeExact(); - } - } catch (Throwable error) { - throw new IllegalAccessException( - "could not create pointer codec call adapter: " + error); - } - } - - private byte[] directUnionBytes(Object value) { - if (unionBytes == null || value == null - || !encodeParameterType.isInstance(value)) { - return null; - } - try { - return (byte[]) unionBytes.invokeExact(value); - } catch (RuntimeException | Error error) { - throw error; - } catch (Throwable error) { - throw new IllegalStateException( - "could not access generated Rust union storage", error); - } - } - - private Object[] directUnionObjects(Object value) { - if (unionObjects == null || value == null - || !encodeParameterType.isInstance(value)) { - return null; - } - try { - return (Object[]) unionObjects.invokeExact(value); - } catch (RuntimeException | Error error) { - throw error; - } catch (Throwable error) { - throw new IllegalStateException( - "could not access generated Rust union references", error); - } - } - - private boolean isArrayCodecFor(Object value) { - return arrayElementSize > 0 - && value != null - && value.getClass().isArray(); - } - } - /** Two-entry per-thread cache for the codecs used by tight pointer loops. */ private static final class CodecPlanCache { private String firstName; - private CodecPlan firstPlan; + private MemoryCodec firstPlan; private String secondName; - private CodecPlan secondPlan; + private MemoryCodec secondPlan; - private CodecPlan get(String name) { + private MemoryCodec get(String name) { if (name == firstName || (firstName != null && firstName.equals(name))) { return firstPlan; } if (name == secondName || (secondName != null && secondName.equals(name))) { String previousName = firstName; - CodecPlan previousPlan = firstPlan; + MemoryCodec previousPlan = firstPlan; firstName = secondName; firstPlan = secondPlan; secondName = previousName; @@ -1801,7 +1643,7 @@ private CodecPlan get(String name) { return null; } - private void remember(String name, CodecPlan plan) { + private void remember(String name, MemoryCodec plan) { secondName = firstName; secondPlan = firstPlan; firstName = name; @@ -1849,40 +1691,6 @@ private void remember(String name, ClassLoader loader, Class resolved) { } } - private interface CodecEncoder { - byte[] encode(Object value); - } - - private interface CodecDecoder { - Object decode(byte[] bytes); - } - - private interface CodecRangeDecoder { - Object decode(byte[] bytes, int offset); - } - - private interface CodecBinder { - void bind(Pointer pointer, Object value); - } - - private static Object lambdaAdapter( - MethodHandles.Lookup lookup, - String methodName, - Class interfaceType, - MethodType erasedType, - MethodHandle implementation, - MethodType instantiatedType) - throws Throwable { - CallSite site = LambdaMetafactory.metafactory( - lookup, - methodName, - MethodType.methodType(interfaceType), - erasedType, - implementation, - instantiatedType); - return site.getTarget().invoke(); - } - private static final class RepeatedArrayState { private final Object template; @@ -2504,11 +2312,13 @@ public static void encodeArrayMemory( elementCodec) .flushAllMemoryViews(); } - if (array instanceof byte[] && elementSize == 1 && elementCodec == null) { - System.arraycopy(array, 0, bytes, offset, length); + Class component = array.getClass().getComponentType(); + if (elementCodec == null && component.isPrimitive() + && elementSize == ArrayMemoryCodec.elementSize(component)) { + ArrayMemoryCodec.write(array, bytes, offset); return; } - CodecPlan aggregatePlan = isGeneratedAggregateCodec(elementCodec) + MemoryCodec aggregatePlan = isGeneratedAggregateCodec(elementCodec) ? codecPlan(elementCodec) : null; for (int index = 0; index < length; index++) { @@ -2545,11 +2355,12 @@ public static void encodeArrayMemory( public static void decodeArrayMemory( byte[] bytes, int offset, Object array, int elementSize, String elementCodec) { int length = Array.getLength(array); - if (array instanceof byte[] && elementSize == 1 && elementCodec == null) { - System.arraycopy(bytes, offset, array, 0, length); + Class componentType = array.getClass().getComponentType(); + if (elementCodec == null && componentType.isPrimitive() + && elementSize == ArrayMemoryCodec.elementSize(componentType)) { + ArrayMemoryCodec.read(bytes, offset, array); return; } - Class componentType = array.getClass().getComponentType(); for (int index = 0; index < length; index++) { Object element = decodeMemoryValue( bytes, offset + index * elementSize, elementSize, elementCodec, componentType); @@ -3937,7 +3748,7 @@ private void bindDecodedMemoryView(Object value) { || isBuiltInCodec(viewCodecClassName)) { return; } - CodecBinder bind = codecPlan(viewCodecClassName).bind; + CodecCalls.Binder bind = codecPlan(viewCodecClassName).bind(); if (bind == null) { return; } @@ -8601,7 +8412,7 @@ private static void storePrimitiveArrayByte( } private static long loadArrayCodecBits( - Object array, int byteOffset, int byteCount, CodecPlan arrayPlan) { + Object array, int byteOffset, int byteCount, MemoryCodec arrayPlan) { if (!arrayPlan.isArrayCodecFor(array)) { throw new IllegalArgumentException("pointer codec is not a fixed-array codec"); } @@ -8633,7 +8444,7 @@ private static long loadArrayCodecBits( byte[] image; boolean temporary = false; if (isGeneratedAggregateCodec(arrayPlan.arrayElementCodec)) { - CodecPlan elementPlan = codecPlan(arrayPlan.arrayElementCodec); + MemoryCodec elementPlan = codecPlan(arrayPlan.arrayElementCodec); if (elementPlan.isArrayCodecFor(element)) { long selected = loadArrayCodecBits(element, withinElement, chunk, elementPlan); @@ -8643,7 +8454,7 @@ private static long loadArrayCodecBits( } image = elementPlan.directUnionBytes(element); if (image == null) { - image = elementPlan.encode.encode(element); + image = elementPlan.encode().encode(element); temporary = true; } } else { @@ -8676,7 +8487,7 @@ private static void loadArrayCodecRange( byte[] target, int targetOffset, int byteCount, - CodecPlan arrayPlan) { + MemoryCodec arrayPlan) { if (!arrayPlan.isArrayCodecFor(array)) { throw new IllegalArgumentException("pointer codec is not a fixed-array codec"); } @@ -8711,7 +8522,7 @@ private static void loadArrayCodecRange( byte[] image; boolean temporary = false; if (isGeneratedAggregateCodec(arrayPlan.arrayElementCodec)) { - CodecPlan elementPlan = codecPlan(arrayPlan.arrayElementCodec); + MemoryCodec elementPlan = codecPlan(arrayPlan.arrayElementCodec); if (elementPlan.isArrayCodecFor(element)) { loadArrayCodecRange( element, @@ -8725,7 +8536,7 @@ private static void loadArrayCodecRange( } image = elementPlan.directUnionBytes(element); if (image == null) { - image = elementPlan.encode.encode(element); + image = elementPlan.encode().encode(element); temporary = true; } } else { @@ -8762,7 +8573,7 @@ private static void loadArrayCodecRange( } private static void storeArrayCodecBits( - Object array, int byteOffset, long incoming, int byteCount, CodecPlan arrayPlan) { + Object array, int byteOffset, long incoming, int byteCount, MemoryCodec arrayPlan) { if (!arrayPlan.isArrayCodecFor(array)) { throw new IllegalArgumentException("pointer codec is not a fixed-array codec"); } @@ -8799,7 +8610,7 @@ private static void storeArrayCodecBits( } if (isGeneratedAggregateCodec(arrayPlan.arrayElementCodec)) { - CodecPlan elementPlan = codecPlan(arrayPlan.arrayElementCodec); + MemoryCodec elementPlan = codecPlan(arrayPlan.arrayElementCodec); if (elementPlan.isArrayCodecFor(element)) { storeArrayCodecBits( element, withinElement, selected, chunk, elementPlan); @@ -8850,7 +8661,7 @@ private static void storeArrayCodecRange( byte[] source, int sourceOffset, int byteCount, - CodecPlan arrayPlan) { + MemoryCodec arrayPlan) { if (!arrayPlan.isArrayCodecFor(array)) { throw new IllegalArgumentException("pointer codec is not a fixed-array codec"); } @@ -8889,7 +8700,7 @@ private static void storeArrayCodecRange( } if (isGeneratedAggregateCodec(arrayPlan.arrayElementCodec)) { - CodecPlan elementPlan = codecPlan(arrayPlan.arrayElementCodec); + MemoryCodec elementPlan = codecPlan(arrayPlan.arrayElementCodec); if (elementPlan.isArrayCodecFor(element)) { storeArrayCodecRange( element, @@ -9096,14 +8907,14 @@ private long loadUnsignedAtSlow(long absoluteByteOffset, int byteCount) { encoded = encodeFatPointer( value, allocationElementSize, allocationCodecClassName); } else if (isGeneratedAggregateCodec(allocationCodecClassName)) { - CodecPlan plan = codecPlan(allocationCodecClassName); + MemoryCodec plan = codecPlan(allocationCodecClassName); if (plan.isArrayCodecFor(value)) { return loadArrayCodecBits( value, withinElement, byteCount, plan); } directUnionBytes = plan.directUnionBytes(value); encoded = directUnionBytes == null - ? plan.encode.encode(value) + ? plan.encode().encode(value) : directUnionBytes; } if (encoded != null) { @@ -9194,12 +9005,12 @@ private int loadByte(long absoluteByteOffset, boolean flushMemoryViews) { return result; } if (isGeneratedAggregateCodec(allocationCodecClassName)) { - CodecPlan plan = codecPlan(allocationCodecClassName); + MemoryCodec plan = codecPlan(allocationCodecClassName); if (plan.isArrayCodecFor(value)) { return (int) loadArrayCodecBits(value, withinElement, 1, plan); } byte[] direct = plan.directUnionBytes(value); - byte[] bytes = direct == null ? plan.encode.encode(value) : direct; + byte[] bytes = direct == null ? plan.encode().encode(value) : direct; if (withinElement >= bytes.length) { if (direct == null) { discardEncodedReferences(bytes); @@ -9305,7 +9116,7 @@ private byte[] loadRange(int byteCount) { image = encodeFatPointer(value, allocationElementSize, allocationCodecClassName); } else if (!directPrimitiveArray && isGeneratedAggregateCodec(allocationCodecClassName)) { - CodecPlan plan = codecPlan(allocationCodecClassName); + MemoryCodec plan = codecPlan(allocationCodecClassName); if (plan.isArrayCodecFor(value)) { loadArrayCodecRange( value, @@ -9330,7 +9141,7 @@ && isGeneratedAggregateCodec(allocationCodecClassName)) { consumed += chunk; continue; } - image = plan.encode.encode(value); + image = plan.encode().encode(value); } if (image != null) { if (withinElement + chunk > image.length) { @@ -9412,7 +9223,7 @@ private void storeBytesAt(long absoluteByteOffset, long bits, int byteCount) { encoded = encodeFatPointer( current, allocationElementSize, allocationCodecClassName); } else if (isGeneratedAggregateCodec(allocationCodecClassName)) { - CodecPlan plan = codecPlan(allocationCodecClassName); + MemoryCodec plan = codecPlan(allocationCodecClassName); if (plan.isArrayCodecFor(current)) { discardEncodedPointers( allocation, absoluteByteOffset, byteCount); @@ -9428,7 +9239,7 @@ private void storeBytesAt(long absoluteByteOffset, long bits, int byteCount) { direct[0] = (byte) bits; return; } - encoded = plan.encode.encode(current); + encoded = plan.encode().encode(current); } if (encoded != null) { if (withinElement + byteCount > encoded.length) { @@ -9543,7 +9354,7 @@ private void storeByte(long absoluteByteOffset, int value) { return; } if (isGeneratedAggregateCodec(allocationCodecClassName)) { - CodecPlan plan = codecPlan(allocationCodecClassName); + MemoryCodec plan = codecPlan(allocationCodecClassName); if (plan.isArrayCodecFor(current)) { prepareMemoryWrite(absoluteByteOffset, 1); discardEncodedPointers(allocation, absoluteByteOffset, 1); @@ -9558,7 +9369,7 @@ private void storeByte(long absoluteByteOffset, int value) { direct[0] = (byte) value; return; } - byte[] bytes = plan.encode.encode(current); + byte[] bytes = plan.encode().encode(current); if (withinElement >= bytes.length) { throw new IndexOutOfBoundsException("aggregate codec returned a short memory image"); } @@ -10312,10 +10123,10 @@ public Object directAggregate(Class type) { public Object getObjectCopyAs(String targetClassName) { if (allocation instanceof byte[] && rareState == null && viewSize > 0 && isGeneratedAggregateCodec(viewCodecClassName)) { - CodecPlan plan = codecPlan(viewCodecClassName); + MemoryCodec plan = codecPlan(viewCodecClassName); if (targetClassName != null && matchesBinaryClassName(targetClassName, plan.encodeParameterType.getName())) { - if (plan.decodeAt != null) { + if (allocation instanceof byte[] && rareState == null && plan.decodeAt() != null) { byte[] bytes = (byte[]) allocation; int offset = Math.toIntExact(byteOffset); int size = materializedViewSize(); @@ -10324,7 +10135,7 @@ && matchesBinaryClassName(targetClassName, plan.encodeParameterType.getName())) "aggregate read exceeds byte-addressable Rust storage"); } flushMemoryViewsOverlapping(byteOffset, size); - return plan.decodeAt.decode(bytes, offset); + return plan.decodeAt().decode(bytes, offset); } // Pointer-bearing and external codecs still use an independent // image carrying the source's reference/provenance metadata. @@ -10800,8 +10611,8 @@ private static boolean tryCopyDirectUnionRange( if (sourceCarrier == null || destinationCarrier == null) { return false; } - CodecPlan sourcePlan = codecPlan(source.allocationCodecClassName); - CodecPlan destinationPlan = codecPlan(destination.allocationCodecClassName); + MemoryCodec sourcePlan = codecPlan(source.allocationCodecClassName); + MemoryCodec destinationPlan = codecPlan(destination.allocationCodecClassName); byte[] sourceBytes = sourcePlan.directUnionBytes(sourceCarrier); byte[] destinationBytes = destinationPlan.directUnionBytes(destinationCarrier); if (sourceBytes == null || destinationBytes == null) { @@ -11554,7 +11365,7 @@ private void storeRange(byte[] source) { Math.toIntExact(Math.floorDiv(absoluteOffset, allocationElementSize)); int withinElement = (int) Math.floorMod(absoluteOffset, allocationElementSize); Object current = readElement(elementIndex); - CodecPlan plan = isBuiltInCodec(allocationCodecClassName) + MemoryCodec plan = isBuiltInCodec(allocationCodecClassName) ? null : codecPlan(allocationCodecClassName); int chunk = Math.min( source.length - consumed, @@ -11659,84 +11470,31 @@ && hasProjectedFieldCells(current)) { } } - private static CodecPlan codecPlan(String codecClassName) { + private static MemoryCodec codecPlan(String codecClassName) { if (codecClassName == null) { throw new IllegalStateException("pointer view has no aggregate codec"); } CodecPlanCache recent = RECENT_CODEC_PLANS.get(); - CodecPlan local = recent.get(codecClassName); + MemoryCodec local = recent.get(codecClassName); if (local != null) { return local; } - CodecPlan cached = CODEC_METHODS.get(codecClassName); + MemoryCodec cached = CODEC_METHODS.get(codecClassName); if (cached != null) { recent.remember(codecClassName, cached); return cached; } try { - if (codecClassName.startsWith(ZERO_SIZED_CODEC_PREFIX)) { - Class valueType = resolvedRuntimeClass( - codecClassName.substring(ZERO_SIZED_CODEC_PREFIX.length())); - return rememberCodecPlan(codecClassName, new CodecPlan(valueType), recent); - } - Class codec = resolvedRuntimeClass(codecClassName); - Method encode = null; - Method decode = null; - Method decodeAt = null; - Method bind = null; - Method arrayElementSize = null; - Method arrayElementCodec = null; - for (Method method : codec.getMethods()) { - if (method.getName().equals("encode") && method.getParameterTypes().length == 1) { - encode = method; - } else if (method.getName().equals("decode") - && method.getParameterTypes().length == 1) { - decode = method; - } else if (method.getName().equals("decodeAt") - && java.util.Arrays.equals(method.getParameterTypes(), - new Class[] {byte[].class, int.class})) { - decodeAt = method; - } else if (method.getName().equals("bind") - && method.getParameterTypes().length == 2) { - bind = method; - } else if (method.getName().equals("_rustArrayElementSize") - && method.getParameterTypes().length == 0) { - arrayElementSize = method; - } else if (method.getName().equals("_rustArrayElementCodec") - && method.getParameterTypes().length == 0) { - arrayElementCodec = method; - } - } - if (encode == null || decode == null) { - throw new NoSuchMethodException( - "pointer codec must define public static encode/decode methods"); - } - encode.setAccessible(true); - decode.setAccessible(true); - if (decodeAt != null) { - decodeAt.setAccessible(true); - } - if (bind != null) { - bind.setAccessible(true); - } - if (arrayElementSize != null) { - arrayElementSize.setAccessible(true); - } - if (arrayElementCodec != null) { - arrayElementCodec.setAccessible(true); - } - CodecPlan plan = new CodecPlan( - encode, decode, decodeAt, bind, arrayElementSize, arrayElementCodec); - return rememberCodecPlan(codecClassName, plan, recent); + return rememberCodecPlan(codecClassName, MemoryCodec.load(codecClassName, RUNTIME_CLASS_LOADER), recent); } catch (ReflectiveOperationException error) { throw new IllegalStateException("could not load Rust pointer codec " + codecClassName, error); } } - private static CodecPlan rememberCodecPlan( - String name, CodecPlan plan, CodecPlanCache recent) { - CodecPlan previous = CODEC_METHODS.putIfAbsent(name, plan); - CodecPlan result = previous == null ? plan : previous; + private static MemoryCodec rememberCodecPlan( + String name, MemoryCodec plan, CodecPlanCache recent) { + MemoryCodec previous = CODEC_METHODS.putIfAbsent(name, plan); + MemoryCodec result = previous == null ? plan : previous; recent.remember(name, result); return result; } @@ -11810,7 +11568,7 @@ public static Object[] emptyUnionObjectStorage() { private static byte[] encodeAggregate(String codecClassName, Object value) { try { - return codecPlan(codecClassName).encode.encode(value); + return codecPlan(codecClassName).encode().encode(value); } catch (Throwable error) { throw new IllegalStateException( "could not encode Rust aggregate memory with " @@ -11823,7 +11581,7 @@ private static byte[] encodeAggregate(String codecClassName, Object value) { private static Object decodeAggregate(String codecClassName, byte[] bytes) { try { - return codecPlan(codecClassName).decode.decode(bytes); + return codecPlan(codecClassName).decode().decode(bytes); } catch (Throwable error) { throw new IllegalStateException("could not decode Rust aggregate memory", error); } diff --git a/tests/integration/pointer_provenance/CodecAdapters.java b/tests/integration/pointer_provenance/CodecAdapters.java new file mode 100644 index 00000000..d9e9c127 --- /dev/null +++ b/tests/integration/pointer_provenance/CodecAdapters.java @@ -0,0 +1,66 @@ +import java.lang.reflect.Field; +import java.util.Map; +import org.rustlang.runtime.Pointer; + +/** Adding a layout must not generate a fresh JVM class for each operation. */ +public final class CodecAdapters { + public static final class OtherCodec { + public static byte[] e$other(MemoryViews.Pair value) { return MemoryViews.PairCodec.e$pair(value); } + public static MemoryViews.Pair d$other(byte[] bytes) { return MemoryViews.PairCodec.d$pair(bytes); } + public static MemoryViews.Pair a$other(byte[] bytes, int offset) { return MemoryViews.PairCodec.a$pair(bytes, offset); } + public static void w$other(MemoryViews.Pair value, byte[] bytes, int offset) { MemoryViews.PairCodec.w$pair(value, bytes, offset); } + } + + public static final class RangeOnly { + public static MemoryViews.Pair a$range(byte[] bytes, int offset) { return MemoryViews.PairCodec.a$pair(bytes, offset); } + public static void w$range(MemoryViews.Pair value, byte[] bytes, int offset) { MemoryViews.PairCodec.w$pair(value, bytes, offset); } + } + + public static void check() throws Exception { + String first = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;"; + String second = "CodecAdapters$OtherCodec#other#LMemoryViews$Pair;"; + for (String recipe : new String[] {first, second}) { + Pointer memory = Pointer.array(new byte[8], 0, 1).retype(8, recipe); + memory.set(new MemoryViews.Pair(37, 41)); + MemoryViews.Pair copy = (MemoryViews.Pair) memory.getObjectCopyAs(MemoryViews.Pair.class.getName()); + if (copy.first != 37 || copy.second != 41) throw new AssertionError("codec changed its value"); + } + Field cache = Pointer.class.getDeclaredField("CODEC_METHODS"); + cache.setAccessible(true); + Map plans = (Map) cache.get(null); + Object left = plans.get(first), right = plans.get(second); + Field initialized = right.getClass().getDeclaredField("initialized"); + initialized.setAccessible(true); + int used = initialized.getInt(right); + if ((used & (1 << 3)) == 0 || (used & (1 << 2)) != 0) { + throw new AssertionError("range decode resolved an unused whole-buffer decoder"); + } + for (String name : new String[] {"encode", "encodeAt", "decode", "decodeAt"}) { + java.lang.reflect.Method method = left.getClass().getDeclaredMethod(name); + method.setAccessible(true); + if (method.invoke(left).getClass() != method.invoke(right).getClass()) { + throw new AssertionError("codec operation " + name + " generated a class per layout"); + } + } + String range = "CodecAdapters$RangeOnly#range#LMemoryViews$Pair;#8"; + Pointer typed = Pointer.cell(new MemoryViews.Pair(19, 23), 8, range); + // Neither the whole encoder nor decoder has a generated wrapper. + if (typed.retype(4, null).getI32() != 19) throw new AssertionError("shared whole encoder"); + typed.retype(4, null).set(29); + MemoryViews.Pair value = (MemoryViews.Pair) typed.getObjectCopyAs(MemoryViews.Pair.class.getName()); + if (value.first != 29 || value.second != 23) throw new AssertionError("range codec round trip"); + Object plan = plans.get(range); + java.lang.reflect.Method encode = plan.getClass().getDeclaredMethod("encode"); + java.lang.reflect.Method decode = plan.getClass().getDeclaredMethod("decode"); + encode.setAccessible(true); decode.setAccessible(true); + Object encoder = encode.invoke(plan), decoder = decode.invoke(plan); + java.lang.reflect.Method call = encoder.getClass().getDeclaredMethod("encode", Object.class); + call.setAccessible(true); + byte[] bytes = (byte[]) call.invoke(encoder, new MemoryViews.Pair(31, 37)); + if (bytes.length != 8) throw new AssertionError("exact whole encoder byte size"); + call = decoder.getClass().getDeclaredMethod("decode", byte[].class); + call.setAccessible(true); + value = (MemoryViews.Pair) call.invoke(decoder, bytes); + if (value.first != 31 || value.second != 37) throw new AssertionError("shared whole decoder"); + } +} diff --git a/tests/integration/pointer_provenance/Main.java b/tests/integration/pointer_provenance/Main.java index bb7d3458..6e41af94 100644 --- a/tests/integration/pointer_provenance/Main.java +++ b/tests/integration/pointer_provenance/Main.java @@ -6,6 +6,7 @@ public class Main { public static void main(String[] args) throws Exception { ArrayViews.check(); MemoryViews.check(); + CodecAdapters.check(); StructuralViews.check(); Field field = Pointer.class.getDeclaredField("EXPOSED_ADDRESSES"); field.setAccessible(true); diff --git a/tests/integration/pointer_provenance/MemoryViews.java b/tests/integration/pointer_provenance/MemoryViews.java index 324d27f5..5e057e8a 100644 --- a/tests/integration/pointer_provenance/MemoryViews.java +++ b/tests/integration/pointer_provenance/MemoryViews.java @@ -6,13 +6,21 @@ public final class MemoryViews { public static final class Pair { public int first; public int second; + + public Pair() { } + + public Pair(int first, int second) { + this.first = first; + this.second = second; + } } public static final class PairCodec { static int encodes; + static int rangeEncodes; static byte[] decodedStorage; - public static byte[] encode(Pair value) { + public static byte[] e$pair(Pair value) { encodes++; byte[] bytes = new byte[8]; MemoryBytes.write(bytes, 0, 4, value.first); @@ -20,25 +28,34 @@ public static byte[] encode(Pair value) { return bytes; } - public static Pair decode(byte[] bytes) { - return decodeAt(bytes, 0); + public static Pair d$pair(byte[] bytes) { + return a$pair(bytes, 0); } - public static Pair decodeAt(byte[] bytes, int offset) { + public static Pair a$pair(byte[] bytes, int offset) { decodedStorage = bytes; Pair value = new Pair(); value.first = (int) MemoryBytes.read(bytes, offset, 4); value.second = (int) MemoryBytes.read(bytes, offset + 4, 4); return value; } + + public static void w$pair(Pair value, byte[] bytes, int offset) { + rangeEncodes++; + MemoryBytes.write(bytes, offset, 4, value.first); + MemoryBytes.write(bytes, offset + 4, 4, value.second); + } } + private static final String PAIR_CODEC = + "MemoryViews$PairCodec#pair#LMemoryViews$Pair;"; + public static void main(String[] args) { check(); } private static Pair view(Pointer bytes) { - return (Pair) bytes.retype(8, PairCodec.class.getName()).getObject(); + return (Pair) bytes.retype(8, PAIR_CODEC).getObject(); } public static void check() { @@ -78,7 +95,7 @@ public static void check() { } left.first = 71; - Pointer copiedPointer = bytes.retype(8, PairCodec.class.getName()); + Pointer copiedPointer = bytes.retype(8, PAIR_CODEC); Pair snapshot = (Pair) copiedPointer.getObjectCopyAs(Pair.class.getName()); if (PairCodec.decodedStorage != storage) { throw new AssertionError("plain aggregate read copied its source buffer"); @@ -90,7 +107,7 @@ public static void check() { if (view(bytes).first != 71) { throw new AssertionError("value copy retained a live alias"); } - Pair offsetCopy = (Pair) bytes.byte_offset(8).retype(8, PairCodec.class.getName()) + Pair offsetCopy = (Pair) bytes.byte_offset(8).retype(8, PAIR_CODEC) .getObjectCopyAs(Pair.class.getName()); if (offsetCopy.first != 101 || offsetCopy.second != 59 || PairCodec.decodedStorage != storage) { @@ -119,7 +136,7 @@ public static void check() { Pair fields = new Pair(); fields.first = 17; fields.second = 23; - Pointer root = Pointer.cell(fields, 8, PairCodec.class.getName()); + Pointer root = Pointer.cell(fields, 8, PAIR_CODEC); Pointer first = root.projectStructField(Pair.class.getName(), "first", 0, 4, null); if (first.add(0).offset_from(first) != 0 || first.add(1).getI32() != 23) { throw new AssertionError("pointer arithmetic lost the containing allocation"); From 154b7693ad994d01c0ee6d28a7c78d488851273d Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Tue, 29 Sep 2026 22:10:39 +1000 Subject: [PATCH 03/61] bound weak metadata caches --- runtime/src/Pointer.java | 448 +++---------------------------- runtime/src/WeakIdentityMap.java | 287 ++++++++++++++++++++ 2 files changed, 320 insertions(+), 415 deletions(-) create mode 100644 runtime/src/WeakIdentityMap.java diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index 89b01f77..10541a22 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -13,7 +13,6 @@ import java.lang.invoke.MethodType; import java.math.BigInteger; import java.nio.charset.StandardCharsets; -import java.util.AbstractMap; import java.util.Arrays; import java.util.HashMap; import java.util.HashSet; @@ -474,7 +473,7 @@ private RebuildableIdentityFilter(int wordCount, long rebuildMarks) { private static final RebuildableIdentityFilter MEMORY_VIEW_FILTER = new RebuildableIdentityFilter(); private static final RebuildableIdentityFilter MEMORY_VIEW_ORIGIN_FILTER = - new RebuildableIdentityFilter(); + new RebuildableIdentityFilter(REPEATED_ARRAY_FILTER_WORDS, MEMORY_VIEW_ORIGIN_FILTER_REBUILD_MARKS); private static final AtomicLongArray MEMORY_VIEW_EPOCHS = new AtomicLongArray(STATE_STRIPE_COUNT); private static final RebuildableIdentityFilter ENCODED_REFERENCE_FILTER = @@ -580,7 +579,6 @@ private static void retainEncodedReference(Object owner, Object referencedAlloca || owner == referencedAllocation) { return; } - markIdentityFilter(ENCODED_REFERENCE_FILTER, owner); Map stripe = stateStripe(ENCODED_REFERENCES, owner); synchronized (stripe) { @@ -597,8 +595,9 @@ private static void retainEncodedReference(Object owner, Object referencedAlloca } references.allocations.put(referencedAllocation, Boolean.TRUE); } + markIdentityFilter(ENCODED_REFERENCE_FILTER, owner); } - maybeRebuildEncodedReferenceFilter(); + maybeRebuildIdentityFilter(ENCODED_REFERENCE_FILTER, ENCODED_REFERENCES); } private static void transferEncodedReferences(Object sourceOwner, Object targetOwner) { @@ -927,7 +926,6 @@ private static void rememberEncodedPointer( if (owner == null || target == null || size <= 0 || codec == null) { return; } - markIdentityFilter(ENCODED_POINTER_FILTER, owner); Map> stripe = stateStripe(ENCODED_POINTERS, owner); synchronized (stripe) { @@ -940,8 +938,9 @@ private static void rememberEncodedPointer( pointers.put( offset, new EncodedPointerState(owner, size, codec, target)); + markIdentityFilter(ENCODED_POINTER_FILTER, owner); } - maybeRebuildEncodedPointerFilter(); + maybeRebuildIdentityFilter(ENCODED_POINTER_FILTER, ENCODED_POINTERS); } private static Pointer encodedPointer( @@ -1016,7 +1015,6 @@ private static void transferEncodedPointers( if (copied == null) { return; } - markIdentityFilter(ENCODED_POINTER_FILTER, targetOwner); Map> targetStripe = stateStripe(ENCODED_POINTERS, targetOwner); synchronized (targetStripe) { @@ -1038,6 +1036,7 @@ private static void transferEncodedPointers( entry.codec, entry.target)); } + markIdentityFilter(ENCODED_POINTER_FILTER, targetOwner); } } @@ -1091,293 +1090,6 @@ private static void discardEncodedPointers(Object owner, long offset, int size) } } - private static final class WeakIdentityMap extends AbstractMap { - private static final class OpenWeakReference extends WeakReference { - private final int identityHash; - - private OpenWeakReference(Object key, ReferenceQueue queue) { - super(key, queue); - identityHash = System.identityHashCode(key); - } - } - - private final ReferenceQueue collectedKeys = new ReferenceQueue<>(); - private OpenWeakReference[] keys; - private Object[] values; - private int size; - private int used; - private int readsUntilCleanup = 256; - - private WeakIdentityMap() { - this(16); - } - - private WeakIdentityMap(int initialCapacity) { - int capacity = 16; - while (capacity < initialCapacity) { - capacity <<= 1; - } - keys = newTable(capacity); - values = new Object[capacity]; - } - - private static OpenWeakReference[] newTable(int length) { - return new OpenWeakReference[length]; - } - - private static int tableIndex(int hash, int length) { - hash ^= hash >>> 16; - hash *= 0x7feb352d; - hash ^= hash >>> 15; - return hash & (length - 1); - } - - /** Removes an entry without leaving a tombstone in its probe chain. */ - private void deleteEntry(int deleted) { - if (values[deleted] != null) { - size--; - } - keys[deleted] = null; - values[deleted] = null; - used--; - - int mask = keys.length - 1; - for (int index = (deleted + 1) & mask; - keys[index] != null; - index = (index + 1) & mask) { - OpenWeakReference reference = keys[index]; - int home = tableIndex(reference.identityHash, keys.length); - if ((index < home && (home <= deleted || deleted <= index)) - || (home <= deleted && deleted <= index)) { - keys[deleted] = reference; - values[deleted] = values[index]; - keys[index] = null; - values[index] = null; - deleted = index; - } - } - } - - private void reset() { - keys = newTable(16); - values = new Object[16]; - size = 0; - used = 0; - readsUntilCleanup = 256; - while (collectedKeys.poll() != null) { - // Entries no longer exist after reset. - } - } - - private void discardCollectedKeys() { - OpenWeakReference collected; - while ((collected = (OpenWeakReference) collectedKeys.poll()) != null) { - int index = tableIndex(collected.identityHash, keys.length); - while (keys[index] != null) { - if (keys[index] == collected) { - deleteEntry(index); - break; - } - index = (index + 1) & (keys.length - 1); - } - } - } - - private void maybeDiscardCollectedKeys() { - if (--readsUntilCleanup == 0) { - discardCollectedKeys(); - readsUntilCleanup = 256; - } - } - - private void rehashForInsert() { - discardCollectedKeys(); - readsUntilCleanup = 256; - int newLength = - size * 4 >= keys.length * 3 - ? keys.length << 1 - : keys.length; - OpenWeakReference[] oldKeys = keys; - Object[] oldValues = values; - keys = newTable(newLength); - values = new Object[newLength]; - size = 0; - used = 0; - for (int oldIndex = 0; oldIndex < oldKeys.length; oldIndex++) { - OpenWeakReference reference = oldKeys[oldIndex]; - Object key = reference == null ? null : reference.get(); - if (key == null) { - continue; - } - int index = tableIndex(reference.identityHash, keys.length); - while (keys[index] != null) { - index = (index + 1) & (keys.length - 1); - } - keys[index] = reference; - values[index] = oldValues[oldIndex]; - size++; - used++; - } - } - - @SuppressWarnings("unchecked") - private V valueAt(int index) { - return (V) values[index]; - } - - @Override - public V get(Object key) { - maybeDiscardCollectedKeys(); - int index = tableIndex(System.identityHashCode(key), keys.length); - while (true) { - OpenWeakReference reference = keys[index]; - if (reference == null) { - return null; - } - Object live = reference.get(); - if (live == null) { - // The key can be collected after the batched queue drain; - // leave a tombstone until the next maintenance pass. - } else if (live == key) { - return valueAt(index); - } - index = (index + 1) & (keys.length - 1); - } - } - - @Override - public V put(Object key, V value) { - discardCollectedKeys(); - readsUntilCleanup = 256; - if (used * 4 >= keys.length * 3) { - rehashForInsert(); - } - int index = tableIndex(System.identityHashCode(key), keys.length); - while (true) { - OpenWeakReference reference = keys[index]; - if (reference == null) { - used++; - keys[index] = new OpenWeakReference(key, collectedKeys); - values[index] = value; - size++; - return null; - } - Object live = reference.get(); - if (live == null) { - deleteEntry(index); - continue; - } else if (live == key) { - V previous = valueAt(index); - values[index] = value; - return previous; - } - index = (index + 1) & (keys.length - 1); - } - } - - @Override - public V remove(Object key) { - discardCollectedKeys(); - readsUntilCleanup = 256; - int index = tableIndex(System.identityHashCode(key), keys.length); - while (true) { - OpenWeakReference reference = keys[index]; - if (reference == null) { - return null; - } - Object live = reference.get(); - if (live == null) { - deleteEntry(index); - continue; - } else if (live == key) { - V previous = valueAt(index); - deleteEntry(index); - if (size == 0) { - reset(); - } - return previous; - } - index = (index + 1) & (keys.length - 1); - } - } - - @Override - public int size() { - discardCollectedKeys(); - for (int index = 0; index < keys.length; ) { - OpenWeakReference reference = keys[index]; - if (reference != null && reference.get() == null) { - deleteEntry(index); - } else { - index++; - } - } - return size; - } - - @Override - public boolean isEmpty() { - discardCollectedKeys(); - if (size == 0) { - return true; - } - for (int index = 0; index < keys.length; ) { - OpenWeakReference reference = keys[index]; - if (reference == null) { - index++; - continue; - } - if (reference.get() != null) { - return false; - } - deleteEntry(index); - } - reset(); - return true; - } - - @Override - public Set> entrySet() { - discardCollectedKeys(); - Set> entries = new HashSet<>(); - for (int index = 0; index < keys.length; ) { - OpenWeakReference reference = keys[index]; - Object key = reference == null ? null : reference.get(); - if (key == null) { - if (reference != null) { - deleteEntry(index); - continue; - } - } else { - entries.add(new java.util.AbstractMap.SimpleImmutableEntry<>( - key, valueAt(index))); - } - index++; - } - return entries; - } - - private long markLiveKeys(AtomicLongArray filter) { - discardCollectedKeys(); - long count = 0; - for (int index = 0; index < keys.length; ) { - OpenWeakReference reference = keys[index]; - Object key = reference == null ? null : reference.get(); - if (key == null) { - if (reference != null) { - deleteEntry(index); - continue; - } - } else { - markIdentityFilter(filter, key); - count++; - } - index++; - } - return count; - } - } - /** Must be called while holding {@link #ALLOCATIONS}. */ private static AllocationInfo allocationInfo(Object allocation) { AllocationInfo info = ALLOCATIONS.get(allocation); @@ -3171,7 +2883,7 @@ private StructuralViewState structuralViewState(Object source, boolean create) { } } if (created) { - maybeRebuildStructuralViewFilter(); + maybeRebuildIdentityFilter(STRUCTURAL_VIEW_FILTER, STRUCTURAL_VIEWS); } return state; } @@ -3221,12 +2933,17 @@ private static boolean mayBeInIdentityFilter(AtomicLongArray filter, Object valu } private static void markIdentityFilter(RebuildableIdentityFilter filter, Object value) { - AtomicLongArray primary = filter.primary; - boolean added = markIdentityFilter(primary, value); - AtomicLongArray secondary = filter.secondary; - if (secondary != null && secondary != primary) { - markIdentityFilter(secondary, value); - } + // The caller holds the owning stripe during publication. + // Also mark the replacement filter if a rebuild has passed that stripe. + // Recheck primary because the rebuild can publish it and clear secondary between reads. + boolean added = false; + AtomicLongArray primary; + do { + primary = filter.primary; + added |= markIdentityFilter(primary, value); + AtomicLongArray secondary = filter.secondary; + if (secondary != null && secondary != primary) markIdentityFilter(secondary, value); + } while (primary != filter.primary); if (added) { filter.marks.incrementAndGet(); } @@ -3244,119 +2961,23 @@ private static boolean mayBeInIdentityFilter( && mayBeInIdentityFilter(secondary, value); } - private static void maybeRebuildMemoryViewFilter() { - if (MEMORY_VIEW_FILTER.marks.get() < MEMORY_VIEW_FILTER.rebuildMarks - || !MEMORY_VIEW_FILTER.rebuilding.compareAndSet(0, 1)) { - return; - } - try { - AtomicLongArray rebuilt = - new AtomicLongArray(MEMORY_VIEW_FILTER.wordCount); - MEMORY_VIEW_FILTER.secondary = rebuilt; - for (Map> stripe : MEMORY_VIEWS) { - synchronized (stripe) { - ((WeakIdentityMap>) stripe) - .markLiveKeys(rebuilt); - } - } - MEMORY_VIEW_FILTER.primary = rebuilt; - MEMORY_VIEW_FILTER.secondary = null; - MEMORY_VIEW_FILTER.marks.set(0); - } finally { - MEMORY_VIEW_FILTER.rebuilding.set(0); - } - } - - private static void maybeRebuildMemoryViewOriginFilter() { - if (MEMORY_VIEW_ORIGIN_FILTER.marks.get() - < MEMORY_VIEW_ORIGIN_FILTER_REBUILD_MARKS - || !MEMORY_VIEW_ORIGIN_FILTER.rebuilding.compareAndSet(0, 1)) { - return; - } - try { - AtomicLongArray rebuilt = - new AtomicLongArray(MEMORY_VIEW_ORIGIN_FILTER.wordCount); - MEMORY_VIEW_ORIGIN_FILTER.secondary = rebuilt; - for (Map stripe : MEMORY_VIEW_ORIGINS) { - synchronized (stripe) { - ((WeakIdentityMap) stripe) - .markLiveKeys(rebuilt); - } - } - MEMORY_VIEW_ORIGIN_FILTER.primary = rebuilt; - MEMORY_VIEW_ORIGIN_FILTER.secondary = null; - MEMORY_VIEW_ORIGIN_FILTER.marks.set(0); - } finally { - MEMORY_VIEW_ORIGIN_FILTER.rebuilding.set(0); - } - } - - private static void maybeRebuildStructuralViewFilter() { - if (STRUCTURAL_VIEW_FILTER.marks.get() < STRUCTURAL_VIEW_FILTER.rebuildMarks - || !STRUCTURAL_VIEW_FILTER.rebuilding.compareAndSet(0, 1)) { - return; - } - try { - AtomicLongArray rebuilt = - new AtomicLongArray(STRUCTURAL_VIEW_FILTER.wordCount); - STRUCTURAL_VIEW_FILTER.secondary = rebuilt; - for (Map> stripe : STRUCTURAL_VIEWS) { - synchronized (stripe) { - ((WeakIdentityMap>) stripe) - .markLiveKeys(rebuilt); - } - } - STRUCTURAL_VIEW_FILTER.primary = rebuilt; - STRUCTURAL_VIEW_FILTER.secondary = null; - STRUCTURAL_VIEW_FILTER.marks.set(0); - } finally { - STRUCTURAL_VIEW_FILTER.rebuilding.set(0); - } - } - - private static void maybeRebuildEncodedReferenceFilter() { - if (ENCODED_REFERENCE_FILTER.marks.get() < ENCODED_REFERENCE_FILTER.rebuildMarks - || !ENCODED_REFERENCE_FILTER.rebuilding.compareAndSet(0, 1)) { - return; - } - try { - AtomicLongArray rebuilt = - new AtomicLongArray(ENCODED_REFERENCE_FILTER.wordCount); - ENCODED_REFERENCE_FILTER.secondary = rebuilt; - for (Map stripe : ENCODED_REFERENCES) { - synchronized (stripe) { - ((WeakIdentityMap) stripe).markLiveKeys(rebuilt); - } - } - ENCODED_REFERENCE_FILTER.primary = rebuilt; - ENCODED_REFERENCE_FILTER.secondary = null; - ENCODED_REFERENCE_FILTER.marks.set(0); - } finally { - ENCODED_REFERENCE_FILTER.rebuilding.set(0); - } - } - - private static void maybeRebuildEncodedPointerFilter() { - if (ENCODED_POINTER_FILTER.marks.get() - < ENCODED_POINTER_FILTER.rebuildMarks - || !ENCODED_POINTER_FILTER.rebuilding.compareAndSet(0, 1)) { - return; - } + private static void maybeRebuildIdentityFilter( + RebuildableIdentityFilter filter, Map[] stripes) { + if (filter.marks.get() < filter.rebuildMarks + || !filter.rebuilding.compareAndSet(0, 1)) return; try { - AtomicLongArray rebuilt = - new AtomicLongArray(ENCODED_POINTER_FILTER.wordCount); - ENCODED_POINTER_FILTER.secondary = rebuilt; - for (Map> stripe : ENCODED_POINTERS) { + AtomicLongArray rebuilt = new AtomicLongArray(filter.wordCount); + filter.secondary = rebuilt; + for (Map stripe : stripes) { synchronized (stripe) { - ((WeakIdentityMap>) stripe) - .markLiveKeys(rebuilt); + ((WeakIdentityMap) stripe).forEachLiveKey(key -> markIdentityFilter(rebuilt, key)); } } - ENCODED_POINTER_FILTER.primary = rebuilt; - ENCODED_POINTER_FILTER.secondary = null; - ENCODED_POINTER_FILTER.marks.set(0); + filter.primary = rebuilt; + filter.secondary = null; + filter.marks.set(0); } finally { - ENCODED_POINTER_FILTER.rebuilding.set(0); + filter.rebuilding.set(0); } } @@ -3626,7 +3247,6 @@ private Object decodedMemoryViewLocked() { new MemoryViewState(materializedSize, viewCodecClassName, decoded, image); advanceMemoryViewEpoch(allocation); setDirectCellHasMemoryView(allocation, true); - markIdentityFilter(MEMORY_VIEW_FILTER, allocation); synchronized (stripe) { LongRangeMap views = stripe.get(allocation); if (views == null) { @@ -3634,12 +3254,10 @@ private Object decodedMemoryViewLocked() { stripe.put(allocation, views); } views.put(byteOffset, state); + markIdentityFilter(MEMORY_VIEW_FILTER, allocation); } advanceMemoryViewEpoch(allocation); - // Mark again after publishing the map entry so a concurrent filter - // rebuild cannot miss an insertion whose stripe it already scanned. - markIdentityFilter(MEMORY_VIEW_FILTER, allocation); - maybeRebuildMemoryViewFilter(); + maybeRebuildIdentityFilter(MEMORY_VIEW_FILTER, MEMORY_VIEWS); setBoundMemoryViewState(state); registerMemoryViewOrigin(decoded); bindDecodedMemoryView(decoded); @@ -3680,7 +3298,7 @@ private void registerMemoryViewOrigin(Object value) { markIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, value); previous = stripe.put(value, new MemoryViewOrigin(this)); } - maybeRebuildMemoryViewOriginFilter(); + maybeRebuildIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, MEMORY_VIEW_ORIGINS); Object previousAllocation = previous == null ? null : previous.allocation.get(); if (previousAllocation instanceof FieldCell && previousAllocation != allocation) { removeMemoryOriginView(previousAllocation, value); diff --git a/runtime/src/WeakIdentityMap.java b/runtime/src/WeakIdentityMap.java new file mode 100644 index 00000000..31cad297 --- /dev/null +++ b/runtime/src/WeakIdentityMap.java @@ -0,0 +1,287 @@ +package org.rustlang.runtime; + +import java.lang.ref.ReferenceQueue; +import java.lang.ref.WeakReference; +import java.util.AbstractMap; +import java.util.HashSet; +import java.util.Map; +import java.util.Set; + +/** Identity-keyed weak table. Callers synchronize every operation externally. */ +final class WeakIdentityMap extends AbstractMap { + private static final class OpenWeakReference extends WeakReference { + private final int identityHash; + + private OpenWeakReference(Object key, ReferenceQueue queue) { + super(key, queue); + identityHash = System.identityHashCode(key); + } + } + + private final ReferenceQueue collectedKeys = new ReferenceQueue<>(); + private OpenWeakReference[] keys; + private Object[] values; + private int size; + private int used; + private int readsUntilCleanup = 256; + + WeakIdentityMap() { + this(16); + } + + WeakIdentityMap(int initialCapacity) { + int capacity = 16; + while (capacity < initialCapacity) { + capacity <<= 1; + } + keys = new OpenWeakReference[capacity]; + values = new Object[capacity]; + } + + private static int tableIndex(int hash, int length) { + hash ^= hash >>> 16; + hash *= 0x7feb352d; + hash ^= hash >>> 15; + return hash & (length - 1); + } + + /** Removes an entry without leaving a tombstone in its probe chain. */ + private void deleteEntry(int deleted) { + if (values[deleted] != null) { + size--; + } + keys[deleted] = null; + values[deleted] = null; + used--; + + int mask = keys.length - 1; + for (int index = (deleted + 1) & mask; + keys[index] != null; + index = (index + 1) & mask) { + OpenWeakReference reference = keys[index]; + int home = tableIndex(reference.identityHash, keys.length); + if ((index < home && (home <= deleted || deleted <= index)) + || (home <= deleted && deleted <= index)) { + keys[deleted] = reference; + values[deleted] = values[index]; + keys[index] = null; + values[index] = null; + deleted = index; + } + } + } + + private void reset() { + keys = new OpenWeakReference[16]; + values = new Object[16]; + size = 0; + used = 0; + readsUntilCleanup = 256; + while (collectedKeys.poll() != null) { + // Entries no longer exist after reset. + } + } + + private void discardCollectedKeys() { + OpenWeakReference collected; + while ((collected = (OpenWeakReference) collectedKeys.poll()) != null) { + int index = tableIndex(collected.identityHash, keys.length); + while (keys[index] != null) { + if (keys[index] == collected) { + deleteEntry(index); + break; + } + index = (index + 1) & (keys.length - 1); + } + } + } + + private void maybeDiscardCollectedKeys() { + if (--readsUntilCleanup == 0) { + discardCollectedKeys(); + readsUntilCleanup = 256; + } + } + + private void rehashForInsert() { + discardCollectedKeys(); + readsUntilCleanup = 256; + int newLength = + size * 4 >= keys.length * 3 + ? keys.length << 1 + : keys.length; + OpenWeakReference[] oldKeys = keys; + Object[] oldValues = values; + keys = new OpenWeakReference[newLength]; + values = new Object[newLength]; + size = 0; + used = 0; + for (int oldIndex = 0; oldIndex < oldKeys.length; oldIndex++) { + OpenWeakReference reference = oldKeys[oldIndex]; + Object key = reference == null ? null : reference.get(); + if (key == null) { + continue; + } + int index = tableIndex(reference.identityHash, keys.length); + while (keys[index] != null) { + index = (index + 1) & (keys.length - 1); + } + keys[index] = reference; + values[index] = oldValues[oldIndex]; + size++; + used++; + } + } + + @SuppressWarnings("unchecked") + private V valueAt(int index) { + return (V) values[index]; + } + + @Override + public V get(Object key) { + maybeDiscardCollectedKeys(); + int index = tableIndex(System.identityHashCode(key), keys.length); + while (true) { + OpenWeakReference reference = keys[index]; + if (reference == null) { + return null; + } + Object live = reference.get(); + // Keep dead entries in the probe chain until the next queue drain. + if (live != null && live == key) { + return valueAt(index); + } + index = (index + 1) & (keys.length - 1); + } + } + + @Override + public V put(Object key, V value) { + discardCollectedKeys(); + readsUntilCleanup = 256; + if (used * 4 >= keys.length * 3) { + rehashForInsert(); + } + int index = tableIndex(System.identityHashCode(key), keys.length); + while (true) { + OpenWeakReference reference = keys[index]; + if (reference == null) { + used++; + keys[index] = new OpenWeakReference(key, collectedKeys); + values[index] = value; + size++; + return null; + } + Object live = reference.get(); + if (live == null) { + deleteEntry(index); + continue; + } else if (live == key) { + V previous = valueAt(index); + values[index] = value; + return previous; + } + index = (index + 1) & (keys.length - 1); + } + } + + @Override + public V remove(Object key) { + discardCollectedKeys(); + readsUntilCleanup = 256; + int index = tableIndex(System.identityHashCode(key), keys.length); + while (true) { + OpenWeakReference reference = keys[index]; + if (reference == null) { + return null; + } + Object live = reference.get(); + if (live == null) { + deleteEntry(index); + continue; + } else if (live == key) { + V previous = valueAt(index); + deleteEntry(index); + if (size == 0) { + reset(); + } + return previous; + } + index = (index + 1) & (keys.length - 1); + } + } + + @Override + public int size() { + discardCollectedKeys(); + for (int index = 0; index < keys.length; ) { + OpenWeakReference reference = keys[index]; + if (reference != null && reference.get() == null) { + deleteEntry(index); + } else { + index++; + } + } + return size; + } + + @Override + public boolean isEmpty() { + discardCollectedKeys(); + if (size == 0) { + return true; + } + for (int index = 0; index < keys.length; ) { + OpenWeakReference reference = keys[index]; + if (reference == null) { + index++; + continue; + } + if (reference.get() != null) { + return false; + } + deleteEntry(index); + } + reset(); + return true; + } + + @Override + public Set> entrySet() { + discardCollectedKeys(); + Set> entries = new HashSet<>(); + for (int index = 0; index < keys.length; ) { + OpenWeakReference reference = keys[index]; + Object key = reference == null ? null : reference.get(); + if (key == null) { + if (reference != null) { + deleteEntry(index); + continue; + } + } else { + entries.add(new java.util.AbstractMap.SimpleImmutableEntry<>( + key, valueAt(index))); + } + index++; + } + return entries; + } + + void forEachLiveKey(java.util.function.Consumer visit) { + discardCollectedKeys(); + for (int index = 0; index < keys.length; ) { + OpenWeakReference reference = keys[index]; + Object key = reference == null ? null : reference.get(); + if (key == null) { + if (reference != null) { + deleteEntry(index); + continue; + } + } else { + visit.accept(key); + } + index++; + } + } +} From 146ef8e3b3c9ad45a64476840a374dada1bcc25b Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 00:46:28 +1000 Subject: [PATCH 04/61] cache field projections --- runtime/src/Pointer.java | 332 +++++++++++------- .../pointer_provenance/CyclicFieldViews.java | 51 +++ .../pointer_provenance/FieldProjections.java | 130 +++++++ .../integration/pointer_provenance/Main.java | 2 + 4 files changed, 395 insertions(+), 120 deletions(-) create mode 100644 tests/integration/pointer_provenance/CyclicFieldViews.java create mode 100644 tests/integration/pointer_provenance/FieldProjections.java diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index 10541a22..b47c91b0 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -2360,6 +2360,7 @@ private static Object decodeFatPointer( if (traitInterface == null) { data = data.retype(elementSize, elementCodec); } else { + data = independentFieldMetadata(data); Pointer marker = pointerObjectFromAddress(pointerMetadata); TraitMetadataInfo info = TRAIT_METADATA_INFO.get(marker.numericAddress()); data.traitMetadataMarker(marker); @@ -3485,7 +3486,7 @@ public void commitMemoryView() { // decoding this field separately. Nested direct writes still need // to reach the enclosing enum's original byte storage. if (allocation instanceof FieldCell) { - commitOriginMemoryView(((FieldCell) allocation).owner()); + ((FieldCell) allocation).commitOwners(); } return; } @@ -3556,21 +3557,54 @@ private Object managedViewObject() { // identity filter as short-lived field pointers accumulate in long runs. private static boolean mayHaveFieldCells(Object owner) { if (owner instanceof Cell) { - return ((Cell) owner).hasFieldCells; + return ((Cell) owner).fields != null; } if (owner instanceof FieldCell) { - return ((FieldCell) owner).hasFieldCells; + return ((FieldCell) owner).fields != null; } return owner != null && mayBeInIdentityFilter(FIELD_CELL_FILTER, owner); } - private static void markFieldCells(Object owner) { - if (owner instanceof Cell) { - ((Cell) owner).hasFieldCells = true; - } else if (owner instanceof FieldCell) { - ((FieldCell) owner).hasFieldCells = true; - } else { - markIdentityFilter(FIELD_CELL_FILTER, owner); + private static final class FieldCellCache extends WeakReference { + private volatile FieldCellCache next; + + private FieldCellCache(FieldCell field, FieldCellCache next) { + super(field); + this.next = next; + } + } + + private static FieldCellCache ownedFields(Object owner) { + return owner instanceof Cell ? ((Cell) owner).fields : ((FieldCell) owner).fields; + } + + /** Keep cached field cells weak. A decoded value can refer to its owner. + * A strong field cache would retain that cycle through the global view map. */ + private static FieldCell ownedField(Object owner, FieldAccess access) { + for (FieldCellCache entry = ownedFields(owner); entry != null; entry = entry.next) { + FieldCell field = entry.get(); + if (field != null && (field.access == access || field.access.cacheKey.equals(access.cacheKey))) + return field; + } + synchronized (owner) { + FieldCellCache head = ownedFields(owner), previous = null; + for (FieldCellCache entry = head; entry != null; entry = entry.next) { + FieldCell field = entry.get(); + if (field != null) { + if (field.access == access || field.access.cacheKey.equals(access.cacheKey)) return field; + previous = entry; + } else if (previous == null) { + head = entry.next; + } else { + // Keep the next link for readers that still hold this dead entry. + previous.next = entry.next; + } + } + FieldCell field = new FieldCell(owner, access, true); + FieldCellCache entry = new FieldCellCache(field, head); + if (owner instanceof Cell) ((Cell) owner).fields = entry; + else ((FieldCell) owner).fields = entry; + return field; } } @@ -3605,6 +3639,12 @@ private static java.util.List collectProjectedFieldViews( if (!mayHaveProjectedViews(owner)) { return projected; } + if (owner instanceof Cell || owner instanceof FieldCell) { + for (FieldCellCache entry = ownedFields(owner); entry != null; entry = entry.next) { + projected = collectProjectedFieldView(entry.get(), projected); + } + return projected; + } Map>> fieldStripe = stateStripe(FIELD_CELLS, owner); synchronized (fieldStripe) { @@ -3613,22 +3653,22 @@ private static java.util.List collectProjectedFieldViews( return projected; } for (WeakReference reference : fields.values()) { - FieldCell cell = reference.get(); - if (cell != null - && (cell.hasStructuralView - || cell.hasMemoryView - || cell.hasMemoryOrigins - || cell.hasProjectedViews)) { - if (projected == null) { - projected = new java.util.ArrayList<>(); - } - projected.add(cell); - } + projected = collectProjectedFieldView(reference.get(), projected); } } return projected; } + private static java.util.List collectProjectedFieldView( + FieldCell cell, java.util.List projected) { + if (cell != null && (cell.hasStructuralView || cell.hasMemoryView + || cell.hasMemoryOrigins || cell.hasProjectedViews)) { + if (projected == null) projected = new java.util.ArrayList<>(); + projected.add(cell); + } + return projected; + } + private static void discardProjectedFieldViews(Object owner) { java.util.List projected = collectProjectedFieldViews(owner, null); if (projected == null) { @@ -3662,6 +3702,12 @@ private static void discardProjectedFieldViews(Object owner) { } private static boolean hasProjectedFieldCells(Object owner) { + if (owner instanceof Cell || owner instanceof FieldCell) { + for (FieldCellCache entry = ownedFields(owner); entry != null; entry = entry.next) { + if (entry.get() != null) return true; + } + return false; + } if (!mayHaveFieldCells(owner)) { return false; } @@ -3892,7 +3938,7 @@ private static final class Cell { private Object value; private volatile boolean hasStructuralView; private volatile boolean hasMemoryView; - private volatile boolean hasFieldCells; + private volatile FieldCellCache fields; private volatile boolean hasProjectedViews; private Cell(Object value) { @@ -3931,13 +3977,17 @@ private static Object memoryViewAllocation(Object value) { } private static final class FieldCell { + // Share only the same immutable origin and offset. + private volatile AddressOriginState cachedAddressOrigin; + // Cache thin projections only. Exclude published addresses, decoded views, and dynamic metadata. + private volatile Pointer cachedProjection; private final Object fixedOwner; private final Object rootOwner; // Direct enum-payload borrows have no parent Pointer. Retain the // decoded owner's storage while this field cell remains reachable. private final Object memoryBacking; private final FieldAccess access; - private volatile boolean hasFieldCells; + private volatile FieldCellCache fields; private volatile boolean hasProjectedViews; private volatile boolean hasMemoryOrigins; private final int fieldNameHash; @@ -3945,25 +3995,13 @@ private static final class FieldCell { private volatile boolean hasMemoryView; private FieldCell(Object owner, FieldAccess access) { - fixedOwner = owner; - rootOwner = null; - memoryBacking = memoryViewAllocation(owner); - this.access = access; - fieldNameHash = access.fieldNameHash; + this(owner, access, false); } - private FieldCell(Cell owner, FieldAccess access) { - fixedOwner = null; - rootOwner = owner; - memoryBacking = null; - this.access = access; - fieldNameHash = access.fieldNameHash; - } - - private FieldCell(FieldCell owner, FieldAccess access) { - fixedOwner = null; - rootOwner = owner; - memoryBacking = null; + private FieldCell(Object owner, FieldAccess access, boolean rooted) { + fixedOwner = rooted ? null : owner; + rootOwner = rooted ? owner : null; + memoryBacking = rooted ? null : memoryViewAllocation(owner); this.access = access; fieldNameHash = access.fieldNameHash; } @@ -4018,10 +4056,23 @@ private void set(Object value) { // of sibling fields (for example an inline vector's data when its // length changes before writing a newly inserted element). discardProjectedFieldViews(this); - commitOriginMemoryView(owner); - Object identity = ownerIdentity(); - if (identity != owner) { - commitOriginMemoryView(identity); + commitOwners(); + } + + private void commitOwners() { + // An enum payload can belong to a decoded enum that owns the byte storage. + // Intermediate field carriers need not have registered views. + FieldCell current = this; + while (true) { + Object owner = current.owner(); + commitOriginMemoryView(owner); + if (current.rootOwner instanceof FieldCell) { + current = (FieldCell) current.rootOwner; + } else { + Object identity = current.ownerIdentity(); + if (identity != owner) commitOriginMemoryView(identity); + return; + } } } @@ -4423,9 +4474,15 @@ private synchronized boolean isEmpty() { private volatile RarePointerState rareState; private long metadata = -1; + /** Immutable provenance can be shared without coupling derived views. */ private static final class AddressOriginState { - private Pointer addressOrigin; - private long addressOriginOffset; + private final Pointer addressOrigin; + private final long addressOriginOffset; + + private AddressOriginState(Pointer origin, long offset) { + addressOrigin = origin; + addressOriginOffset = offset; + } } private static final class RarePointerState { @@ -4444,19 +4501,6 @@ private static final class RarePointerState { private String traitPointeeCodecClassName; } - private AddressOriginState mutableAddressState() { - AddressOriginState current = addressState; - if (current != null) { - return current; - } - synchronized (this) { - if (addressState == null) { - addressState = new AddressOriginState(); - } - return addressState; - } - } - private RarePointerState mutableRareState() { RarePointerState current = rareState; if (current != null) { @@ -4572,9 +4616,21 @@ private long addressOriginOffset() { } private void setAddressOrigin(Pointer origin, long offset) { - AddressOriginState current = mutableAddressState(); - current.addressOrigin = origin; - current.addressOriginOffset = offset; + AddressOriginState current = addressState; + if (current != null && current.addressOrigin == origin && current.addressOriginOffset == offset) { + return; + } + FieldCell field = byteOffset == 0 && allocation instanceof FieldCell ? (FieldCell) allocation : null; + if (field != null) { + current = field.cachedAddressOrigin; + if (current != null && current.addressOrigin == origin && current.addressOriginOffset == offset) { + addressState = current; + return; + } + } + current = new AddressOriginState(origin, offset); + addressState = current; + if (field != null) field.cachedAddressOrigin = current; } private long publishedAddress() { @@ -4820,9 +4876,18 @@ public static Pointer cell(Object value) { /** Returns a stable, write-through pointer to a generated Rust value field. */ public static Pointer field( Object owner, String fieldName, int size, String codecClassName) { + return field(owner, fieldName, size, codecClassName, true); + } + + private static Pointer field( + Object owner, String fieldName, int size, String codecClassName, boolean cache) { if (owner == null) { throw new NullPointerException("Rust field pointer requires an owner"); } + // Erased zero-sized fields have no Java member. Avoid reflection errors and cache entries. + if (size == 0 && optionalInstanceField(owner.getClass(), fieldName) == null) { + return Pointer.cell(null, 0, codecClassName); + } FieldCell cell; Map>> stripe = stateStripe(FIELD_CELLS, owner); @@ -4847,9 +4912,15 @@ public static Pointer field( } fields.put(fieldName, new WeakReference<>(cell)); } - markFieldCells(owner); + markIdentityFilter(FIELD_CELL_FILTER, owner); } - return new Pointer(cell, size, 0, size, codecClassName); + Pointer cached = cell.cachedProjection; + if (cache && cached != null && cached.metadata == -1 && cached.rareState == null + && cached.addressState == null && cached.viewSize == size + && java.util.Objects.equals(cached.viewCodecClassName, codecClassName)) return cached; + Pointer result = new Pointer(cell, size, 0, size, codecClassName); + if (cache) cell.cachedProjection = result; + return result; } public static Pointer field( @@ -4888,37 +4959,46 @@ private static Pointer rootField( String fieldName, long size, String codecClassName) { + return rootField(owner, ownerClass, fieldName, size, codecClassName, null, 0); + } + + private static Pointer rootField(Object owner, Class ownerClass, String fieldName, + long size, String codecClassName, Pointer source, long displacement) { + // Zero-sized fields have no Java member. + if (size == 0 && optionalInstanceField(ownerClass, fieldName) == null) { + return finishFieldProjection(Pointer.cell(null, 0, codecClassName), source, displacement); + } FieldAccess access; try { access = fieldAccess(ownerClass, fieldName); } catch (NoSuchFieldException error) { if (size == 0) { - return Pointer.cell(null, 0, codecClassName); + return finishFieldProjection(Pointer.cell(null, 0, codecClassName), source, displacement); } throw new IllegalArgumentException( "unknown Rust field " + ownerClass.getName() + "." + fieldName, error); } - String cacheKey = access.cacheKey; - FieldCell cell; - Map>> stripe = - stateStripe(FIELD_CELLS, owner); - synchronized (stripe) { - Map> fields = stripe.get(owner); - if (fields == null) { - fields = new HashMap<>(); - stripe.put(owner, fields); - } - WeakReference reference = fields.get(cacheKey); - cell = reference == null ? null : reference.get(); - if (cell == null) { - cell = owner instanceof Cell - ? new FieldCell((Cell) owner, access) - : new FieldCell((FieldCell) owner, access); - fields.put(cacheKey, new WeakReference<>(cell)); - } - markFieldCells(owner); + FieldCell cell = ownedField(owner, access); + Pointer cached = cell.cachedProjection; + boolean reusable = source != null && source.metadata == -1 && source.rareState == null; + if (reusable && cached != null && cached.metadata == -1 && cached.rareState == null + && cached.viewSize == size && java.util.Objects.equals(cached.viewCodecClassName, codecClassName)) { + AddressOriginState origin = source.addressState; + Pointer ownerOrigin = origin == null || origin.addressOrigin == null ? source : origin.addressOrigin; + long offset = origin == null || origin.addressOrigin == null ? displacement + : Math.addExact(origin.addressOriginOffset, displacement); + AddressOriginState previous = cached.addressState; + if (previous != null && previous.addressOrigin == ownerOrigin && previous.addressOriginOffset == offset) + return cached; } - return new Pointer(cell, checkedArrayLength(size), 0, size, codecClassName); + Pointer result = finishFieldProjection( + new Pointer(cell, checkedArrayLength(size), 0, size, codecClassName), source, displacement); + if (reusable) cell.cachedProjection = result; + return result; + } + + private static Pointer finishFieldProjection(Pointer pointer, Pointer source, long displacement) { + return source == null ? pointer : pointer.withMetadata(source.metadata).inheritAddressOrigin(source, displacement); } private Object compatibleStructView(String ownerClassName) { @@ -4990,20 +5070,11 @@ public Pointer projectStructField( long fieldOffset, long fieldSize, String fieldCodecClassName) { - Class ownerClass = null; - Class fieldType = null; - if (fieldSize != 0 || traitMetadataCarrier() != null) { - try { - ownerClass = resolvedRuntimeClass(ownerClassName); - fieldType = instanceField(ownerClass, fieldName).getType(); - } catch (ClassNotFoundException error) { - throw new IllegalArgumentException( - "unknown Rust aggregate class " + ownerClassName, error); - } catch (NoSuchFieldException error) { - throw new IllegalArgumentException( - "unknown Rust field " + ownerClassName + "." + fieldName, error); - } + // Rust field offsets locate byte storage without Java reflection or a decoded parent. + if (allocation instanceof byte[] && traitMetadataCarrier() == null) { + return byteOffsetRetype(fieldOffset, fieldSize, fieldCodecClassName); } + Class ownerClass = null; if ((allocation instanceof Cell || allocation instanceof FieldCell) && byteOffset == 0 && viewSize == allocationElementSize) { @@ -5017,15 +5088,26 @@ public Pointer projectStructField( ownerClass, fieldName, fieldSize, - fieldCodecClassName) - .withMetadata(metadata) - .inheritAddressOrigin(this, fieldOffset); + fieldCodecClassName, this, fieldOffset); } } catch (ClassNotFoundException error) { throw new IllegalArgumentException( "unknown Rust aggregate class " + ownerClassName, error); } } + Class fieldType = null; + if (fieldSize != 0 || traitMetadataCarrier() != null) { + try { + if (ownerClass == null) ownerClass = resolvedRuntimeClass(ownerClassName); + fieldType = instanceField(ownerClass, fieldName).getType(); + } catch (ClassNotFoundException error) { + throw new IllegalArgumentException( + "unknown Rust aggregate class " + ownerClassName, error); + } catch (NoSuchFieldException error) { + throw new IllegalArgumentException( + "unknown Rust field " + ownerClassName + "." + fieldName, error); + } + } boolean managedField = fieldType == null || !fieldType.isPrimitive(); // Replaceable roots and nested fields use rootField above. Only an // existing carrier may supply a live field here: decoding the whole @@ -5050,7 +5132,7 @@ public Pointer projectStructField( "unknown Rust aggregate class " + ownerClassName, error); } if (owner != null) { - return field(owner, fieldName, fieldSize, fieldCodecClassName) + return field(owner, fieldName, checkedArrayLength(fieldSize), fieldCodecClassName, false) .withMetadata(metadata) .inheritAddressOrigin(this, fieldOffset); } @@ -5161,8 +5243,7 @@ public static Pointer array( -1) .withMetadata(origin.metadata); if (activeMemoryViewMatches(array, origin)) { - return source.byte_offset(relativeOffset) - .retype(elementSize, codecClassName); + return source.byteOffsetRetype(relativeOffset, elementSize, codecClassName); } return new Pointer( array, @@ -6257,6 +6338,7 @@ public static Pointer attachTraitObjectCarrier( && (pointeeAlignment & (pointeeAlignment - 1)) != 0)) { throw new IllegalArgumentException("invalid trait-object pointee layout"); } + pointer = independentFieldMetadata(pointer); pointer.traitObjectCarrier(carrier); pointer.traitPointeeSize(pointeeSize); pointer.traitPointeeAlignment(pointeeAlignment); @@ -6308,6 +6390,7 @@ public static Pointer attachStructTailTraitMetadata(Pointer pointer, Pointer tai "struct-tail trait pointer does not carry dynamic metadata"); } TraitObjectCarrier traitCarrier = (TraitObjectCarrier) carrier; + pointer = independentFieldMetadata(pointer); pointer.traitMetadataCarrier(carrier); pointer.traitMetadataMarker(tailPointer.traitMetadataMarker()); pointer.traitPointeeSize(tailPointer.traitPointeeSize() >= 0 @@ -6634,36 +6717,45 @@ private Pointer copyDynamicMetadata(Pointer source) { } private Pointer inheritAddressOrigin(Pointer source, long additionalOffset) { - Pointer origin = source.addressOrigin(); - // A decoded carrier's origin is weakly indexed. Keep its backing cell - // alive while any projected pointer can still mutate that carrier; - // flattening past it lets GC detach live aliases from their storage. - if (origin == null || source.boundMemoryViewState() != null) { + AddressOriginState current = source.addressState; + // Keep decoded backing cells alive for as long as a projected alias. + if (current == null || current.addressOrigin == null || source.boundMemoryViewState() != null) { setAddressOrigin(source, additionalOffset); + } else if (additionalOffset == 0) { + addressState = current; } else { - setAddressOrigin( - origin, - Math.addExact(source.addressOriginOffset(), additionalOffset)); + setAddressOrigin(current.addressOrigin, + Math.addExact(current.addressOriginOffset, additionalOffset)); } return this; } private Pointer copyAddressOrigin(Pointer source, long additionalOffset) { - Pointer origin = source.addressOrigin(); - if (origin != null) { - setAddressOrigin( - origin, - Math.addExact(source.addressOriginOffset(), additionalOffset)); + AddressOriginState current = source.addressState; + if (current != null && current.addressOrigin != null) { + if (additionalOffset == 0) { + addressState = current; + } else { + // Origin offsets are address words and must support wrapping arithmetic. + setAddressOrigin(current.addressOrigin, + current.addressOriginOffset + additionalOffset); + } } return this; } + private static Pointer independentFieldMetadata(Pointer pointer) { + // Detach metadata even if this field pointer is no longer the cached projection. + return pointer.allocation instanceof FieldCell + ? pointer.retype(pointer.viewSize, pointer.viewCodecClassName) : pointer; + } + public static Pointer withMetadata(Pointer pointer, long metadata) { - return pointer.withMetadata(metadata); + return (pointer.metadata == metadata ? pointer : independentFieldMetadata(pointer)).withMetadata(metadata); } public static Pointer withMetadata(Pointer pointer, int metadata) { - return pointer.withMetadata(Integer.toUnsignedLong(metadata)); + return withMetadata(pointer, Integer.toUnsignedLong(metadata)); } public long metadata() { @@ -11474,7 +11566,7 @@ private Pointer nominalManagedPointee(String targetClassName) { if (match == null) { return this; } - return field(value, match.getName(), viewSize, viewCodecClassName) + return field(value, match.getName(), checkedArrayLength(viewSize), viewCodecClassName, false) .withMetadata(metadata) .inheritAddressOrigin(this, 0); } catch (ClassNotFoundException error) { diff --git a/tests/integration/pointer_provenance/CyclicFieldViews.java b/tests/integration/pointer_provenance/CyclicFieldViews.java new file mode 100644 index 00000000..0f69109b --- /dev/null +++ b/tests/integration/pointer_provenance/CyclicFieldViews.java @@ -0,0 +1,51 @@ +import java.lang.ref.WeakReference; +import java.lang.reflect.Field; +import java.util.Map; +import org.rustlang.runtime.Pointer; + +/** A decoded field may contain a reference back to its enclosing allocation. */ +public final class CyclicFieldViews { + public static final class Holder { public long bits = 41; } + public static final class View { + public Pointer parent; + View(Pointer parent) { this.parent = parent; } + } + private static Pointer decodingParent; + public static final class Codec { + public static byte[] e$view(View value) { return new byte[8]; } + public static View d$view(byte[] bytes) { return new View(decodingParent); } + } + + private static WeakReference abandoned() { + Pointer root = Pointer.cell(new Holder(), 8, null); + Pointer field = root.projectStructField(Holder.class.getName(), "bits", 0, 8, null); + // Only the cached decoded view retains this parent after return. + decodingParent = root; + try { + View view = (View) field.retype(8, + "CyclicFieldViews$Codec#view#LCyclicFieldViews$View;").getObject(); + if (view.parent != root) throw new AssertionError("decoded parent"); + } finally { + decodingParent = null; + } + return new WeakReference<>(root); + } + + public static void check() throws Exception { + WeakReference root = abandoned(); + for (int i = 0; i < 20; i++) { + System.gc(); + Thread.sleep(20); + // Drain the weak maps so the field key and then its decoded parent can be collected. + for (String name : new String[] {"MEMORY_VIEWS", "MEMORY_VIEW_ORIGINS"}) { + Field field = Pointer.class.getDeclaredField(name); + field.setAccessible(true); + for (Map stripe : (Map[]) field.get(null)) { + synchronized (stripe) { stripe.size(); } + } + } + if (root.get() == null) return; + } + throw new AssertionError("decoded field cache retained an abandoned parent"); + } +} diff --git a/tests/integration/pointer_provenance/FieldProjections.java b/tests/integration/pointer_provenance/FieldProjections.java new file mode 100644 index 00000000..a3b3542c --- /dev/null +++ b/tests/integration/pointer_provenance/FieldProjections.java @@ -0,0 +1,130 @@ +import org.rustlang.runtime.Pointer; + +/** Repeated field borrows retain a live place without rebuilding its wrapper. */ +public final class FieldProjections { + public static class Pair { + public int first, second; + Pair(int first, int second) { this.first = first; this.second = second; } + } + + public static final class Outer { + public Pair inner; + Outer(Pair inner) { this.inner = inner; } + } + + public static final class Derived extends Pair { + Derived(int first, int second) { super(first, second); } + } + + private static Pointer second(Pointer root) { + return root.projectStructField(Pair.class.getName(), "second", 4, 4, null); + } + + private static boolean thin(Pointer pointer) { + try { + pointer.metadata(); + return false; + } catch (IllegalStateException expected) { + return true; + } + } + + private static java.lang.ref.WeakReference abandonedProjection() { + Pointer root = Pointer.cellAligned(new Pair(107, 109), 8, null, 4); + // Exercise address bookkeeping as well as the owner/projection cycle. + second(root).addr(); + return new java.lang.ref.WeakReference<>(root); + } + + public static void check() throws Exception { + Pair fixedOwner = new Pair(5, 7); + Pointer fixedFirst = Pointer.field(fixedOwner, "second", 4, null); + Pointer fixedSecond = Pointer.field(fixedOwner, "second", 4, null); + Pointer fixedWide = Pointer.withMetadata(fixedFirst, 13); + if (fixedWide.metadata() != 13 || !thin(fixedSecond)) throw new AssertionError("fixed field metadata alias"); + Pointer fixedTrait = Pointer.attachTraitObjectCarrier(fixedFirst, "fixed", 4, 4); + if (!"fixed".equals(fixedTrait.getObject()) || fixedSecond.getI32() != 7) { + throw new AssertionError("fixed field trait alias"); + } + Pointer.field(fixedOwner, "second", 1, null); + Pointer oldWide = Pointer.withMetadata(fixedFirst, 17); + if (oldWide.metadata() != 17 || !thin(fixedSecond)) throw new AssertionError("evicted fixed field metadata alias"); + + Pointer root = Pointer.cellAligned(new Pair(17, 31), 8, null, 4); + Pointer field = second(root); + if (second(root) != field) throw new AssertionError("unchanged field projection was rebuilt"); + field.set(43); + if (second(root).getI32() != 43) throw new AssertionError("field write"); + root.set(new Pair(53, 59)); + if (field.getI32() != 59) throw new AssertionError("parent replacement"); + Pointer priorCached = second(root); + Pointer narrow = root.projectStructField(Pair.class.getName(), "second", 4, 1, null); + Pointer detached = Pointer.withMetadata(priorCached, 5); + if (detached.metadata() != 5 || !thin(field)) throw new AssertionError("evicted cache metadata alias"); + Pointer trait = Pointer.attachTraitObjectCarrier(priorCached, "carrier", 4, 4); + if (!"carrier".equals(trait.getObject()) || !Integer.valueOf(59).equals(priorCached.getObject())) { + throw new AssertionError("trait attachment changed an existing field borrow"); + } + if (second(root).getI32() != 59 || narrow.getI8() != 59) { + throw new AssertionError("projection width changed another view"); + } + Pointer earlier = second(root); + field = Pointer.withMetadata(field, 7); + if (!thin(earlier)) throw new AssertionError("late metadata changed an existing borrow"); + Pointer other = second(root); + if (!thin(other) || field.metadata() != 7) throw new AssertionError("projection metadata"); + Pointer.withMetadata(root, 11); + Pointer metadata = second(root); + if (metadata.metadata() != 11 || !thin(other)) throw new AssertionError("parent metadata"); + + Pointer outer = Pointer.cellAligned(new Outer(new Pair(67, 71)), 8, null, 4); + Pointer inner = outer.projectStructField(Outer.class.getName(), "inner", 0, 8, null); + Pointer nested = second(inner); + outer.set(new Outer(new Pair(73, 79))); + if (nested.getI32() != 79) throw new AssertionError("nested parent replacement"); + long address = nested.addr(); + if (second(inner).addr() != address || nested.byte_offset(-4).addr() != address - 4) { + throw new AssertionError("projection address identity"); + } + inner = null; + outer = null; + System.gc(); + if (nested.getI32() != 79 || nested.addr() != address) { + throw new AssertionError("projection backing lifetime"); + } + Pointer inherited = Pointer.cellAligned(new Derived(97, 101), 8, null, 4); + Pointer baseField = second(inherited); + Pointer derivedField = inherited.projectStructField( + Derived.class.getName(), "second", 4, 4, null); + if (baseField != derivedField) throw new AssertionError("inherited field lost storage identity"); + derivedField.set(103); + if (baseField.getI32() != 103) throw new AssertionError("inherited field alias"); + Pointer concurrent = Pointer.cellAligned(new Pair(83, 89), 8, null, 4); + java.util.concurrent.atomic.AtomicReference failure = new java.util.concurrent.atomic.AtomicReference<>(); + Thread[] threads = new Thread[4]; + for (int i = 0; i < threads.length; i++) { + threads[i] = new Thread(() -> { + try { + for (int j = 0; j < 5000; j++) { + if (second(concurrent).getI32() != 89) throw new AssertionError("concurrent read"); + Pointer first = concurrent.projectStructField( + Pair.class.getName(), "first", 0, 4, null); + if (first.getI32() != 83) throw new AssertionError("concurrent second field"); + } + } catch (Throwable error) { + failure.set(error); + } + }); + } + for (Thread thread : threads) thread.start(); + for (Thread thread : threads) thread.join(); + if (failure.get() != null) throw new AssertionError(failure.get()); + java.lang.ref.WeakReference abandoned = abandonedProjection(); + for (int i = 0; i < 10; i++) { + System.gc(); + if (abandoned.get() == null) return; + Thread.sleep(20); + } + throw new AssertionError("field cache retained an abandoned allocation"); + } +} diff --git a/tests/integration/pointer_provenance/Main.java b/tests/integration/pointer_provenance/Main.java index 6e41af94..5ee2bf59 100644 --- a/tests/integration/pointer_provenance/Main.java +++ b/tests/integration/pointer_provenance/Main.java @@ -6,6 +6,8 @@ public class Main { public static void main(String[] args) throws Exception { ArrayViews.check(); MemoryViews.check(); + FieldProjections.check(); + CyclicFieldViews.check(); CodecAdapters.check(); StructuralViews.check(); Field field = Pointer.class.getDeclaredField("EXPOSED_ADDRESSES"); From 202bc00cc1e3a3919d04194e8b1cafce2e8ea683 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 02:34:39 +1000 Subject: [PATCH 05/61] use component storage and shared carriers --- runtime/src/BorrowedFieldPath.java | 40 + runtime/src/BorrowedStorage.java | 47 + runtime/src/KotlinFutureInterop.java | 26 +- runtime/src/Pointer.java | 1662 +++++++---------- runtime/src/RustClasses.java | 31 + runtime/src/RustField.java | 181 ++ runtime/src/SliceView.java | 100 + runtime/src/Storage.java | 46 + runtime/src/StorageLayout.java | 170 ++ runtime/src/TaggedLong.java | 22 + runtime/src/Utf8View.java | 27 + .../pointer_provenance/BorrowedFields.java | 181 ++ .../integration/pointer_provenance/Main.java | 3 + .../pointer_provenance/RangeCodecs.java | 35 + .../pointer_provenance/TypedFields.java | 118 ++ .../integration/pointer_provenance/src/lib.rs | 22 + 16 files changed, 1738 insertions(+), 973 deletions(-) create mode 100644 runtime/src/BorrowedFieldPath.java create mode 100644 runtime/src/BorrowedStorage.java create mode 100644 runtime/src/RustClasses.java create mode 100644 runtime/src/RustField.java create mode 100644 runtime/src/SliceView.java create mode 100644 runtime/src/Storage.java create mode 100644 runtime/src/StorageLayout.java create mode 100644 runtime/src/TaggedLong.java create mode 100644 runtime/src/Utf8View.java create mode 100644 tests/integration/pointer_provenance/BorrowedFields.java create mode 100644 tests/integration/pointer_provenance/RangeCodecs.java create mode 100644 tests/integration/pointer_provenance/TypedFields.java diff --git a/runtime/src/BorrowedFieldPath.java b/runtime/src/BorrowedFieldPath.java new file mode 100644 index 00000000..d30867c3 --- /dev/null +++ b/runtime/src/BorrowedFieldPath.java @@ -0,0 +1,40 @@ +package org.rustlang.runtime; + +/** One allocation-owned projection recipe, shared by every reference to a field. */ +final class BorrowedFieldPath { + final int offset, size; + final boolean view; + final String codec; + final RustField[] fields; + + BorrowedFieldPath(Class type, int offset, int size, String path, boolean view, String codec) + throws ReflectiveOperationException { + this.offset = offset; + this.size = size; + this.view = view; + this.codec = codec; + String[] names = path.split("/"); + fields = new RustField[names.length]; + for (int i = 0; i < names.length; i++) { + fields[i] = RustField.find(type, names[i]); + type = fields[i].getType(); + } + } + + RustField field() { return fields[fields.length - 1]; } + + private Object owner(Object root) throws IllegalAccessException { + for (int i = 0; i < fields.length - 1; i++) root = fields[i].get(root); + return root; + } + + Object read(Object root, long[] metadata) { + try { return field().borrowedParts(owner(root), metadata); } + catch (IllegalAccessException error) { throw new IllegalStateException(error); } + } + + void write(Object root, Object backing, long offset, long length) { + try { field().setBorrowedParts(owner(root), backing, offset, length); } + catch (IllegalAccessException error) { throw new IllegalStateException(error); } + } +} diff --git a/runtime/src/BorrowedStorage.java b/runtime/src/BorrowedStorage.java new file mode 100644 index 00000000..23f32169 --- /dev/null +++ b/runtime/src/BorrowedStorage.java @@ -0,0 +1,47 @@ +package org.rustlang.runtime; + +/** Stores local reference components until a boundary is needed. + * After materialization, Cell.value holds the value shared by all aliases. + */ +final class BorrowedStorage extends Storage { + private Object root; + private long offset, length; + private int plan; + final boolean view; + boolean split; + + BorrowedStorage(Object value, int size, String codec, long metadata, String layout, boolean view) { + super(value, size, codec, metadata, layout); + this.view = view; + } + + void store(Object root, long offset, long length, int plan) { + this.root = root; + this.offset = offset; + this.length = length; + this.plan = plan; + value = null; + split = true; + } + + Object read(long[] metadata, boolean view) { + if (split) { + metadata[0] = offset; + if (view) metadata[1] = length; + return root; + } + metadata[0] = view ? RustField.viewStart((SliceView) value) : 0; + if (view) metadata[1] = RustField.viewLength((SliceView) value); + return view ? RustField.viewRoot((SliceView) value) : value; + } + + @Override void materialize() { + if (!split) return; + if (plan < 0) value = plan == -2 + ? new Utf8View(root, Math.toIntExact(offset), length) + : new SliceView(root, Math.toIntExact(offset), length); + else value = root == null && offset == 0 ? null : Pointer.addressFromParts(root, offset, plan); + split = false; + root = null; + } +} diff --git a/runtime/src/KotlinFutureInterop.java b/runtime/src/KotlinFutureInterop.java index 37574e54..9446a0fb 100644 --- a/runtime/src/KotlinFutureInterop.java +++ b/runtime/src/KotlinFutureInterop.java @@ -1,7 +1,6 @@ package org.rustlang.runtime; import java.lang.reflect.Constructor; -import java.lang.reflect.Field; import java.lang.reflect.InvocationHandler; import java.lang.reflect.InvocationTargetException; import java.lang.reflect.Method; @@ -77,7 +76,6 @@ private static final class PollAdapter { private final Constructor rawWakerConstructor; private final Constructor wakerConstructor; private final Constructor contextConstructor; - private final Constructor pinConstructor; private final Method resume; private final Pointer vtablePointer; private final int futureSize; @@ -102,10 +100,7 @@ private PollAdapter( "org.rustlang.core.task.wake.Waker", true, loader); Class contextClass = Class.forName( "org.rustlang.core.task.wake.Context", true, loader); - Class pinClass = Class.forName( - "org.rustlang.core.pin.Pin_MutRef" + futureClass.getSimpleName(), - true, - loader); + resume = findResumeMethod(futureClass); rawWakerConstructor = rawWakerClass.getConstructor(Pointer.class, Pointer.class); Constructor vtableConstructor = constructorWithArity(rawWakerVTableClass, 4); @@ -163,8 +158,6 @@ public Object invoke(Object proxy, Method method, Object[] args) { wakerConstructor = wakerClass.getConstructor(rawWakerClass); contextConstructor = rustConstructor(contextClass); - pinConstructor = pinClass.getConstructor(Pointer.class); - resume = findResumeMethod(futureClass, pinClass); this.futureSize = futureSize; this.futureCodec = futureCodec; this.futureAlignment = futureAlignment; @@ -190,14 +183,13 @@ private Object poll(Object future, Runnable wake) { Object context = contextConstructor.newInstance(contextArguments); Pointer futurePointer = Pointer.receiverCellAligned( future, futureSize, futureCodec, futureAlignment); - Object pin = pinConstructor.newInstance(futurePointer); - Object result = resume.invoke(null, pin, Pointer.cell(context)); + Object result = resume.invoke(null, futurePointer, Pointer.cell(context)); - if (result.getClass().getName().endsWith("$Pending")) { + if (result.getClass().getSimpleName().equals("Pending")) { return PENDING; } try { - Field output = result.getClass().getField("value"); + RustField output = RustField.find(result.getClass(), "value"); return output.get(result); } catch (NoSuchFieldException noOutput) { return UNIT; @@ -235,8 +227,7 @@ private static Constructor rustConstructor(Class type) throws NoSuchMethod private static Object defaultRustValue(Class type) throws ReflectiveOperationException { if (type.isInterface()) { - Class none = Class.forName( - type.getName() + "$None", true, type.getClassLoader()); + Class none = RustClasses.nested(type, "None"); return none.getConstructor().newInstance(); } try { @@ -247,17 +238,18 @@ private static Object defaultRustValue(Class type) throws ReflectiveOperation } } - private static Method findResumeMethod(Class futureClass, Class pinClass) + private static Method findResumeMethod(Class futureClass) throws ReflectiveOperationException { + // Body is an executable namespace, never a shared storage carrier. Class bodyClass = Class.forName( futureClass.getName() + "$Body", true, futureClass.getClassLoader()); for (Method method : bodyClass.getMethods()) { Class[] parameters = method.getParameterTypes(); if (Modifier.isStatic(method.getModifiers()) && parameters.length == 2 - && parameters[0] == pinClass + && parameters[0] == Pointer.class && parameters[1] == Pointer.class - && method.getReturnType().getName().contains(".task.poll.Poll_")) { + && RustClasses.hasNested(method.getReturnType(), "Ready", "Pending")) { return method; } } diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index b47c91b0..bb989ed9 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -4,7 +4,6 @@ import java.lang.ref.WeakReference; import java.lang.reflect.Array; import java.lang.reflect.Constructor; -import java.lang.reflect.Field; import java.lang.reflect.InvocationTargetException; import java.lang.reflect.Method; import java.lang.reflect.Modifier; @@ -27,10 +26,7 @@ import java.util.concurrent.atomic.AtomicLongArray; public final class Pointer { - private static final String RELATIVE_POINTER_ELEMENT_OFFSET_SUFFIX = - "$rcj$elementOffset"; - private static final String RELATIVE_POINTER_BYTE_OFFSET_SUFFIX = - "$rcj$byteOffset"; + private static Object arrayGet(Object array, int index) { if (array instanceof byte[]) { @@ -128,44 +124,6 @@ public static void dropRustValue(Object value) { } } - private static void dropTraitPointer(Object pointer) { - if (pointer == null) { - return; - } - if (pointer instanceof TraitObjectCarrier) { - dropRustValue(((TraitObjectCarrier) pointer).rustTraitObjectPayload()); - return; - } - if (pointer instanceof Pointer) { - Object payload = ((Pointer) pointer).getObject(); - if (payload != pointer) { - dropRustValue(payload); - } - return; - } - if (isSliceViewCarrierType(pointer.getClass())) { - try { - SliceAccess access = SLICE_ACCESSES.get(pointer.getClass()); - Object backing = access.array(pointer); - int offset = access.offset(pointer); - if (backing instanceof Pointer) { - Object payload = ((Pointer) backing).add(offset).directCellValueOrSelf(); - if (payload != backing) { - dropRustValue(payload); - } - return; - } - if (backing != null && backing.getClass().isArray()) { - dropRustValue(arrayGet(backing, offset)); - return; - } - } catch (ReflectiveOperationException error) { - throw new IllegalStateException("invalid Rust trait-object pointer", error); - } - } - dropRustValue(pointer); - } - public static boolean catchUnwind(Object tryFunction, Pointer data, Object catchFunction) { try { invokeRustFunction(tryFunction, data); @@ -279,10 +237,6 @@ protected ResolvedClassCache initialValue() { new ConcurrentHashMap<>(); private static final Map SHARED_CONSTANT_ARRAYS = new IdentityHashMap<>(); - private static final ConcurrentHashMap DROP_METHOD_HANDLES = - new ConcurrentHashMap<>(); - private static final ConcurrentHashMap DROP_FIELDS_METHOD_HANDLES = - new ConcurrentHashMap<>(); private static final ConcurrentHashMap CODEC_DESCRIPTORS = new ConcurrentHashMap<>(); private static final ConcurrentHashMap BINARY_CLASS_NAMES = @@ -297,13 +251,6 @@ protected ResolvedClassCache initialValue() { new ConcurrentHashMap<>(); private static final Map>> RESOLVED_CLASSES = new IdentityHashMap<>(); - private static final ClassValue> INSTANCE_FIELDS = - new ClassValue>() { - @Override - protected ConcurrentHashMap computeValue(Class type) { - return new ConcurrentHashMap<>(); - } - }; private static final ClassValue> FIELD_ACCESSORS = new ClassValue>() { @Override @@ -311,37 +258,7 @@ protected ConcurrentHashMap computeValue(Class type) { return new ConcurrentHashMap<>(); } }; - private static final ClassValue SLICE_ACCESSES = - new ClassValue() { - @Override - protected SliceAccess computeValue(Class type) { - return new SliceAccess(type); - } - }; - private static final ClassValue PUBLIC_INSTANCE_FIELDS = - new ClassValue() { - @Override - protected Field[] computeValue(Class type) { - Field[] all = type.getFields(); - int count = 0; - for (Field field : all) { - if (!Modifier.isStatic(field.getModifiers()) - && !field.isSynthetic()) { - count++; - } - } - Field[] fields = new Field[count]; - int index = 0; - for (Field field : all) { - if (!Modifier.isStatic(field.getModifiers()) - && !field.isSynthetic()) { - field.setAccessible(true); - fields[index++] = field; - } - } - return fields; - } - }; + private static final ClassValue PUBLIC_INSTANCE_FIELDS = RustField.ALL; private static final ClassValue> PUBLIC_CONSTRUCTORS_BY_ARITY = new ClassValue>() { @Override @@ -361,36 +278,6 @@ protected Map computeValue(Class type) { return constructors; } }; - private static final ClassValue> SLICE_VIEW_CONSTRUCTORS = - new ClassValue>() { - @Override - protected Constructor computeValue(Class type) { - try { - Constructor constructor = - type.getConstructor(Object.class, int.class, int.class); - constructor.setAccessible(true); - return constructor; - } catch (NoSuchMethodException error) { - throw new IllegalStateException( - "Rust slice view has no array/offset/length constructor", error); - } - } - }; - private static final ClassValue> LONG_SLICE_VIEW_CONSTRUCTORS = - new ClassValue>() { - @Override - protected Constructor computeValue(Class type) { - try { - Constructor constructor = - type.getConstructor(Object.class, int.class, long.class); - constructor.setAccessible(true); - return constructor; - } catch (NoSuchMethodException error) { - throw new IllegalStateException( - "Rust slice view has no long-length constructor", error); - } - } - }; private static final ClassValue RUST_FUNCTION_POINTER_TYPES = new ClassValue() { @Override @@ -408,7 +295,7 @@ protected Boolean computeValue(Class type) { new ClassValue() { @Override protected ManagedCopyPlan computeValue(Class type) { - Field[] fields = PUBLIC_INSTANCE_FIELDS.get(type); + RustField[] fields = PUBLIC_INSTANCE_FIELDS.get(type); ManagedFieldPlan[] fieldPlans = new ManagedFieldPlan[fields.length]; for (int index = 0; index < fields.length; index++) { try { @@ -1140,16 +1027,11 @@ private static void recordAlignment(Object allocation, int alignment) { } private static boolean isSliceViewType(Class type) { - return type != null && SLICE_VIEW_CLASS_NAME.equals(type.getName()); + return type == SliceView.class; } private static boolean isSliceViewCarrierType(Class type) { - for (Class current = type; current != null; current = current.getSuperclass()) { - if (SLICE_VIEW_CLASS_NAME.equals(current.getName())) { - return true; - } - } - return false; + return type != null && SliceView.class.isAssignableFrom(type); } /** @@ -1290,11 +1172,11 @@ private static final class ManagedFieldPlan { private final MethodHandle getter; private final MethodHandle setter; - private ManagedFieldPlan(Field field) throws IllegalAccessException { + private ManagedFieldPlan(RustField field) throws IllegalAccessException { MethodHandles.Lookup lookup = MethodHandles.lookup(); - getter = lookup.unreflectGetter(field).asType(MethodType.methodType( + getter = field.getter().asType(MethodType.methodType( Object.class, Object.class)); - setter = lookup.unreflectSetter(field).asType(MethodType.methodType( + setter = field.setter().asType(MethodType.methodType( void.class, Object.class, Object.class)); } @@ -1435,224 +1317,6 @@ public static void overwriteManagedObject(Object target, Object replacement) { copyStructuralFields(replacement, target); } - /** Drops a Box pointee and always runs the Box deallocator during unwinding. */ - public static void dropBoxWithCleanup( - Object pointee, - Object box, - String pointeeOwner, - String pointeeMethod, - String pointeeDescriptor, - String boxOwner, - String boxMethod, - String boxDescriptor) { - Throwable pendingDropFailure = null; - try { - if (pointeeOwner.isEmpty()) { - dropTraitPointer(pointee); - } else { - invokeDropMethod(pointee, pointeeOwner, pointeeMethod, pointeeDescriptor); - } - } catch (Throwable failure) { - PanicSupport.abortIfStackOverflow(failure); - if (Boolean.getBoolean("org.rustlang.debugUnwind")) { - failure.printStackTrace(System.err); - } - pendingDropFailure = failure; - } - try { - invokeDropMethod(box, boxOwner, boxMethod, boxDescriptor); - } catch (Throwable failure) { - PanicSupport.abortIfStackOverflow(failure); - if (Boolean.getBoolean("org.rustlang.debugUnwind")) { - failure.printStackTrace(System.err); - } - if (pendingDropFailure != null) { - Runtime.getRuntime().halt(134); - } - pendingDropFailure = failure; - } - if (pendingDropFailure != null) { - rethrowUnchecked(pendingDropFailure); - } - } - - /** Runs a custom destructor and always drops the value's fields afterward. */ - public static void dropAdtWithCleanup( - Object value, - String owner, - String dropMethod, - String dropDescriptor, - String fieldsMethod) { - Throwable pendingDropFailure = null; - try { - invokeDropMethod(value, owner, dropMethod, dropDescriptor); - } catch (Throwable failure) { - PanicSupport.abortIfStackOverflow(failure); - if (Boolean.getBoolean("org.rustlang.debugUnwind")) { - failure.printStackTrace(System.err); - } - pendingDropFailure = failure; - } - try { - String key = value.getClass().getName() + '\0' + fieldsMethod; - MethodHandle handle = DROP_FIELDS_METHOD_HANDLES.get(key); - if (handle == null) { - MethodHandle resolved = MethodHandles.publicLookup().findVirtual( - value.getClass(), fieldsMethod, MethodType.methodType(void.class)); - MethodHandle previous = DROP_FIELDS_METHOD_HANDLES.putIfAbsent(key, resolved); - handle = previous == null ? resolved : previous; - } - handle.invokeWithArguments(value); - } catch (Throwable failure) { - PanicSupport.abortIfStackOverflow(failure); - if (Boolean.getBoolean("org.rustlang.debugUnwind")) { - failure.printStackTrace(System.err); - } - if (pendingDropFailure != null) { - Runtime.getRuntime().halt(134); - } - pendingDropFailure = failure; - } - if (pendingDropFailure != null) { - rethrowUnchecked(pendingDropFailure); - } - } - - private static void invokeDropMethod( - Object value, String ownerName, String methodName, String descriptor) - throws Throwable { - String key = ownerName + '\0' + methodName + '\0' + descriptor; - MethodHandle handle = DROP_METHOD_HANDLES.get(key); - if (handle == null) { - Class owner = resolvedRuntimeClass(ownerName); - MethodType methodType = - MethodType.fromMethodDescriptorString(descriptor, owner.getClassLoader()); - MethodHandle resolved = - MethodHandles.publicLookup().findStatic(owner, methodName, methodType); - MethodHandle previous = DROP_METHOD_HANDLES.putIfAbsent(key, resolved); - handle = previous == null ? resolved : previous; - } - Class parameterType = handle.type().parameterType(0); - Object argument = parameterType.isInstance(value) - ? value - : parameterType == Pointer.class && isSliceViewCarrierType(value.getClass()) - ? fromSlice(value) - : adaptStructuralField(value, parameterType); - handle.invokeWithArguments(argument); - } - - /** Runs generated Rust element drop glue over a dynamically sized slice. */ - public static void dropSlice(Object slice, String ownerClassName, String methodName) { - dropSlice( - slice, - ownerClassName, - methodName, - "(Lorg/rustlang/runtime/Pointer;)V", - -1, - null); - } - - /** Runs generated Rust element drop glue using its actual JVM carrier ABI. */ - public static void dropSlice( - Object slice, String ownerClassName, String methodName, String descriptor) { - dropSlice(slice, ownerClassName, methodName, descriptor, -1, null); - } - - /** Runs slice drop glue through the element type's logical memory view. */ - public static void dropSlice( - Object slice, - String ownerClassName, - String methodName, - String descriptor, - long elementViewSize, - String elementViewCodec) { - if (slice == null) { - return; - } - try { - Class sliceClass = slice.getClass(); - Object array = instanceField(sliceClass, "array").get(slice); - int offset = instanceField(sliceClass, "offset").getInt(slice); - int length = instanceField(sliceClass, "length").getInt(slice); - Pointer data = array instanceof Pointer - ? ((Pointer) array).sliceElementView().add(offset) - : null; - if (data == null && (array == null || !array.getClass().isArray())) { - throw new IllegalArgumentException("Rust slice drop requires array-backed storage"); - } - MethodHandle drop = null; - Throwable pendingDropFailure = null; - for (int index = 0; index < length; index++) { - Object stored = data == null ? arrayGet(array, offset + index) : null; - boolean alreadyTyped = stored != null - && elementViewSize >= 0 - && isGeneratedAggregateCodec(elementViewCodec) - && codecPlan(elementViewCodec).encodeParameterType.isInstance(stored); - Pointer element = alreadyTyped - ? Pointer.cell(stored, elementViewSize, elementViewCodec) - : data == null ? Pointer.cell(stored) : data.add(index); - if (!alreadyTyped && (elementViewSize >= 0 || elementViewCodec != null)) { - element = element.retype( - elementViewSize >= 0 ? elementViewSize : element.viewSize, - elementViewCodec); - if (elementViewSize == 0 && elementViewCodec != null) { - // `MaybeUninit` slice drop is an explicit initialized - // `T` view. Do not let the erased ZST source wrapper win - // over the target codec merely because both occupy zero bytes. - element.clearZeroSizedSourceView(); - } - } - Object managed = element.getObject(); - try { - if (managed instanceof RustDrop) { - ((RustDrop) managed).rustDrop(); - } else { - if (drop == null) { - Class owner = resolvedRuntimeClass(ownerClassName); - MethodType methodType = MethodType.fromMethodDescriptorString( - descriptor, owner.getClassLoader()); - drop = MethodHandles.publicLookup().findStatic( - owner, - methodName, - methodType); - } - Class parameterType = drop.type().parameterType(0); - Object argument = parameterType == Pointer.class - ? element - : adaptStructuralField(element.getObject(), parameterType); - drop.invokeWithArguments(argument); - } - } catch (Throwable failure) { - PanicSupport.abortIfStackOverflow(failure); - if (pendingDropFailure != null) { - Runtime.getRuntime().halt(134); - } - pendingDropFailure = failure; - } - } - if (pendingDropFailure instanceof RuntimeException) { - throw (RuntimeException) pendingDropFailure; - } - if (pendingDropFailure instanceof Error) { - throw (Error) pendingDropFailure; - } - if (pendingDropFailure != null) { - throw new IllegalStateException("Rust slice element drop failed", pendingDropFailure); - } - } catch (ReflectiveOperationException error) { - throw new IllegalStateException("could not invoke Rust slice element drop glue", error); - } catch (Throwable error) { - PanicSupport.abortIfStackOverflow(error); - if (error instanceof RuntimeException) { - throw (RuntimeException) error; - } - if (error instanceof Error) { - throw (Error) error; - } - throw new IllegalStateException("Rust slice element drop failed", error); - } - } - /** Creates an independent JVM carrier for a copied Rust aggregate value. */ public static Object copyManagedValue(Object value) { if (value == null) { @@ -2409,14 +2073,7 @@ private static Object decodeFatPointer( data = typedPointerObjectFromAddress(dataAddress, codec); } data = data.retype(elementSize, elementCodec); - try { - Class viewClass = resolvedRuntimeClass(descriptor[0]); - return LONG_SLICE_VIEW_CONSTRUCTORS - .get(viewClass) - .newInstance(data, 0, pointerMetadata); - } catch (ReflectiveOperationException error) { - throw new IllegalStateException("could not reconstruct Rust slice fat pointer", error); - } + return SliceView.create(descriptor[0], data, 0, pointerMetadata); } /** Writes a JVM fat-pointer carrier using Rust's native two-word layout. */ @@ -2551,28 +2208,13 @@ private static Class resolvedRuntimeClass(String className) return resolvedClass(className, RUNTIME_CLASS_LOADER); } - private static Field instanceField(Class owner, String name) + private static RustField instanceField(Class owner, String name) throws NoSuchFieldException { - ConcurrentHashMap fields = INSTANCE_FIELDS.get(owner); - Field cached = fields.get(name); - if (cached != null) { - return cached; - } - Field field = owner.getField(name); - if (Modifier.isStatic(field.getModifiers())) { - throw new NoSuchFieldException(owner.getName() + "." + name + " is static"); - } - field.setAccessible(true); - Field previous = fields.putIfAbsent(name, field); - return previous == null ? field : previous; + return RustField.find(owner, name); } - private static Field optionalInstanceField(Class owner, String name) { - try { - return instanceField(owner, name); - } catch (NoSuchFieldException ignored) { - return null; - } + private static RustField optionalInstanceField(Class owner, String name) { + return RustField.optional(owner, name); } private static FieldAccess fieldAccess(Class owner, String name) @@ -2589,15 +2231,15 @@ private static FieldAccess fieldAccess(Class owner, String name) private static long sliceLogicalLength(Object slice) throws ReflectiveOperationException { - return SLICE_ACCESSES.get(slice.getClass()).rustLength(slice); + return ((SliceView) slice).rustLength; } private static Object sliceBackingForArray(Object slice, Class targetArrayType) throws ReflectiveOperationException { - SliceAccess access = SLICE_ACCESSES.get(slice.getClass()); - Object backing = access.array(slice); - int offset = access.offset(slice); - int length = access.length(slice); + SliceView view = (SliceView) slice; + Object backing = view.array; + int offset = view.offset; + int length = view.length; if (backing instanceof Pointer) { Object result = ((Pointer) backing).backingArrayRange(targetArrayType, offset, length); if (result != null) { @@ -2740,8 +2382,7 @@ private static Object adaptStructuralField( return value; } if (isSliceViewCarrierType(targetType) && value.getClass().isArray()) { - Constructor constructor = SLICE_VIEW_CONSTRUCTORS.get(targetType); - return constructor.newInstance(value, 0, Array.getLength(value)); + return SliceView.create(targetType, value, 0, Array.getLength(value)); } if (targetType.isArray() && isSliceViewCarrierType(value.getClass())) { return sliceBackingForArray(value, targetType); @@ -2761,7 +2402,7 @@ private static Object constructStructuralView( Object source, Class targetClass, Object traitTailCarrier) { try { Object transparentInner = null; - for (Field field : PUBLIC_INSTANCE_FIELDS.get(source.getClass())) { + for (RustField field : PUBLIC_INSTANCE_FIELDS.get(source.getClass())) { if (transparentInner != null) { transparentInner = null; break; @@ -2771,7 +2412,7 @@ private static Object constructStructuralView( if (transparentInner != null && targetClass.isInstance(transparentInner)) { return transparentInner; } - Field[] targetFields = PUBLIC_INSTANCE_FIELDS.get(targetClass); + RustField[] targetFields = PUBLIC_INSTANCE_FIELDS.get(targetClass); ConstructorPlan constructor = structuralConstructor(targetClass, targetFields.length); Object[] args = new Object[constructor.parameterTypes.length]; java.lang.reflect.Parameter[] parameters = constructor.reflection.getParameters(); @@ -2788,7 +2429,7 @@ private static Object constructStructuralView( ? parameters[index].getName() : targetFields[index].getName(); try { - Field sourceField = instanceField(source.getClass(), fieldName); + RustField sourceField = instanceField(source.getClass(), fieldName); args[index] = adaptStructuralField( sourceField.get(source), parameters[index].getType(), index == parameters.length - 1 ? traitTailCarrier : null); @@ -2814,12 +2455,12 @@ private static Object constructStructuralView( private static void copyStructuralFields(Object source, Object target) { try { - Field[] targetFields = PUBLIC_INSTANCE_FIELDS.get(target.getClass()); + RustField[] targetFields = PUBLIC_INSTANCE_FIELDS.get(target.getClass()); if (targetFields.length == 1 && targetFields[0].getType().isInstance(source)) { targetFields[0].set(target, source); return; } - Field[] sourceFields = PUBLIC_INSTANCE_FIELDS.get(source.getClass()); + RustField[] sourceFields = PUBLIC_INSTANCE_FIELDS.get(source.getClass()); if (sourceFields.length == 1) { Object inner = sourceFields[0].get(source); if (inner != null && target.getClass().isInstance(inner)) { @@ -2827,8 +2468,8 @@ private static void copyStructuralFields(Object source, Object target) { return; } } - for (Field targetField : targetFields) { - Field sourceField = instanceField(source.getClass(), targetField.getName()); + for (RustField targetField : targetFields) { + RustField sourceField = instanceField(source.getClass(), targetField.getName()); Object sourceValue = sourceField.get(source); Object targetValue = targetField.get(target); Object adapted; @@ -3270,11 +2911,11 @@ private void attachDecodedTraitTail(Object decoded) { if (decoded == null || carrier == null) { return; } - Field[] fields = PUBLIC_INSTANCE_FIELDS.get(decoded.getClass()); + RustField[] fields = PUBLIC_INSTANCE_FIELDS.get(decoded.getClass()); if (fields.length == 0) { return; } - Field tail = fields[fields.length - 1]; + RustField tail = fields[fields.length - 1]; if (!tail.getType().isInstance(carrier)) { return; } @@ -3747,7 +3388,7 @@ private Object transparentManagedView(Class targetClass) { } Object match = null; try { - for (Field field : PUBLIC_INSTANCE_FIELDS.get(owner.getClass())) { + for (RustField field : PUBLIC_INSTANCE_FIELDS.get(owner.getClass())) { Object candidate = field.get(owner); if (candidate != null && targetClass.isInstance(candidate)) { if (match != null) { @@ -3824,14 +3465,7 @@ private Object structTailSourceObject(String[] descriptor) { int elementSize = structTailPointerElementSize(descriptor); String elementCodec = descriptor[4].isEmpty() ? null : descriptor[4]; Pointer data = byteOffsetRetype(prefixSize, elementSize, elementCodec); - try { - Class viewClass = resolvedRuntimeClass(descriptor[2]); - return LONG_SLICE_VIEW_CONSTRUCTORS - .get(viewClass) - .newInstance(data, 0, metadata()); - } catch (ReflectiveOperationException error) { - throw new IllegalStateException("could not construct Rust struct-tail source view", error); - } + return SliceView.create(descriptor[2], data, 0, metadata()); } private Object structuralSourceObject() { @@ -3869,7 +3503,7 @@ private static long inferredStructuralMetadata(Object value) { if (isSliceViewType(value.getClass())) { return sliceLogicalLength(value); } - for (Field field : PUBLIC_INSTANCE_FIELDS.get(value.getClass())) { + for (RustField field : PUBLIC_INSTANCE_FIELDS.get(value.getClass())) { if (!mayContainStructuralMetadata(field.getType())) { continue; } @@ -3920,7 +3554,7 @@ private static boolean mayContainStructuralMetadata( return false; } boolean result = false; - for (Field field : PUBLIC_INSTANCE_FIELDS.get(type)) { + for (RustField field : PUBLIC_INSTANCE_FIELDS.get(type)) { if (mayContainStructuralMetadata( field.getType(), visiting, encounteredCycle)) { result = true; @@ -3934,14 +3568,14 @@ private static boolean mayContainStructuralMetadata( return result; } - private static final class Cell { - private Object value; + static class Cell { + Object value; private volatile boolean hasStructuralView; private volatile boolean hasMemoryView; private volatile FieldCellCache fields; private volatile boolean hasProjectedViews; - private Cell(Object value) { + Cell(Object value) { this.value = value; } } @@ -4023,14 +3657,6 @@ private Object get() { try { Object owner = owner(); Object value = (Object) access.getter.invokeExact(owner); - if (value instanceof Pointer && access.elementOffsetGetter != null) { - long elementOffset = - (long) access.elementOffsetGetter.invokeExact(owner); - long byteOffset = - (long) access.byteOffsetGetter.invokeExact(owner); - return materializeRelative( - (Pointer) value, elementOffset, byteOffset); - } return value; } catch (RuntimeException | Error error) { throw error; @@ -4043,10 +3669,6 @@ private void set(Object value) { Object owner = owner(); try { access.setter.invokeExact(owner, value); - if (access.elementOffsetSetter != null) { - access.elementOffsetSetter.invokeExact(owner, 0L); - access.byteOffsetSetter.invokeExact(owner, 0L); - } } catch (RuntimeException | Error error) { throw error; } catch (Throwable error) { @@ -4079,126 +3701,32 @@ private void commitOwners() { } private static final class FieldAccess { + private final RustField field; private final MethodHandle getter; private final MethodHandle setter; - private final MethodHandle elementOffsetGetter; - private final MethodHandle byteOffsetGetter; - private final MethodHandle elementOffsetSetter; - private final MethodHandle byteOffsetSetter; private final int fieldNameHash; private final boolean primitive; - // Field metadata is shared; constructing this per projection hashes - // long generated class names on every Rust aggregate field access. + private final boolean borrowed; + // Reuse field metadata to avoid hashing generated class names on each projection. private final String cacheKey; - private FieldAccess(Field field) { + private FieldAccess(RustField field) { + this.field = field; cacheKey = field.getDeclaringClass().getName() + '\n' + field.getName(); fieldNameHash = field.getName().hashCode(); primitive = field.getType().isPrimitive(); + borrowed = field.isBorrowed(); try { - MethodHandles.Lookup lookup = MethodHandles.lookup(); - getter = lookup.unreflectGetter(field).asType( + getter = field.getter().asType( MethodType.methodType(Object.class, Object.class)); - setter = lookup.unreflectSetter(field).asType( + setter = field.setter().asType( MethodType.methodType(void.class, Object.class, Object.class)); - if (field.getType() == Pointer.class) { - Field elementOffset = optionalInstanceField( - field.getDeclaringClass(), - field.getName() + RELATIVE_POINTER_ELEMENT_OFFSET_SUFFIX); - Field byteOffset = optionalInstanceField( - field.getDeclaringClass(), - field.getName() + RELATIVE_POINTER_BYTE_OFFSET_SUFFIX); - if (elementOffset != null && byteOffset != null) { - elementOffsetGetter = lookup.unreflectGetter(elementOffset).asType( - MethodType.methodType(long.class, Object.class)); - byteOffsetGetter = lookup.unreflectGetter(byteOffset).asType( - MethodType.methodType(long.class, Object.class)); - elementOffsetSetter = lookup.unreflectSetter(elementOffset).asType( - MethodType.methodType(void.class, Object.class, long.class)); - byteOffsetSetter = lookup.unreflectSetter(byteOffset).asType( - MethodType.methodType(void.class, Object.class, long.class)); - } else { - elementOffsetGetter = null; - byteOffsetGetter = null; - elementOffsetSetter = null; - byteOffsetSetter = null; - } - } else { - elementOffsetGetter = null; - byteOffsetGetter = null; - elementOffsetSetter = null; - byteOffsetSetter = null; - } } catch (IllegalAccessException error) { throw new IllegalStateException("could not access Rust field pointer", error); } } } - private static final class SliceAccess { - private final MethodHandle array; - private final MethodHandle offset; - private final MethodHandle length; - private final MethodHandle rustLength; - - private SliceAccess(Class type) { - try { - MethodHandles.Lookup lookup = MethodHandles.lookup(); - array = lookup.unreflectGetter(instanceField(type, "array")).asType( - MethodType.methodType(Object.class, Object.class)); - offset = lookup.unreflectGetter(instanceField(type, "offset")).asType( - MethodType.methodType(int.class, Object.class)); - length = lookup.unreflectGetter(instanceField(type, "length")).asType( - MethodType.methodType(int.class, Object.class)); - rustLength = lookup.unreflectGetter(instanceField(type, "rustLength")).asType( - MethodType.methodType(long.class, Object.class)); - } catch (ReflectiveOperationException error) { - throw new IllegalArgumentException( - "invalid Rust slice view " + type.getName(), error); - } - } - - private Object array(Object slice) throws ReflectiveOperationException { - try { - return (Object) array.invokeExact(slice); - } catch (RuntimeException | Error error) { - throw error; - } catch (Throwable error) { - throw new IllegalStateException("could not read Rust slice backing", error); - } - } - - private int offset(Object slice) throws ReflectiveOperationException { - try { - return (int) offset.invokeExact(slice); - } catch (RuntimeException | Error error) { - throw error; - } catch (Throwable error) { - throw new IllegalStateException("could not read Rust slice offset", error); - } - } - - private int length(Object slice) throws ReflectiveOperationException { - try { - return (int) length.invokeExact(slice); - } catch (RuntimeException | Error error) { - throw error; - } catch (Throwable error) { - throw new IllegalStateException("could not read Rust slice storage length", error); - } - } - - private long rustLength(Object slice) throws ReflectiveOperationException { - try { - return (long) rustLength.invokeExact(slice); - } catch (RuntimeException | Error error) { - throw error; - } catch (Throwable error) { - throw new IllegalStateException("could not read Rust slice length", error); - } - } - } - private static final class AllocationInfo { private Long base; private int alignment = 16; @@ -4755,6 +4283,11 @@ public static Pointer cell(Object value, long size, String codecClassName) { return cellAligned(value, checkedArrayLength(size), codecClassName, 16); } + public static Pointer cellAligned( + Object value, int size, String codec, int alignment, String scalarLayout) { + return cellAligned(value, size, codec, alignment); + } + public static Pointer cellAligned( Object value, int size, @@ -4776,6 +4309,71 @@ public static Pointer cellAligned( return metadata < 0 ? pointer : pointer.withMetadata(metadata); } + public static Object storage(Object value, int size, String codec) { + return storageAligned(value, size, codec, 16); + } + + public static Object storage(Object value, long size, String codec) { + return storageAligned(value, checkedArrayLength(size), codec, 16); + } + + /** Compiler-proven, nonzero typed storage, retaining the original identity. */ + public static Object storageAligned(Object value, int size, String codec, int alignment) { + return storageAligned(value, size, codec, alignment, null); + } + + public static Object storageAligned(Object value, int size, String codec, int alignment, String scalarLayout) { + return storageAligned(value, size, codec, alignment, scalarLayout, -1); + } + + public static Object borrowedStorage(Object value, int size, String codec, int kind) { + return borrowedStorageAligned(value, size, codec, 16, kind); + } + + public static Object borrowedStorage(Object value, long size, String codec, int kind) { + return borrowedStorage(value, checkedArrayLength(size), codec, kind); + } + + public static Object borrowedStorageAligned(Object value, int size, String codec, int alignment, int kind) { + return borrowedStorageAligned(value, size, codec, alignment, null, kind); + } + + public static Object borrowedStorageAligned(Object value, int size, String codec, int alignment, String layout, int kind) { + return storageAligned(value, size, codec, alignment, layout, kind); + } + + private static Object storageAligned(Object value, int size, String codec, int alignment, + String scalarLayout, int borrowed) { + if (size <= 0 || alignment <= 0 || (alignment & (alignment - 1)) != 0) { + throw new IllegalArgumentException("invalid typed Rust storage layout"); + } + long metadata = mayCarryStructuralMetadata(value, size) + ? inferredStructuralMetadata(value) : -1; + // Trait and DST addresses need the full codec carrier to preserve dynamic metadata. + Storage storage = borrowed >= 0 && (borrowed != 0 || size == 8) + ? new BorrowedStorage(value, size, codec, metadata, scalarLayout, borrowed != 0) + : new Storage(value, size, codec, metadata, scalarLayout); + recordAlignment(storage, alignment); + return storage; + } + + static Pointer storageBoundary(Storage storage) { + Pointer result = new Pointer(storage, storage.size, 0, storage.size, storage.codec); + return storage.metadata < 0 ? result : result.withMetadata(storage.metadata); + } + + private static boolean directStorage(Storage storage) { + return unescapedStorage(storage) + && (!(storage instanceof BorrowedStorage) || !((BorrowedStorage) storage).split); + } + + private static boolean unescapedStorage(Storage storage) { + Cell cell = storage; + return storage.boundary == null && !cell.hasStructuralView + && !cell.hasMemoryView && cell.fields == null && !cell.hasProjectedViews + && !isStructuralViewCodec(storage.codec); + } + /** Populate published static storage without changing the address captured * by self-references or by another static's initializer. Includes ZST carriers. */ public void initializeStatic(Object value) { @@ -5026,10 +4624,10 @@ private boolean hasStableManagedCarrier() { protected Boolean computeValue(Class type) { Set> visited = new HashSet<>(); while (visited.add(type)) { - Field[] fields = PUBLIC_INSTANCE_FIELDS.get(type); + RustField[] fields = PUBLIC_INSTANCE_FIELDS.get(type); if (fields.length == 2) { - Field bytes = optionalInstanceField(type, "_bytes"); - Field objects = optionalInstanceField(type, "_objects"); + RustField bytes = optionalInstanceField(type, "_bytes"); + RustField objects = optionalInstanceField(type, "_objects"); return bytes != null && bytes.getType() == byte[].class && objects != null && objects.getType() == Object[].class; } @@ -5179,14 +4777,7 @@ public Object projectStructSliceField( } Pointer data = byteOffsetRetype(fieldOffset, elementSize, elementCodecClassName); - try { - Class sliceView = resolvedRuntimeClass(SLICE_VIEW_CLASS_NAME); - return LONG_SLICE_VIEW_CONSTRUCTORS - .get(sliceView) - .newInstance(data, 0, metadata()); - } catch (ReflectiveOperationException error) { - throw new IllegalStateException("could not construct Rust DST slice view", error); - } + return SliceView.create(SLICE_VIEW_CLASS_NAME, data, 0, metadata()); } public Object projectStructStrField( @@ -5203,14 +4794,7 @@ public Object projectStructStrField( } Pointer data = byteOffsetRetype(fieldOffset, 1, null); - try { - Class utf8View = resolvedRuntimeClass(UTF8_VIEW_CLASS_NAME); - return LONG_SLICE_VIEW_CONSTRUCTORS - .get(utf8View) - .newInstance(data, 0, metadata()); - } catch (ReflectiveOperationException error) { - throw new IllegalStateException("could not construct Rust DST string view", error); - } + return SliceView.create(UTF8_VIEW_CLASS_NAME, data, 0, metadata()); } public static Pointer array( @@ -5605,61 +5189,63 @@ public static Pointer fromSlice( if (sliceView instanceof Pointer) { return ((Pointer) sliceView).retype(elementSize, codecClassName); } - try { - SliceAccess access = SLICE_ACCESSES.get(sliceView.getClass()); - Object backing = access.array(sliceView); - int offset = access.offset(sliceView); - long length = access.rustLength(sliceView); - if (backing instanceof Pointer) { - return ((Pointer) backing) - .sliceStorageView(elementSize, codecClassName) - .add(offset) - .withMetadata(length); + SliceView view = (SliceView) sliceView; + return fromSliceParts(view.array, view.offset, view.rustLength, elementSize, codecClassName); + } + + /** Extract an address from SSA view components without constructing a view object. */ + public static Pointer fromSliceParts( + Object backing, int offset, long length, int elementSize, String codecClassName) { + if (backing == null) { + return withoutProvenance(Math.multiplyExact((long) offset, elementSize), + elementSize, codecClassName).withMetadata(length); + } + if (backing instanceof Pointer) { + return ((Pointer) backing) + .sliceStorageView(elementSize, codecClassName) + .add(offset) + .withMetadata(length); + } + MemoryViewOrigin origin = null; + if (mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, backing)) { + Map originStripe = + stateStripe(MEMORY_VIEW_ORIGINS, backing); + synchronized (originStripe) { + origin = originStripe.get(backing); } - MemoryViewOrigin origin = null; - if (mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, backing)) { - Map originStripe = - stateStripe(MEMORY_VIEW_ORIGINS, backing); - synchronized (originStripe) { - origin = originStripe.get(backing); + } + if (origin != null) { + Object allocation = origin.allocation.get(); + if (allocation != null) { + long relativeOffset = Math.multiplyExact((long) offset, elementSize); + Pointer source = new Pointer( + allocation, + origin.allocationElementSize, + origin.byteOffset, + origin.viewSize, + origin.allocationCodecClassName, + origin.allocationCodecClassName, + -1) + .withMetadata(origin.metadata); + if (activeMemoryViewMatches(backing, origin)) { + return source.byteOffsetRetype(relativeOffset, elementSize, codecClassName) + .withMetadata(length); } + return new Pointer( + backing, + elementSize, + relativeOffset, + elementSize, + codecClassName) + .inheritAddressOrigin(source, relativeOffset) + .withMetadata(length); } - if (origin != null) { - Object allocation = origin.allocation.get(); - if (allocation != null) { - long relativeOffset = Math.multiplyExact((long) offset, elementSize); - Pointer source = new Pointer( - allocation, - origin.allocationElementSize, - origin.byteOffset, - origin.viewSize, - origin.allocationCodecClassName, - origin.allocationCodecClassName, - -1) - .withMetadata(origin.metadata); - if (activeMemoryViewMatches(backing, origin)) { - return source.byte_offset(relativeOffset) - .retype(elementSize, codecClassName) - .withMetadata(length); - } - return new Pointer( - backing, - elementSize, - relativeOffset, - elementSize, - codecClassName) - .inheritAddressOrigin(source, relativeOffset) - .withMetadata(length); - } - } - return array( - backing, - offset, - elementSize, - codecClassName).withMetadata(length); - } catch (ReflectiveOperationException error) { - throw new IllegalArgumentException("invalid Rust slice view", error); } + return array( + backing, + offset, + elementSize, + codecClassName).withMetadata(length); } public static Pointer fromSlice(Object sliceView, int elementSize) { @@ -5674,19 +5260,12 @@ public static Pointer fromSlice( if (sliceView instanceof Pointer) { return ((Pointer) sliceView).retype(elementSize, codecClassName); } - try { - SliceAccess access = SLICE_ACCESSES.get(sliceView.getClass()); - Object backing = access.array(sliceView); - if (backing instanceof Pointer) { - int offset = access.offset(sliceView); - long length = access.rustLength(sliceView); - return ((Pointer) backing) - .sliceStorageView(elementSize, codecClassName) - .add(offset) - .withMetadata(length); - } - } catch (ReflectiveOperationException error) { - throw new IllegalArgumentException("invalid Rust slice view", error); + SliceView view = (SliceView) sliceView; + if (view.array instanceof Pointer) { + return ((Pointer) view.array) + .sliceStorageView(elementSize, codecClassName) + .add(view.offset) + .withMetadata(view.rustLength); } return fromSlice(sliceView, checkedArrayLength(elementSize), codecClassName); } @@ -5699,17 +5278,12 @@ public static Pointer fromSlice(Object sliceView) { if (sliceView == null) { return nullPointer(); } - try { - Object array = SLICE_ACCESSES.get(sliceView.getClass()).array(sliceView); - if (array instanceof Pointer) { - Pointer pointer = (Pointer) array; - return fromSlice( - sliceView, pointer.viewSize, pointer.viewCodecClassName); - } - return fromSlice(sliceView, inferredArrayElementSize(array)); - } catch (ReflectiveOperationException error) { - throw new IllegalArgumentException("invalid Rust slice view", error); + Object array = ((SliceView) sliceView).array; + if (array instanceof Pointer) { + Pointer pointer = (Pointer) array; + return fromSlice(sliceView, pointer.viewSize, pointer.viewCodecClassName); } + return fromSlice(sliceView, inferredArrayElementSize(array)); } public static Pointer nullPointer(int viewSize) { @@ -6063,15 +5637,7 @@ public static Object stringView(String value, String viewClassName) { if (cached != null) { return cached; } - try { - Class viewClass = resolvedRuntimeClass(viewClassName); - cached = SLICE_VIEW_CONSTRUCTORS.get(viewClass) - .newInstance(views.bytes, 0, views.bytes.length); - } catch (ReflectiveOperationException error) { - throw new IllegalStateException( - "could not construct cached Rust string view " + viewClassName, - error); - } + cached = SliceView.create(viewClassName, views.bytes, 0, views.bytes.length); if (utf8) { views.utf8 = cached; } else { @@ -6459,6 +6025,7 @@ public static Pointer retypeStructTailWithMetadataOf( String newViewCodecClassName) { Pointer result = attachStructTailTraitMetadata( pointer.retype(newViewSize, newViewCodecClassName), metadataSource); + result.traitObjectCarrier(null); if (result.traitAdapterClassName() == null) { return result; } @@ -6480,10 +6047,8 @@ public static Pointer retypeStructTailFromTraitPointer( Pointer pointer, long newViewSize, String newViewCodecClassName) { - Pointer result = pointer.retype(newViewSize, newViewCodecClassName); - result.traitObjectCarrier(null); - result.traitMetadataCarrier(null); - return attachStructTailTraitMetadata(result, pointer); + return retypeStructTailWithMetadataOf( + pointer, pointer, newViewSize, newViewCodecClassName); } /** Creates a coherent JVM carrier view for a Rust struct-tail unsizing coercion. */ @@ -6518,7 +6083,7 @@ public static Pointer unsizeStruct( try { Object source = pointer.getObject(); Class targetClass = resolvedRuntimeClass(targetClassName); - for (Field targetField : PUBLIC_INSTANCE_FIELDS.get(targetClass)) { + for (RustField targetField : PUBLIC_INSTANCE_FIELDS.get(targetClass)) { if (!isSliceViewCarrierType(targetField.getType())) { continue; } @@ -6618,14 +6183,7 @@ public static Object restoreErasedSliceView(Pointer pointer, String viewClassNam throw new IllegalArgumentException("invalid erased Rust slice view"); } Pointer data = restoreErasedView(pointer); - try { - Class viewClass = resolvedRuntimeClass(viewClassName); - return LONG_SLICE_VIEW_CONSTRUCTORS - .get(viewClass) - .newInstance(data, 0, pointer.metadata()); - } catch (ReflectiveOperationException error) { - throw new IllegalStateException("could not rebuild erased Rust slice view", error); - } + return SliceView.create(viewClassName, data, 0, pointer.metadata()); } private static Pointer traitMetadataPointer(Object metadata, int depth) @@ -6636,7 +6194,7 @@ private static Pointer traitMetadataPointer(Object metadata, int depth) if (metadata == null || depth == 0) { return null; } - for (Field field : PUBLIC_INSTANCE_FIELDS.get(metadata.getClass())) { + for (RustField field : PUBLIC_INSTANCE_FIELDS.get(metadata.getClass())) { Object nested = field.get(metadata); Pointer marker = traitMetadataPointer(nested, depth - 1); if (marker != null) { @@ -6688,6 +6246,15 @@ public static Pointer fromRawTraitParts(Pointer data, Object metadata) { } } + /** Rebuilds a trait-tailed struct using the supplied vtable, not stale data metadata. */ + public static Pointer fromRawStructTraitParts( + Pointer data, long prefixSize, String prefixCodec, Object metadata) { + Pointer tail = fromRawTraitParts(data.byte_offset(prefixSize), metadata); + Pointer result = data.retype(prefixSize, prefixCodec); + result.traitObjectCarrier(null); + return attachStructTailTraitMetadata(result, tail); + } + private Pointer withMetadata(long metadata) { this.metadata = metadata; return this; @@ -7321,45 +6888,6 @@ private boolean samePointerMetadataWords(Pointer other) { return markerAdapter != null && markerAdapter.equals(unmaterializedAdapter); } - /** Compares deferred pointer data and metadata words without materialization. */ - public static boolean samePointerRelative( - Pointer left, - long leftElementOffset, - long leftByteOffset, - Pointer right, - long rightElementOffset, - long rightByteOffset) { - if (left == null || right == null) { - return false; - } - if (left.traitObjectCarrier() != null - || left.traitAdapterClassName() != null - || left.traitMetadataMarker() != null - || right.traitObjectCarrier() != null - || right.traitAdapterClassName() != null - || right.traitMetadataMarker() != null) { - return materializeRelative(left, leftElementOffset, leftByteOffset) - .samePointer(materializeRelative( - right, rightElementOffset, rightByteOffset)); - } - long leftDisplacement = Math.addExact( - Math.multiplyExact(leftElementOffset, left.viewSize), leftByteOffset); - long rightDisplacement = Math.addExact( - Math.multiplyExact(rightElementOffset, right.viewSize), rightByteOffset); - boolean sameAddress; - if (left.allocation != null && left.allocation == right.allocation) { - sameAddress = Math.addExact(left.byteOffset, leftDisplacement) - == Math.addExact(right.byteOffset, rightDisplacement); - } else if (left.allocation == null && right.allocation == null) { - sameAddress = Math.addExact(left.exposedAddress, leftDisplacement) - == Math.addExact(right.exposedAddress, rightDisplacement); - } else { - sameAddress = Math.addExact(left.numericAddress(), leftDisplacement) - == Math.addExact(right.numericAddress(), rightDisplacement); - } - return sameAddress && left.samePointerMetadataWords(right); - } - /** Compares both words of a Rust slice/str fat pointer. */ public static boolean fatPointerEquals(Object left, Object right) { if (left == right) { @@ -7483,6 +7011,26 @@ public boolean greaterOrEqual(Pointer other) { return compareAddress(other) >= 0; } + /** Test a reference or NonNull niche without exposing an address or creating a boundary. */ + public static long nullableLocationTag(Object root, long offset) { + if (root == null) return offset == 0 ? 0 : 1; + if (!(root instanceof Pointer)) return 1; + Pointer pointer = (Pointer) root; + if (pointer.allocation != null) return 1; + long address = pointer.addressOrigin() == null + ? pointer.exposedAddress : pointer.numericAddress(); + return address + offset == 0 ? 0 : 1; + } + + public static long nullableTag(Pointer pointer) { + return nullableLocationTag(pointer, 0); + } + + /** Test the data address for null. An empty array or string still has a non-null root. */ + public static long nullableViewLocationTag(Object root, int start) { + return start == 0 ? nullableLocationTag(root, 0) : 1; + } + public static boolean is_null(Pointer pointer) { return pointer == null || pointer.numericAddress() == 0; } @@ -7569,9 +7117,10 @@ private static void discardCollectedTypedExposedTargets(int limit) { } } - public static Object asRefOption(Pointer pointer, String optionClassName) { - String variantName = optionClassName + (is_null(pointer) ? "$None" : "$Some"); + public static Object asRefOption(Pointer pointer, String someClassName, String noneClassName) { + String variantName = is_null(pointer) ? noneClassName : someClassName; try { + // Shared enum interfaces need not own the payload class. Use the exact compiler-supplied variant. Class variant = resolvedRuntimeClass(variantName); if (is_null(pointer)) { return constructorWithArity(variant, 0).newInstance(); @@ -7590,7 +7139,7 @@ public static Object asRefOption(Pointer pointer, String optionClassName) { return constructor.newInstance(referent); } catch (ReflectiveOperationException error) { throw new IllegalStateException( - "could not construct Rust pointer option " + optionClassName, error); + "could not construct Rust pointer option " + variantName, error); } } @@ -7837,62 +7386,6 @@ public static long address(Pointer pointer) { return pointer == null ? 0L : pointer.address(); } - public static long addressRelative( - Pointer base, long elementOffset, long byteOffset) { - if (base == null) { - return 0L; - } - long displacement = Math.addExact( - Math.multiplyExact(elementOffset, base.viewSize), byteOffset); - if (displacement == 0) { - return base.address(); - } - Pointer origin = base.addressOrigin(); - if (origin != null) { - return Math.addExact( - origin.address(), - Math.addExact(base.addressOriginOffset(), displacement)); - } - if (base.allocation == null) { - return Math.addExact(base.exposedAddress, displacement); - } - synchronized (ALLOCATIONS) { - AllocationInfo info = allocationInfo(base.allocation); - long address = Math.addExact( - base.publishAllocationRange(info), - Math.addExact(base.byteOffset, displacement)); - registerExposedTarget(address, base.exposedTargetAt(displacement)); - return address; - } - } - - public static long addrRelative( - Pointer base, long elementOffset, long byteOffset) { - if (base == null) { - return 0L; - } - long displacement = Math.addExact( - Math.multiplyExact(elementOffset, base.viewSize), byteOffset); - return Math.addExact(base.numericAddress(), displacement); - } - - public static boolean isNullRelative( - Pointer base, long elementOffset, long byteOffset) { - return base == null || addrRelative(base, elementOffset, byteOffset) == 0; - } - - public static boolean isAlignedRelative( - Pointer base, - long elementOffset, - long byteOffset, - long alignment) { - if (alignment <= 0 || (alignment & (alignment - 1)) != 0) { - throw new IllegalArgumentException( - "is_aligned_to: align is not a power-of-two"); - } - return (addrRelative(base, elementOffset, byteOffset) & (alignment - 1)) == 0; - } - /** * Publishes an opaque token for an erased trait-object data pointer. * Native fat pointers carry concrete dispatch identity in their metadata; @@ -8106,7 +7599,7 @@ private static int primitiveArrayElementSize(Object value, int aggregateSize) { private static int loadPrimitiveArrayByte(Object array, int byteOffset, int elementSize) { int elementIndex = byteOffset / elementSize; int withinElement = byteOffset % elementSize; - long bits = valueBits(arrayGet(array, elementIndex), elementSize); + long bits = primitiveArrayBits(array, elementIndex); return (int) ((bits >>> (withinElement * 8)) & 0xffL); } @@ -8114,11 +7607,10 @@ private static void storePrimitiveArrayByte( Object array, int byteOffset, int elementSize, int incoming) { int elementIndex = byteOffset / elementSize; int withinElement = byteOffset % elementSize; - Object current = arrayGet(array, elementIndex); - long bits = valueBits(current, elementSize); + long bits = primitiveArrayBits(array, elementIndex); long mask = 0xffL << (withinElement * 8); long updated = (bits & ~mask) | (((long) incoming & 0xffL) << (withinElement * 8)); - arraySet(array, elementIndex, carrierFromBits(current, updated, elementSize)); + storePrimitiveArrayBits(array, elementIndex, updated); } private static long loadArrayCodecBits( @@ -9367,125 +8859,510 @@ public static void atomicFence(int ordering) { } } - private void requireScalarViewSize(int expectedSize, String scalarType) { - if (viewSize == expectedSize) { + private static Object normalizeLocationOrigin(Object root) { + if (root == null || root instanceof Pointer || root instanceof Storage + || !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, root)) return root; + Map stripe = stateStripe(MEMORY_VIEW_ORIGINS, root); + synchronized (stripe) { + MemoryViewOrigin origin = stripe.get(root); + // Filter matches can be stale or false. Only a live origin needs normalization. + if (origin == null || origin.allocation.get() == null) return root; + } + return array(root, 0, inferredArrayElementSize(root)); + } + + /** Pointer differences use allocation identity, never exposed address lookup. */ + public static long locationStride(Object root) { + return root instanceof Storage ? ((Storage) root).size : ((Pointer) root).viewSize; + } + + /** The root carries the complete layout and provenance of a non-scalar view. */ + public static Pointer fromStorageLocation(Object root, long offset) { + if (root instanceof Storage) root = ((Storage) root).boundary(); + // Retain the pointer that binds a decoded view. A later commit must use the same owner. + if (offset == 0) return (Pointer) root; + return ((Pointer) root).byte_offset(offset); + } + + /** The physical ABI records a reconstruction plan, never an inferred JVM layout. */ + public static Pointer addressFromParts(Object root, long offset, int plan) { + if (plan == 64 || plan == 128) return fromBorrowedStorageLocation(root, offset, plan == 64); + return plan == 0 ? fromStorageLocation(root, offset) : fromLocation(root, offset, plan); + } + + private static BorrowedFieldPath borrowedPath(Storage storage, long offset, boolean view) { + StorageLayout layout = storage.layout(((Cell) storage).value); + return layout == null ? null : layout.borrowedAt(offset, view); + } + + private static Pointer fromBorrowedStorageLocation(Object root, long offset, boolean view) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + BorrowedFieldPath path = borrowedPath(storage, offset, view); + if (path != null) { + // Delay path materialization until a boundary needs it. Preserve replaceable parents after escape. + Pointer base = storage.boundary(); + Object owner = storage; + Pointer result = null; + for (int i = 0; i < path.fields.length; i++) { + RustField field = path.fields[i]; + boolean last = i == path.fields.length - 1; + result = rootField(owner, field.getDeclaringClass(), field.getName(), + last ? path.size : 0, last ? path.codec : null); + owner = result.allocation; + } + return result.inheritAddressOrigin(base, offset); + } + } + return fromStorageLocation(root, offset); + } + + /** The address of a stored borrow keeps its enclosing allocation root. */ + public static Object storageBorrowedFieldRoot(Object root, long offset, String owner, + String field, long fieldOffset, long size, String codec) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (directStorage(storage) && isGeneratedAggregateCodec(storage.codec)) { + long absolute = Math.addExact(offset, fieldOffset); + StorageLayout layout = storage.layout(((Cell) storage).value); + if (layout != null) { + BorrowedFieldPath path = layout.borrowedAt(absolute, true); + if (path == null) path = layout.borrowedAt(absolute, false); + if (path != null && path.size == size && path.field().getName().equals(field) + && matchesBinaryClassName(owner, path.field().getDeclaringClass().getName()) + && java.util.Objects.equals(codec, path.codec)) return storage; + } + } + } + return fromStorageLocation(root, offset).projectStructField(owner, field, fieldOffset, size, codec); + } + + /** Register the containing allocation when an escaping array borrow exposes its JVM array. */ + public static Object loadStorageArray(Object root, long offset, String target) { + return fromStorageLocation(root, offset).getObject(); + } + + public static Object loadStorageLocation(Object root, long offset, String target) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (offset == 0 && directStorage(storage)) return ((Cell) storage).value; + root = storage.boundary(); + } + Pointer base = (Pointer) root; + if (base.rareState == null && base.addressState == null) { + return loadObjectLocation(base, offset, target); + } + Pointer view = fromStorageLocation(root, offset); + return target == null ? view.getObject() : view.getObjectAs(target); + } + + /** A stored borrow returns its components in the ordinary caller-owned ABI. */ + public static Object loadBorrowedView(Object root, long offset, long[] metadata) { + if (root instanceof BorrowedStorage && offset == 0 && unescapedStorage((Storage) root) + && ((BorrowedStorage) root).view) { + return ((BorrowedStorage) root).read(metadata, true); + } + if (root instanceof Storage) { + Storage storage = (Storage) root; + BorrowedFieldPath path = borrowedPath(storage, offset, true); + if (path != null && directStorage(storage)) return path.read(((Cell) storage).value, metadata); + root = fromBorrowedStorageLocation(root, offset, true); + offset = 0; + } + FieldCell field = directBorrowedField(root, offset, true); + if (field != null) return readBorrowedField(field, metadata); + SliceView view = (SliceView) loadStorageLocation(root, offset, SLICE_VIEW_CLASS_NAME); + metadata[0] = RustField.viewStart(view); + metadata[1] = RustField.viewLength(view); + return RustField.viewRoot(view); + } + + public static Object loadBorrowedAddress(Object root, long offset, long[] metadata) { + if (root instanceof BorrowedStorage && offset == 0 && unescapedStorage((Storage) root) + && !((BorrowedStorage) root).view) { + return ((BorrowedStorage) root).read(metadata, false); + } + if (root instanceof Storage) { + Storage storage = (Storage) root; + BorrowedFieldPath path = borrowedPath(storage, offset, false); + if (path != null && directStorage(storage)) return path.read(((Cell) storage).value, metadata); + root = fromBorrowedStorageLocation(root, offset, false); + offset = 0; + } + FieldCell field = directBorrowedField(root, offset, false); + if (field != null) return readBorrowedField(field, metadata); + Object value = loadStorageLocation(root, offset, null); + metadata[0] = 0; + return value; + } + + public static void storeBorrowedView(Object root, long offset, Object backing, int start, long length) { + storeBorrowedView(root, offset, backing, start, length, false); + } + + public static void storeBorrowedUtf8(Object root, long offset, Object backing, int start, long length) { + storeBorrowedView(root, offset, backing, start, length, true); + } + + private static void storeBorrowedView(Object root, long offset, Object backing, int start, long length, boolean utf8) { + if (root instanceof BorrowedStorage && offset == 0 && unescapedStorage((Storage) root) + && ((BorrowedStorage) root).view) { + ((BorrowedStorage) root).store(backing, start, length, utf8 ? -2 : -1); return; } - String sourceView = zeroSizedSourceViewSize() >= 0 - ? "; recorded erased source view is " + zeroSizedSourceViewSize() + " bytes" - : ""; - throw new IllegalStateException( - scalarType + " load requires a " + expectedSize + "-byte view, but pointer has a " - + viewSize + "-byte view" + sourceView); + if (root instanceof Storage) { + Storage storage = (Storage) root; + BorrowedFieldPath path = borrowedPath(storage, offset, true); + if (path != null && directStorage(storage)) { + path.write(((Cell) storage).value, backing, start, length); + return; + } + root = fromBorrowedStorageLocation(root, offset, true); + offset = 0; + } + FieldCell field = directBorrowedField(root, offset, true); + if (field != null) { + writeBorrowedField(field, backing, start, length); + return; + } + storeStorageLocation(root, offset, utf8 ? new Utf8View(backing, start, length) + : new SliceView(backing, start, length)); } - private long relativeByteOffset(long elementOffset, long additionalByteOffset) { - return Math.addExact( - byteOffset, - Math.addExact( - Math.multiplyExact(elementOffset, viewSize), - additionalByteOffset)); + public static void storeBorrowedAddress(Object root, long offset, Object backing, long displacement, int size) { + if (root instanceof BorrowedStorage && offset == 0 && unescapedStorage((Storage) root) + && !((BorrowedStorage) root).view && !(backing instanceof Storage)) { + ((BorrowedStorage) root).store(backing, displacement, 0, size); + return; + } + // Resolve referenced boundaries before assignment. Lazy materialization under owner locks could recurse through cycles. + if (root instanceof Storage) { + Storage storage = (Storage) root; + BorrowedFieldPath path = borrowedPath(storage, offset, false); + if (path != null && directStorage(storage)) { + path.write(((Cell) storage).value, backing, displacement, 0); + return; + } + root = fromBorrowedStorageLocation(root, offset, false); + offset = 0; + } + FieldCell field = directBorrowedField(root, offset, false); + if (field != null) { + writeBorrowedField(field, backing, displacement, 0); + return; + } + Pointer value = backing == null && displacement == 0 ? null + : addressFromParts(backing, displacement, size); + storeStorageLocation(root, offset, value); + } + + private static void writeBorrowedField(FieldCell cell, Object root, long offset, long length) { + try { cell.access.field.setBorrowedParts(cell.owner(), root, offset, length); } + catch (IllegalAccessException error) { + throw new IllegalStateException("could not write borrowed Rust field", error); + } + discardProjectedFieldViews(cell); + } + + private static Object readBorrowedField(FieldCell cell, long[] metadata) { + try { return cell.access.field.borrowedParts(cell.owner(), metadata); } + catch (IllegalAccessException error) { + throw new IllegalStateException("could not read borrowed Rust field", error); + } + } + + private static FieldCell directBorrowedField(Object root, long offset, boolean view) { + if (!(root instanceof Pointer)) return null; + Pointer pointer = (Pointer) root; + if (offset != 0 || pointer.byteOffset != 0 + || pointer.viewSize != pointer.allocationElementSize + || !pointer.isDirectAllocationView() + || pointer.rareState != null || !(pointer.allocation instanceof FieldCell)) return null; + FieldCell field = (FieldCell) pointer.allocation; + if (!field.access.field.borrowedShape(view)) return null; + Object owner = field; + for (int depth = 0; depth < 16; depth++) { + if (owner instanceof FieldCell) { + FieldCell current = (FieldCell) owner; + if (current.hasMemoryView || current.hasStructuralView || current.hasProjectedViews + || current.hasMemoryOrigins) return null; + owner = current.rootOwner == null ? current.fixedOwner : current.rootOwner; + } else if (owner instanceof Cell) { + Cell current = (Cell) owner; + if (current.hasMemoryView || current.hasStructuralView || current.hasProjectedViews) return null; + owner = current.value; + } else { + return mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, owner) ? null : field; + } + } + return null; } - /** - * Materializes a compiler-deferred derived pointer at an ABI or semantic boundary. - * A zero displacement preserves object identity and allocates nothing. - */ - public static Pointer materializeRelative( - Pointer base, long elementOffset, long byteOffset) { - if (elementOffset == 0 && byteOffset == 0) { - return base; + /** Copying a value never establishes a mutable decoded-view binding. */ + private boolean storeAggregateRange(long offset, long byteSize, String codec, Object value) { + if (!(allocation instanceof byte[]) || rareState != null || addressState != null + || byteSize <= 0 || !isGeneratedAggregateCodec(codec) + || mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, value)) { + return false; } - long displacement = Math.addExact( - Math.multiplyExact(elementOffset, base.viewSize), byteOffset); - return base.byte_offset(displacement); + CodecCalls.RangeEncoder encode = codecPlan(codec).encodeAt(); + if (encode == null) return false; + byte[] bytes = (byte[]) allocation; + int size = checkedArrayLength(byteSize); + int start = Math.toIntExact(Math.addExact(byteOffset, offset)); + if (start < 0 || start > bytes.length - size) { + throw new IndexOutOfBoundsException("aggregate store exceeds byte-addressable Rust storage"); + } + // Save pending writes outside this range before replacement. Live views use set() to preserve their origin. + prepareMemoryWrite(start, size); + discardEncodedPointers(bytes, start, size); + encode.encode(value, bytes, start); + return true; } - /** Retypes a deferred derived pointer with a single allocation. */ - public static Pointer retypeRelative( - Pointer base, - long elementOffset, - long byteOffset, - long newViewSize, - String newViewCodecClassName) { - long displacement = Math.addExact( - Math.multiplyExact(elementOffset, base.viewSize), byteOffset); - if (displacement == 0) { - return base.retype(newViewSize, newViewCodecClassName); + public static Object storageFieldRoot(Object root, long offset, String owner, + String field, long fieldOffset, long size, String codec) { + // Scalar components supply their own width. Byte storage needs no layout carrier. + if (codec == null && size > 0 && size <= 8 && hasByteStorage(root) + && (!(root instanceof Pointer) || ((Pointer) root).traitMetadataCarrier() == null)) { + return root; } - return base.byteOffsetRetype( - displacement, newViewSize, newViewCodecClassName); + if (root instanceof Storage && size > 0 && size <= 8) { + Storage storage = (Storage) root; + // Prepare byte storage before a later byte alias uses this root and offset. + if (directStorage(storage) && isGeneratedAggregateCodec(storage.codec)) { + StorageLayout layout = storage.layout(((Cell) storage).value); + if (layout != null && layout.at(Math.addExact(offset, fieldOffset), (int) size) != null) { + return storage; + } + } + } + return fromStorageLocation(root, offset).projectStructField(owner, field, fieldOffset, size, codec); } - public static boolean getBooleanRelative( - Pointer base, long elementOffset, long byteOffset) { - base.requireScalarViewSize(1, "bool"); - return base.loadUnsignedAt( - base.relativeByteOffset(elementOffset, byteOffset), 1) - != 0; + public static long storageFieldOffset(Object root, Object base, long offset) { + return root == base ? offset : 0; } - public static byte getI8Relative( - Pointer base, long elementOffset, long byteOffset) { - base.requireScalarViewSize(1, "i8/u8"); - return (byte) base.loadUnsignedAt( - base.relativeByteOffset(elementOffset, byteOffset), 1); + public static void storeStorageLocation(Object root, long offset, Object value) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (offset == 0 && directStorage(storage)) { + Cell cell = storage; + cell.value = convertDirectValue(cell.value, value, storage.size); + return; + } + root = storage.boundary(); + } + Pointer pointer = (Pointer) root; + if (pointer.storeAggregateRange(offset, pointer.viewSize, + pointer.viewCodecClassName, value)) { + return; + } + if (offset == 0) pointer.set(value); + else pointer.byte_offset(offset).set(value); + } + + public static void commitStorageLocation(Object root, long offset) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (offset == 0 && directStorage(storage)) { + discardProjectedFieldViews(((Cell) storage).value); + return; + } + root = storage.boundary(); + } + Pointer pointer = (Pointer) root; + if (offset == 0) pointer.commitMemoryView(); + else pointer.byte_offset(offset).commitMemoryView(); } - public static short getI16Relative( - Pointer base, long elementOffset, long byteOffset) { - base.requireScalarViewSize(2, "i16/u16/f16"); - return (short) base.loadUnsignedAt( - base.relativeByteOffset(elementOffset, byteOffset), 2); + /** Materialize a scalar address only where a JVM object is required. */ + public static Pointer fromLocation(Object root, long offset, int size) { + return fromTypedStorageLocation(root, offset, size, null); } - public static int getI32Relative( - Pointer base, long elementOffset, long byteOffset) { - base.requireScalarViewSize(4, "i32/u32/char"); - return (int) base.loadUnsignedAt( - base.relativeByteOffset(elementOffset, byteOffset), 4); + /** Materialize an exact aggregate view only at a carrier boundary. */ + public static Pointer fromTypedStorageLocation(Object root, long offset, int size, String codec) { + if (root instanceof Storage) root = ((Storage) root).boundary(); + if (root == null && offset == 0) return null; + if (root == null) return fromUnprovenancedAddress(offset, size, codec); + if (!(root instanceof Pointer) && root.getClass().isArray() + && !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, root)) { + return new Pointer(root, inferredArrayElementSize(root), offset, size, null, codec, -1); + } + Pointer pointer = root instanceof Pointer ? (Pointer) root + : array(root, 0, inferredArrayElementSize(root)); + Pointer escaped = pointer.escapedFieldStorage(offset); + if (escaped != null) return escaped.retype(size, codec); + // Combine offset and retype without an intermediate Pointer. + // Keep the result independent because later metadata changes must not affect the source. + Pointer result = new Pointer( + pointer.allocation, + pointer.allocationElementSize, + pointer.byteOffset + offset, + size, + pointer.allocationCodecClassName, + codec, + pointer.allocation == null ? pointer.exposedAddress + offset : -1) + .withMetadata(pointer.metadata) + .copyDynamicMetadata(pointer) + .copyAddressOrigin(pointer, offset); + if (offset == 0 && pointer.zeroSizedSourceViewSize() >= 0) { + result.setZeroSizedSourceView(pointer.zeroSizedSourceViewSize(), + pointer.zeroSizedSourceViewCodecClassName()); + } else if (pointer.viewSize == 0 && pointer.viewCodecClassName != null) { + result.setZeroSizedSourceView(pointer.viewSize, pointer.viewCodecClassName); + } + return result; } - public static long getI64Relative( - Pointer base, long elementOffset, long byteOffset) { - base.requireScalarViewSize(8, "i64/u64"); - return base.loadUnsignedAt( - base.relativeByteOffset(elementOffset, byteOffset), 8); + public static boolean hasByteStorage(Object root) { + return root instanceof byte[] + || (root instanceof Pointer && ((Pointer) root).allocation instanceof byte[]); } - public static float getF32Relative( - Pointer base, long elementOffset, long byteOffset) { - base.requireScalarViewSize(4, "f32"); - return Float.intBitsToFloat((int) base.loadUnsignedAt( - base.relativeByteOffset(elementOffset, byteOffset), 4)); + /** Scalar field fallback after generated code has tried the managed carrier. */ + public static long loadScalarField(Object root, long offset, String owner, + String field, long fieldOffset, int size) { + if (hasByteStorage(root)) { + return loadLocationBits(root, Math.addExact(offset, fieldOffset), size); + } + Pointer projected = fromStorageLocation(root, offset) + .projectStructField(owner, field, fieldOffset, size, null); + return loadLocationBits(projected, 0, size); } - public static double getF64Relative( - Pointer base, long elementOffset, long byteOffset) { - base.requireScalarViewSize(8, "f64"); - return Double.longBitsToDouble(base.loadUnsignedAt( - base.relativeByteOffset(elementOffset, byteOffset), 8)); + public static void storeScalarField(Object root, long offset, String owner, + String field, long fieldOffset, long bits, int size) { + if (hasByteStorage(root)) { + storeLocationBits(root, Math.addExact(offset, fieldOffset), bits, size); + } else { + Pointer projected = fromStorageLocation(root, offset) + .projectStructField(owner, field, fieldOffset, size, null); + storeLocationBits(projected, 0, bits, size); + } + } + + public static long loadLocationBits(Object root, long offset, int size) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + Object value = ((Cell) storage).value; + if (directStorage(storage)) { + StorageLayout layout = storage.layout(value); + StorageLayout.Leaf leaf = layout == null ? null : layout.at(offset, size); + if (leaf != null) return leaf.read(value, offset, size); + } + root = storage.boundary(); + } + if (root instanceof Pointer) { + Pointer pointer = (Pointer) root; + if (pointer.hasDirectPrimitiveArrayStorage()) { + return loadLocationBits(pointer.allocation, Math.addExact(pointer.byteOffset, offset), size); + } + Pointer escaped = pointer.escapedFieldStorage(offset); + return escaped != null ? escaped.loadUnsigned(size) + : pointer.loadUnsignedAt(Math.addExact(pointer.byteOffset, offset), size); + } + int elementSize = inferredArrayElementSize(root); + if (mayBeInIdentityFilter(MEMORY_VIEW_FILTER, root)) { + // Keep the array used by direct JVM access. + // array() can redirect a decoded view to its byte origin and detach later array reads. + return new Pointer(root, elementSize, 0, elementSize, null).loadUnsignedAt(offset, size); + } + if (root instanceof byte[]) return MemoryBytes.read((byte[]) root, Math.toIntExact(offset), size); + int index = Math.toIntExact(offset / elementSize); + int within = (int) (offset % elementSize); + if (within == 0 && size == elementSize) return primitiveArrayBits(root, index); + long bits = 0; + for (int i = 0; i < size; i++) { + bits |= (long) loadPrimitiveArrayByte(root, Math.toIntExact(offset + i), elementSize) << (8 * i); + } + return bits; + } + + public static void storeLocationBits(Object root, long offset, long bits, int size) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + Object value = ((Cell) storage).value; + if (directStorage(storage)) { + StorageLayout layout = storage.layout(value); + StorageLayout.Leaf leaf = layout == null ? null : layout.at(offset, size); + if (leaf != null) { + leaf.write(value, offset, bits, size); + return; + } + } + root = storage.boundary(); + } + if (root instanceof Pointer) { + Pointer pointer = (Pointer) root; + if (pointer.hasDirectPrimitiveArrayStorage()) { + storeLocationBits(pointer.allocation, Math.addExact(pointer.byteOffset, offset), bits, size); + return; + } + Pointer escaped = pointer.escapedFieldStorage(offset); + if (escaped != null) escaped.storeBytes(bits, size); + else pointer.storeBytesAt(Math.addExact(pointer.byteOffset, offset), bits, size); + return; + } + int elementSize = inferredArrayElementSize(root); + if (hasScalarWriteTracking(root)) { + new Pointer(root, elementSize, 0, elementSize, null).storeBytesAt(offset, bits, size); + return; + } + if (root instanceof byte[]) { + MemoryBytes.write((byte[]) root, Math.toIntExact(offset), size, bits); + return; + } + int index = Math.toIntExact(offset / elementSize); + int within = (int) (offset % elementSize); + if (within == 0 && size == elementSize) { + storePrimitiveArrayBits(root, index, bits); + return; + } + for (int i = 0; i < size; i++) { + storePrimitiveArrayByte(root, Math.toIntExact(offset + i), elementSize, (int) (bits >>> (8 * i)) & 255); + } + } + + private boolean hasDirectPrimitiveArrayStorage() { + return allocation != null && rareState == null && addressState == null + && allocation.getClass().isArray() + && allocation.getClass().getComponentType().isPrimitive() + && allocationElementSize == inferredArrayElementSize(allocation) + && !hasScalarWriteTracking(allocation); + } + + private void requireScalarViewSize(int expectedSize, String scalarType) { + if (viewSize == expectedSize) { + return; + } + String sourceView = zeroSizedSourceViewSize() >= 0 + ? "; recorded erased source view is " + zeroSizedSourceViewSize() + " bytes" + : ""; + throw new IllegalStateException( + scalarType + " load requires a " + expectedSize + "-byte view, but pointer has a " + + viewSize + "-byte view" + sourceView); } /** * Loads a reference-valued pointee without allocating its derived pointer * when the JVM allocation already stores the requested Rust value directly. */ - public static Object getObjectRelative( + private static Object loadObjectLocation( Pointer base, - long elementOffset, long byteOffset, String targetClassName) { if (base == null) { throw new NullPointerException("attempted to dereference a null Rust pointer"); } long absoluteByteOffset = - base.relativeByteOffset(elementOffset, byteOffset); + Math.addExact(base.byteOffset, byteOffset); boolean sameCodec = base.allocationCodecClassName == null ? base.viewCodecClassName == null : base.allocationCodecClassName.equals(base.viewCodecClassName); - if (elementOffset == 0 - && byteOffset == 0 + if (byteOffset == 0 && base.allocation instanceof Cell && base.byteOffset == 0 && base.viewSize == base.allocationElementSize @@ -9544,7 +9421,7 @@ && matchesBinaryClassName( error); } } - Pointer pointer = materializeRelative(base, elementOffset, byteOffset); + Pointer pointer = fromStorageLocation(base, byteOffset); return targetClassName == null || targetClassName.isEmpty() ? pointer.getObject() : pointer.getObjectAs(targetClassName); @@ -10167,6 +10044,11 @@ public void set(Object value) { if (isDirectAllocationView()) { clearStructuralViewState(); int elementIndex = Math.toIntExact(byteOffset / allocationElementSize); + if (allocation instanceof FieldCell && ((FieldCell) allocation).access.borrowed) { + // Do not read the old borrowed value. Partially initialized storage may not contain a valid carrier. + writeElement(elementIndex, value); + return; + } Object current = readElement(elementIndex); writeElement( elementIndex, @@ -10416,100 +10298,6 @@ public static void copyElements( copy(source, destination, checkedElementByteCount(source, elementCount)); } - private static void copyRelativeBytes( - Pointer source, - long sourceElementOffset, - long sourceByteOffset, - Pointer destination, - long destinationElementOffset, - long destinationByteOffset, - int byteCount, - boolean nonOverlapping) { - long absoluteSourceByteOffset = - source.relativeByteOffset(sourceElementOffset, sourceByteOffset); - long absoluteDestinationByteOffset = - destination.relativeByteOffset(destinationElementOffset, destinationByteOffset); - if (nonOverlapping && source.allocation == destination.allocation) { - long sourceEnd = Math.addExact(absoluteSourceByteOffset, byteCount); - long destinationEnd = Math.addExact(absoluteDestinationByteOffset, byteCount); - if (absoluteSourceByteOffset < destinationEnd - && absoluteDestinationByteOffset < sourceEnd) { - throw new IllegalArgumentException("copy_nonoverlapping regions overlap"); - } - } - if (tryCopyScalarRange( - source, - absoluteSourceByteOffset, - destination, - absoluteDestinationByteOffset, - byteCount)) { - return; - } - if (tryCopyDirectUnionRange( - source, - absoluteSourceByteOffset, - destination, - absoluteDestinationByteOffset, - byteCount)) { - return; - } - if (tryCopyPrimitiveArrayRange( - source, - absoluteSourceByteOffset, - destination, - absoluteDestinationByteOffset, - byteCount)) { - return; - } - Pointer materializedSource = - materializeRelative(source, sourceElementOffset, sourceByteOffset); - Pointer materializedDestination = - materializeRelative(destination, destinationElementOffset, destinationByteOffset); - if (nonOverlapping) { - copyNonOverlapping(materializedSource, materializedDestination, byteCount); - } else { - copy(materializedSource, materializedDestination, byteCount); - } - } - - public static void copyRelative( - Pointer source, - long sourceElementOffset, - long sourceByteOffset, - Pointer destination, - long destinationElementOffset, - long destinationByteOffset, - long byteCount) { - copyRelativeBytes( - source, - sourceElementOffset, - sourceByteOffset, - destination, - destinationElementOffset, - destinationByteOffset, - checkedArrayLength(byteCount), - false); - } - - public static void copyElementsRelative( - Pointer source, - long sourceElementOffset, - long sourceByteOffset, - Pointer destination, - long destinationElementOffset, - long destinationByteOffset, - long elementCount) { - copyRelativeBytes( - source, - sourceElementOffset, - sourceByteOffset, - destination, - destinationElementOffset, - destinationByteOffset, - checkedElementByteCount(source, elementCount), - false); - } - public static void copyNonOverlapping(Pointer source, Pointer destination, int byteCount) { if (source.allocation == destination.allocation) { long sourceEnd = source.byteOffset + byteCount; @@ -10532,44 +10320,6 @@ public static void copyNonOverlappingElements( source, destination, checkedElementByteCount(source, elementCount)); } - public static void copyNonOverlappingRelative( - Pointer source, - long sourceElementOffset, - long sourceByteOffset, - Pointer destination, - long destinationElementOffset, - long destinationByteOffset, - long byteCount) { - copyRelativeBytes( - source, - sourceElementOffset, - sourceByteOffset, - destination, - destinationElementOffset, - destinationByteOffset, - checkedArrayLength(byteCount), - true); - } - - public static void copyNonOverlappingElementsRelative( - Pointer source, - long sourceElementOffset, - long sourceByteOffset, - Pointer destination, - long destinationElementOffset, - long destinationByteOffset, - long elementCount) { - copyRelativeBytes( - source, - sourceElementOffset, - sourceByteOffset, - destination, - destinationElementOffset, - destinationByteOffset, - checkedElementByteCount(source, elementCount), - true); - } - private static void swapBytes(Pointer left, Pointer right, int byteCount) { if (trySwapAlignedElements(left, right, byteCount)) { return; @@ -10645,7 +10395,7 @@ private static int rustIntegerCarrierValue(Object value) { if (current == null) { throw new IllegalArgumentException("Rust integer carrier was null"); } - Field[] fields = PUBLIC_INSTANCE_FIELDS.get(current.getClass()); + RustField[] fields = PUBLIC_INSTANCE_FIELDS.get(current.getClass()); if (fields.length != 1) { throw new IllegalArgumentException( "Rust integer carrier does not have one transparent field: " @@ -10716,7 +10466,7 @@ private static Object transparentSliceView(Pointer storage, Object value) { if (isSliceViewType(value.getClass())) { return value; } - Field[] fields = PUBLIC_INSTANCE_FIELDS.get(value.getClass()); + RustField[] fields = PUBLIC_INSTANCE_FIELDS.get(value.getClass()); if (fields.length != 1) { return null; } @@ -10822,6 +10572,12 @@ private Pointer sliceStorageView(long elementSize, String elementCodecClassName) return storage.retype(checkedElementSize, elementCodecClassName); } + private static boolean hasScalarWriteTracking(Object root) { + return mayBeInIdentityFilter(MEMORY_VIEW_FILTER, root) + || mayBeInIdentityFilter(ENCODED_POINTER_FILTER, root) + || mayBeInIdentityFilter(ENCODED_REFERENCE_FILTER, root); + } + public static boolean sliceGetBoolean(Object backing, int index) { if (backing instanceof boolean[]) { return ((boolean[]) backing)[index]; @@ -11466,13 +11222,7 @@ private static Object decodeArrayReference(long address, String codec, Class private static Object decodeArrayReference( Pointer pointer, String codec, Class targetClass) { Pointer data = decodedRawPointer(pointer, codec); - try { - return LONG_SLICE_VIEW_CONSTRUCTORS - .get(targetClass) - .newInstance(data, 0, arrayReferenceLength(codec)); - } catch (ReflectiveOperationException error) { - throw new IllegalStateException("could not reconstruct fixed-array reference", error); - } + return SliceView.create(targetClass, data, 0, arrayReferenceLength(codec)); } private static long rawPointerPointeeSize(String codec) { @@ -11554,8 +11304,8 @@ private Pointer nominalManagedPointee(String targetClassName) { if (targetClass.isInstance(value)) { return this; } - Field match = null; - for (Field candidate : PUBLIC_INSTANCE_FIELDS.get(value.getClass())) { + RustField match = null; + for (RustField candidate : PUBLIC_INSTANCE_FIELDS.get(value.getClass())) { if (targetClass.isAssignableFrom(candidate.getType())) { if (match != null) { return this; diff --git a/runtime/src/RustClasses.java b/runtime/src/RustClasses.java new file mode 100644 index 00000000..fa95bfcb --- /dev/null +++ b/runtime/src/RustClasses.java @@ -0,0 +1,31 @@ +package org.rustlang.runtime; + +import java.util.HashMap; +import java.util.Map; + +/** Nested storage identities come from class metadata, not binary-name suffixes. */ +final class RustClasses { + private static final ClassValue>> NESTED = new ClassValue>>() { + @Override + protected Map> computeValue(Class owner) { + Map> classes = new HashMap<>(); + for (Class nested : owner.getDeclaredClasses()) { + classes.put(nested.getSimpleName(), nested); + } + return classes; + } + }; + + private RustClasses() {} + + static boolean hasNested(Class owner, String first, String second) { + Map> classes = NESTED.get(owner); + return classes.containsKey(first) && classes.containsKey(second); + } + + static Class nested(Class owner, String name) throws ClassNotFoundException { + Class nested = NESTED.get(owner).get(name); + if (nested == null) throw new ClassNotFoundException(owner.getName() + " nested " + name); + return nested; + } +} diff --git a/runtime/src/RustField.java b/runtime/src/RustField.java new file mode 100644 index 00000000..37c8437a --- /dev/null +++ b/runtime/src/RustField.java @@ -0,0 +1,181 @@ +package org.rustlang.runtime; + +import java.lang.invoke.MethodHandle; +import java.lang.invoke.MethodHandles; +import java.lang.invoke.MethodType; +import java.lang.reflect.Field; +import java.lang.reflect.Modifier; +import java.util.LinkedHashMap; +import java.util.Map; + +/** Logical fields of Rust storage, independent of their physical JVM slots. */ +public final class RustField { + private static final ClassValue> FIELDS = + new ClassValue>() { + protected Map computeValue(Class type) { + Map result = new LinkedHashMap<>(); + Field[] fields = type.getFields(); + for (Field field : fields) { + if (Modifier.isStatic(field.getModifiers()) || field.isSynthetic()) continue; + Field displacement = null; + String shape = null; + if (field.getType() == Object.class || field.getType() == long.class) { + for (Field candidate : fields) { + String name = candidate.getName(); + if (!candidate.isSynthetic() || Modifier.isStatic(candidate.getModifiers()) + || !name.startsWith("$rust$")) continue; + int separator = name.indexOf('$', 6); + if (separator < 0 || !name.substring(separator + 1).equals(field.getName())) continue; + String kind = name.substring(6, separator); + if (kind.equals("l")) continue; + if (field.getType() == long.class && !kind.equals("t")) continue; + if (field.getType() == Object.class && kind.equals("t")) continue; + displacement = candidate; + shape = kind; + break; + } + } + result.put(field.getName(), new RustField(field, displacement, shape)); + } + return result; + } + }; + + static final ClassValue ALL = new ClassValue() { + protected RustField[] computeValue(Class type) { + return FIELDS.get(type).values().toArray(new RustField[0]); + } + }; + + public static RustField find(Class owner, String name) throws NoSuchFieldException { + RustField field = optional(owner, name); + if (field == null) throw new NoSuchFieldException(owner.getName() + "." + name); + return field; + } + + static RustField optional(Class owner, String name) { + return FIELDS.get(owner).get(name); + } + + private final Field root; + private final Field displacement; + private final Field length; + private final int size; + private final String codec; + private final boolean typedAddress; + private final Class logicalType; + + private RustField(Field root, Field displacement, String shape) { + this.root = root; + this.displacement = displacement; + root.setAccessible(true); + if (displacement != null) displacement.setAccessible(true); + typedAddress = shape != null && shape.startsWith("a"); + String layoutCodec = null; + try { + if (typedAddress) { + size = Integer.parseInt(shape.substring(1)); + length = null; + logicalType = Pointer.class; + try { + Field metadata = root.getDeclaringClass().getField("$rust$c$" + root.getName()); + if (metadata.getType() != String.class || !Modifier.isStatic(metadata.getModifiers())) { + throw new IllegalStateException("invalid address layout " + metadata); + } + layoutCodec = (String) metadata.get(null); + } catch (NoSuchFieldException absent) { + // Scalar or untyped byte layouts have no codec constant. + } + } else if ("t".equals(shape)) { + size = 0; length = null; logicalType = TaggedLong.class; + } else if ("s".equals(shape) || "u".equals(shape)) { + size = 0; + length = root.getDeclaringClass().getField("$rust$l$" + root.getName()); + length.setAccessible(true); + logicalType = shape.equals("u") ? Utf8View.class : SliceView.class; + } else { + size = shape == null ? 0 : Integer.parseInt(shape); + length = null; + logicalType = displacement == null ? root.getType() : Pointer.class; + } + } catch (ReflectiveOperationException error) { + throw new IllegalStateException("invalid borrowed field " + root, error); + } + codec = layoutCodec; + } + + public String getName() { return root.getName(); } + public Class getType() { return logicalType; } + public Class getDeclaringClass() { return root.getDeclaringClass(); } + boolean isBorrowed() { return displacement != null && logicalType != TaggedLong.class; } + boolean borrowedShape(boolean view) { + return isBorrowed() && (length != null) == view; + } + + /** Read the physical fields without constructing a logical carrier. */ + Object borrowedParts(Object owner, long[] metadata) throws IllegalAccessException { + Object value = root.get(owner); + metadata[0] = length == null ? displacement.getLong(owner) : displacement.getInt(owner); + if (length != null) metadata[1] = length.getLong(owner); + return value; + } + void setBorrowedParts(Object owner, Object value, long offset, long count) throws IllegalAccessException { + if (length == null) displacement.setLong(owner, offset); + else { + displacement.setInt(owner, Math.toIntExact(offset)); + length.setLong(owner, count); + } + root.set(owner, value); + } + MethodHandle getter() throws IllegalAccessException { + if (displacement == null) return MethodHandles.lookup().unreflectGetter(root); + return access("get", MethodType.methodType(Object.class, Object.class)); + } + MethodHandle setter() throws IllegalAccessException { + if (displacement == null) return MethodHandles.lookup().unreflectSetter(root); + return access("set", MethodType.methodType(void.class, Object.class, Object.class)); + } + private MethodHandle access(String name, MethodType type) throws IllegalAccessException { + try { return MethodHandles.lookup().findVirtual(RustField.class, name, type).bindTo(this); } + catch (NoSuchMethodException error) { throw new AssertionError(error); } + } + public int getInt(Object owner) throws IllegalAccessException { return root.getInt(owner); } + + public Object get(Object owner) throws IllegalAccessException { + Object value = root.get(owner); + if (displacement == null) return value; + if (logicalType == TaggedLong.class) return new TaggedLong(((Long) value).longValue(), displacement.getLong(owner)); + if (length == null) return typedAddress + ? Pointer.fromTypedStorageLocation(value, displacement.getLong(owner), size, codec) + : Pointer.addressFromParts(value, displacement.getLong(owner), size); + return SliceView.create(logicalType, value, displacement.getInt(owner), length.getLong(owner)); + } + + public void set(Object owner, Object value) throws IllegalAccessException { + if (logicalType == TaggedLong.class && displacement != null) { + TaggedLong tagged = (TaggedLong) value; + root.setLong(owner, TaggedLong.value(tagged)); + displacement.setLong(owner, TaggedLong.tag(tagged)); + return; + } + if (length != null) { + displacement.setInt(owner, viewStart(value)); + length.setLong(owner, viewLength(value)); + value = viewRoot(value); + } else if (displacement != null) { + value = (Pointer) value; + displacement.setLong(owner, 0); + } + root.set(owner, value); + } + + public static Object viewRoot(Object value) { + return value == null ? null : ((SliceView) value).array; + } + public static int viewStart(Object value) { + return value == null ? 0 : ((SliceView) value).offset; + } + public static long viewLength(Object value) { + return value == null ? 0 : ((SliceView) value).rustLength; + } +} diff --git a/runtime/src/SliceView.java b/runtime/src/SliceView.java new file mode 100644 index 00000000..dd8ec2be --- /dev/null +++ b/runtime/src/SliceView.java @@ -0,0 +1,100 @@ +package org.rustlang.runtime; + +import java.lang.reflect.Array; +import java.nio.charset.StandardCharsets; +import java.util.Objects; + +/** Carries slices across Java and opaque boundaries. Rust calls use components. */ +public class SliceView { + public final Object array; + public final int offset; + public final int length; + public final long rustLength; + + public SliceView(Object array, int offset, int length) { + this(array, offset, (long) length); + } + + public SliceView(Object array, int offset, long length) { + this.array = array; + this.offset = offset; + this.length = (int) length; + this.rustLength = length; + } + + static SliceView create(String className, Object array, int offset, long length) { + if (className.equals("org/rustlang/runtime/SliceView") + || className.equals("org.rustlang.runtime.SliceView")) { + return new SliceView(array, offset, length); + } + if (className.equals("org/rustlang/runtime/Utf8View") + || className.equals("org.rustlang.runtime.Utf8View")) { + return new Utf8View(array, offset, length); + } + throw new IllegalArgumentException("unknown Rust view carrier " + className); + } + + static SliceView create(Class type, Object array, int offset, long length) { + if (type == SliceView.class) return new SliceView(array, offset, length); + if (type == Utf8View.class) return new Utf8View(array, offset, length); + // An explicit Java subclass is an interop boundary, not a Rust carrier. + try { + return (SliceView) type.getConstructor(Object.class, int.class, long.class) + .newInstance(array, offset, length); + } catch (ReflectiveOperationException failure) { + throw new IllegalArgumentException("invalid Rust view carrier " + type, failure); + } + } + + public final Object toArray() { + Object result = Array.newInstance(array.getClass().getComponentType(), length); + System.arraycopy(array, offset, result, 0, length); + return result; + } + + public static SliceView fromString(String value) { + return (SliceView) Pointer.stringView(value, "org/rustlang/runtime/SliceView"); + } + + public static String toUtf8String(SliceView value) { + return new String(Pointer.sliceToByteArray(value.array, value.offset, value.length), + StandardCharsets.UTF_8); + } + + public static Utf8View encodeUtf8(int value, SliceView target) { + SliceView bytes = fromString(String.valueOf(Character.toChars(value))); + System.arraycopy(bytes.array, bytes.offset, target.array, target.offset, bytes.length); + return new Utf8View(target.array, target.offset, bytes.length); + } + + public static boolean startsWith(SliceView value, SliceView prefix) { + if (prefix.length > value.length) return false; + for (int i = 0; i < prefix.length; i++) { + if (!Objects.equals(Pointer.sliceGetObject(value.array, value.offset + i), + Pointer.sliceGetObject(prefix.array, prefix.offset + i))) return false; + } + return true; + } + + public static boolean startsWithI8(SliceView value, SliceView prefix) { + if (prefix.length > value.length) return false; + for (int i = 0; i < prefix.length; i++) { + if (Pointer.sliceGetI8(value.array, value.offset + i) + != Pointer.sliceGetI8(prefix.array, prefix.offset + i)) return false; + } + return true; + } + + public static boolean startsWithI32(SliceView value, SliceView prefix) { + if (prefix.length > value.length) return false; + for (int i = 0; i < prefix.length; i++) { + if (Pointer.sliceGetI32(value.array, value.offset + i) + != Pointer.sliceGetI32(prefix.array, prefix.offset + i)) return false; + } + return true; + } + + public static Object $part$array(SliceView value) { return value == null ? null : value.array; } + public static int $part$offset(SliceView value) { return value == null ? 0 : value.offset; } + public static long $part$rustLength(SliceView value) { return value == null ? 0 : value.rustLength; } +} diff --git a/runtime/src/Storage.java b/runtime/src/Storage.java new file mode 100644 index 00000000..b1d3a8fe --- /dev/null +++ b/runtime/src/Storage.java @@ -0,0 +1,46 @@ +package org.rustlang.runtime; + +/** Owns a typed Rust location. Code carries its byte offset separately. + * Create a Pointer only when an operation needs the general memory API. + */ +class Storage extends Pointer.Cell { + final int size; + final String codec; + final long metadata; + final String scalarLayout; + private volatile StorageLayout layout; + volatile Pointer boundary; + + Storage(Object value, int size, String codec, long metadata, String scalarLayout) { + super(value); + this.size = size; + this.codec = codec; + this.metadata = metadata; + this.scalarLayout = scalarLayout; + } + + StorageLayout layout(Object value) { + StorageLayout current = layout; + if (current == null && value != null && scalarLayout != null) { + layout = current = StorageLayout.of(value.getClass(), scalarLayout); + } + return current; + } + + Pointer boundary() { + Pointer result = boundary; + if (result == null) { + synchronized (this) { + result = boundary; + if (result == null) { + materialize(); + result = Pointer.storageBoundary(this); + boundary = result; + } + } + } + return result; + } + + void materialize() {} +} diff --git a/runtime/src/StorageLayout.java b/runtime/src/StorageLayout.java new file mode 100644 index 00000000..4e0138e8 --- /dev/null +++ b/runtime/src/StorageLayout.java @@ -0,0 +1,170 @@ +package org.rustlang.runtime; + +import java.lang.invoke.MethodHandle; +import java.lang.invoke.MethodHandles; +import java.lang.invoke.MethodType; +import java.util.Arrays; +import java.util.ArrayList; +import java.util.Comparator; +import java.util.concurrent.ConcurrentHashMap; + +/** Rust byte offsets mapped to authoritative JVM scalar fields. */ +final class StorageLayout { + private static final ClassValue> PLANS = + new ClassValue>() { + protected ConcurrentHashMap computeValue(Class type) { + return new ConcurrentHashMap<>(); + } + }; + + static StorageLayout of(Class type, String descriptor) { + return PLANS.get(type).computeIfAbsent(descriptor, key -> new StorageLayout(type, key)); + } + + private final Leaf[] leaves; + private final BorrowedFieldPath[] borrows; + + private StorageLayout(Class type, String descriptor) { + String[] entries = descriptor.split("\n", -1); + ArrayList scalars = new ArrayList<>(); + ArrayList borrowed = new ArrayList<>(); + try { + for (int i = 0; i < entries.length; i++) { + String[] parts = entries[i].split(",", -1); + if (parts.length == 4 && (parts[3].startsWith("p:") || parts[3].startsWith("v:"))) { + int lines = Integer.parseInt(parts[3].substring(2)); + if (lines < 1 || lines > entries.length - i - 1) + throw new IllegalArgumentException("invalid borrowed storage codec"); + String codec = String.join("\n", Arrays.copyOfRange(entries, i + 1, i + 1 + lines)); + borrowed.add(new BorrowedFieldPath(type, Integer.parseInt(parts[0]), + Integer.parseInt(parts[1]), parts[2], parts[3].charAt(0) == 'v', codec)); + i += lines; + continue; + } + scalars.add(new Leaf(type, Integer.parseInt(parts[0]), + Integer.parseInt(parts[1]), parts[2], + parts.length == 4 ? Integer.parseInt(parts[3]) : 0)); + } + } catch (ReflectiveOperationException error) { + throw new IllegalStateException("invalid typed Rust storage layout", error); + } + leaves = scalars.toArray(new Leaf[0]); + borrows = borrowed.toArray(new BorrowedFieldPath[0]); + Arrays.sort(leaves, Comparator.comparingInt(leaf -> leaf.offset)); + Arrays.sort(borrows, Comparator.comparingInt(leaf -> leaf.offset)); + } + + BorrowedFieldPath borrowedAt(long offset, boolean view) { + int low = 0, high = borrows.length; + while (low < high) { + int middle = (low + high) >>> 1; + BorrowedFieldPath path = borrows[middle]; + if (path.offset < offset) low = middle + 1; + else if (path.offset > offset) high = middle; + else return path.view == view && path.field().borrowedShape(view) ? path : null; + } + return null; + } + + Leaf at(long offset, int size) { + if (offset < 0 || size <= 0 || size > 8) return null; + int low = 0, high = leaves.length; + while (low < high) { + int middle = (low + high) >>> 1; + if (leaves[middle].offset <= offset) low = middle + 1; + else high = middle; + } + if (low == 0) return null; + Leaf leaf = leaves[low - 1]; + long relative = offset - leaf.offset; + long extent = (long) leaf.size * (leaf.count == 0 ? 1 : leaf.count); + return relative < extent && relative % leaf.size <= leaf.size - size ? leaf : null; + } + + static final class Leaf { + final int offset, size, count; + private final MethodHandle read, write; + + Leaf(Class type, int offset, int size, String path, int count) throws ReflectiveOperationException { + this.offset = offset; + this.size = size; + this.count = count; + String[] names = path.isEmpty() ? new String[0] : path.split("/"); + MethodHandle owner = MethodHandles.identity(Object.class); + int ownerPath = count == 0 ? names.length - 1 : names.length; + for (int i = 0; i < ownerPath; i++) { + RustField field = RustField.find(type, names[i]); + owner = MethodHandles.filterReturnValue(owner, + field.getter().asType(MethodType.methodType(Object.class, Object.class))); + type = field.getType(); + } + MethodHandle getter, setter; + Class scalar; + if (count == 0) { + RustField field = RustField.find(type, names[names.length - 1]); + getter = field.getter(); + setter = field.setter(); + scalar = field.getType(); + } else { + getter = MethodHandles.arrayElementGetter(type); + setter = MethodHandles.arrayElementSetter(type); + scalar = type.getComponentType(); + } + if (scalar == float.class || scalar == double.class) { + boolean single = scalar == float.class; + Class wrapper = single ? Float.class : Double.class; + Class bits = single ? int.class : long.class; + MethodHandle encode = MethodHandles.lookup().findStatic(wrapper, + single ? "floatToRawIntBits" : "doubleToRawLongBits", + MethodType.methodType(bits, scalar)); + MethodHandle decode = MethodHandles.lookup().findStatic(wrapper, + single ? "intBitsToFloat" : "longBitsToDouble", + MethodType.methodType(scalar, bits)); + getter = MethodHandles.filterReturnValue(getter, encode); + setter = MethodHandles.filterArguments(setter, count == 0 ? 1 : 2, decode); + } + read = MethodHandles.filterArguments(MethodHandles.explicitCastArguments(getter, + count == 0 ? MethodType.methodType(long.class, Object.class) + : MethodType.methodType(long.class, Object.class, int.class)), 0, owner); + write = MethodHandles.filterArguments(MethodHandles.explicitCastArguments(setter, + count == 0 ? MethodType.methodType(void.class, Object.class, long.class) + : MethodType.methodType(void.class, Object.class, int.class, long.class)), 0, owner); + } + + long read(Object root, long position, int count) { + try { + return (bits(root, position) >>> (((position - offset) % size) * 8)) & mask(count); + } catch (RuntimeException | Error error) { + throw error; + } catch (Throwable error) { + throw new IllegalStateException("could not read typed Rust field", error); + } + } + + void write(Object root, long position, long bits, int count) { + try { + int within = (int) ((position - offset) % size); + if (within != 0 || count != size) { + int shift = within * 8; + long mask = mask(count) << shift; + bits = (bits(root, position) & ~mask) | ((bits << shift) & mask); + } + if (this.count == 0) write.invokeExact(root, bits); + else write.invokeExact(root, Math.toIntExact((position - offset) / size), bits); + } catch (RuntimeException | Error error) { + throw error; + } catch (Throwable error) { + throw new IllegalStateException("could not write typed Rust field", error); + } + } + + private long bits(Object root, long position) throws Throwable { + return count == 0 ? (long) read.invokeExact(root) + : (long) read.invokeExact(root, Math.toIntExact((position - offset) / size)); + } + + private static long mask(int size) { + return size == 8 ? -1L : (1L << (size * 8)) - 1; + } + } +} diff --git a/runtime/src/TaggedLong.java b/runtime/src/TaggedLong.java new file mode 100644 index 00000000..36184b46 --- /dev/null +++ b/runtime/src/TaggedLong.java @@ -0,0 +1,22 @@ +package org.rustlang.runtime; + +/** Immutable boundary carrier for a scalar payload and independent Rust tag. */ +public final class TaggedLong { + public final long value; + public final long tag; + + public TaggedLong(long value, long tag) { + this.value = value; + this.tag = tag; + } + + public static TaggedLong of(long value, long tag) { return new TaggedLong(value, tag); } + + public static boolean optionEquals(TaggedLong left, TaggedLong right) { + long tag = tag(left); + return tag == tag(right) && (tag == 0 || value(left) == value(right)); + } + + public static long value(TaggedLong value) { return value == null ? 0 : value.value; } + public static long tag(TaggedLong value) { return value == null ? 0 : value.tag; } +} diff --git a/runtime/src/Utf8View.java b/runtime/src/Utf8View.java new file mode 100644 index 00000000..75716298 --- /dev/null +++ b/runtime/src/Utf8View.java @@ -0,0 +1,27 @@ +package org.rustlang.runtime; + +/** A UTF-8-valid boundary view. Length and offsets are measured in bytes. */ +public final class Utf8View extends SliceView { + public Utf8View(Object array, int offset, int length) { super(array, offset, length); } + public Utf8View(Object array, int offset, long length) { super(array, offset, length); } + + public static Utf8View fromJavaString(String value) { + return (Utf8View) Pointer.stringView(value, "org/rustlang/runtime/Utf8View"); + } + + public static String toJavaString(Utf8View value) { return toUtf8String(value); } + public static SliceView asSlice(Utf8View value) { return value; } + public static Utf8View fromSlice(SliceView value) { + return new Utf8View(value.array, value.offset, value.rustLength); + } + public static long len(Utf8View value) { return value.rustLength; } + public static boolean startsWith(Utf8View value, Utf8View prefix) { + return startsWithI8(value, prefix); + } + public static boolean equals(Utf8View value, Utf8View other) { + return value.rustLength == other.rustLength && startsWith(value, other); + } + public static boolean startsWithChar(Utf8View value, int character) { + return toJavaString(value).startsWith(String.valueOf(Character.toChars(character))); + } +} diff --git a/tests/integration/pointer_provenance/BorrowedFields.java b/tests/integration/pointer_provenance/BorrowedFields.java new file mode 100644 index 00000000..8c74780c --- /dev/null +++ b/tests/integration/pointer_provenance/BorrowedFields.java @@ -0,0 +1,181 @@ +import java.io.ByteArrayOutputStream; +import java.io.InputStream; +import java.lang.reflect.Field; +import java.nio.ByteBuffer; +import org.rustlang.runtime.Pointer; +import org.rustlang.runtime.RustField; +import org.rustlang.runtime.SliceView; + +/** Physical borrowed fields and references to them survive parent replacement. */ +public final class BorrowedFields { + private static final String VIEW = "@slice-pointer\norg/rustlang/runtime/SliceView\n1\n"; + private static final String ADDRESS = "@raw-pointer\n4\n\n"; + private static Class fixture; + + public static final class Codec { + public static void w$slots(Object value, byte[] bytes, int start) throws Exception { + Pointer.encodeFatPointerMemory(RustField.find(fixture, "view").get(value), bytes, start, 16, VIEW); + Pointer.array(bytes, 0, 1).byte_offset(start + 16).retype(8, ADDRESS) + .set(RustField.find(fixture, "address").get(value)); + } + public static Object a$slots(byte[] bytes, int start) throws Exception { + Object value = fixture.getConstructor().newInstance(); + RustField.find(fixture, "view").set(value, Pointer.decodeFatPointerMemory(bytes, start, 16, VIEW)); + RustField.find(fixture, "address").set(value, + Pointer.array(bytes, 0, 1).byte_offset(start + 16).retype(8, ADDRESS).getObject()); + return value; + } + } + public static class Slots { + public Object view; + public int $rust$s$view; + public long $rust$l$view; + public Object address; + public long $rust$4$address; + public Object typed; + public long $rust$a8$typed; + public static final String $rust$c$typed = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;#8"; + } + + // javac cannot declare synthetic fields. Set the flags that the backend emits. + public static Class slotsClass() throws Exception { + byte[] data; + try (InputStream input = BorrowedFields.class.getResourceAsStream("BorrowedFields$Slots.class")) { + ByteArrayOutputStream output = new ByteArrayOutputStream(); + byte[] buffer = new byte[4096]; + for (int count; (count = input.read(buffer)) >= 0;) output.write(buffer, 0, count); + data = output.toByteArray(); + } + ByteBuffer bytes = ByteBuffer.wrap(data); + bytes.position(8); + int count = Short.toUnsignedInt(bytes.getShort()); + String[] names = new String[count]; + for (int i = 1; i < count; i++) { + switch (Byte.toUnsignedInt(bytes.get())) { + case 1: + byte[] text = new byte[Short.toUnsignedInt(bytes.getShort())]; + bytes.get(text); + names[i] = new String(text, "UTF-8"); + break; + case 3: case 4: case 9: case 10: case 11: case 12: case 17: case 18: + bytes.position(bytes.position() + 4); break; + case 5: case 6: bytes.position(bytes.position() + 8); i++; break; + case 7: case 8: case 16: case 19: case 20: + bytes.position(bytes.position() + 2); break; + case 15: bytes.position(bytes.position() + 3); break; + default: throw new AssertionError("unknown fixture constant"); + } + } + bytes.position(bytes.position() + 6); + int interfaces = Short.toUnsignedInt(bytes.getShort()); + bytes.position(bytes.position() + interfaces * 2); + int fields = Short.toUnsignedInt(bytes.getShort()); + for (int i = 0; i < fields; i++) { + int position = bytes.position(); + short flags = bytes.getShort(); + String name = names[Short.toUnsignedInt(bytes.getShort())]; + if (name.startsWith("$rust$")) bytes.putShort(position, (short) (flags | 0x1000)); + bytes.getShort(); + int attributes = Short.toUnsignedInt(bytes.getShort()); + for (int j = 0; j < attributes; j++) { + bytes.getShort(); + int length = bytes.getInt(); + bytes.position(bytes.position() + length); + } + } + return new ClassLoader(BorrowedFields.class.getClassLoader()) { + Class define(byte[] bytes) { return defineClass(null, bytes, 0, bytes.length); } + }.define(data); + } + + public static void check() throws Exception { + Class type = slotsClass(); + fixture = type; + Object first = type.getConstructor().newInstance(); + Object second = type.getConstructor().newInstance(); + Field view = type.getField("view"); + Field address = type.getField("address"); + byte[] bytes = {3, 5, 7, 11}; + int[] words = {13, 17, 19}; + view.set(first, bytes); + type.getField("$rust$s$view").setInt(first, 1); + type.getField("$rust$l$view").setLong(first, 3); + address.set(first, words); + type.getField("$rust$4$address").setLong(first, 4); + byte[] typedBytes = new byte[24]; + org.rustlang.runtime.MemoryBytes.write(typedBytes, 8, 4, 101); + org.rustlang.runtime.MemoryBytes.write(typedBytes, 12, 4, 103); + org.rustlang.runtime.MemoryBytes.write(typedBytes, 16, 4, 107); + type.getField("typed").set(first, typedBytes); + type.getField("$rust$a8$typed").setLong(first, 8); + Pointer typed = (Pointer) RustField.find(type, "typed").get(first); + MemoryViews.Pair pair = (MemoryViews.Pair) typed.getObjectCopyAs(MemoryViews.Pair.class.getName()); + MemoryViews.Pair next = (MemoryViews.Pair) typed.add(1).getObjectCopyAs(MemoryViews.Pair.class.getName()); + if (pair.first != 101 || pair.second != 103 || next.first != 107) + throw new AssertionError("stored exact layout lost its stride or codec"); + RustField.find(type, "typed").set(second, typed.add(1)); + if (((MemoryViews.Pair) ((Pointer) RustField.find(type, "typed").get(second)) + .getObjectCopyAs(MemoryViews.Pair.class.getName())).first != 107) + throw new AssertionError("reflective exact-layout store lost its location"); + Pointer owner = Pointer.cell(first, 24, null); + // Supply the exact Class from the child loader. + Field allocation = Pointer.class.getDeclaredField("allocation"); + allocation.setAccessible(true); + java.lang.reflect.Method project = Pointer.class.getDeclaredMethod("rootField", + Object.class, Class.class, String.class, long.class, String.class); + project.setAccessible(true); + Pointer slice = (Pointer) project.invoke(null, allocation.get(owner), type, "view", 16L, null); + Pointer pointer = (Pointer) project.invoke(null, allocation.get(owner), type, "address", 8L, null); + long[] metadata = new long[2]; + if (Pointer.loadBorrowedView(slice, 0, metadata) != bytes || metadata[0] != 1 || metadata[1] != 3) + throw new AssertionError("view components were boxed or changed"); + if (Pointer.loadBorrowedAddress(pointer, 0, metadata) != words || metadata[0] != 4) + throw new AssertionError("address components were boxed or changed"); + owner.set(second); + Pointer.storeBorrowedView(slice, 0, bytes, 2, 2); + Pointer.storeBorrowedAddress(pointer, 0, words, 8, 4); + if (view.get(second) != bytes || address.get(second) != words + || type.getField("$rust$s$view").getInt(second) != 2 + || type.getField("$rust$4$address").getLong(second) != 8 + || type.getField("$rust$s$view").getInt(first) != 1) + throw new AssertionError("borrowed store detached after parent replacement"); + if (Pointer.loadBorrowedAddress(pointer, 0, metadata) != words || metadata[0] != 8) + throw new AssertionError("borrowed read retained replaced parent"); + Pointer.storeBorrowedAddress(pointer, 0, null, 0, 4); + if (Pointer.loadBorrowedAddress(pointer, 0, metadata) != null || metadata[0] != 0) + throw new AssertionError("nullable borrow changed"); + + String codec = "BorrowedFields$Codec#slots#Ljava/lang/Object;#24"; + Object storage = Pointer.storageAligned(first, 24, codec, 8, + "0,16,view,v:4\n" + VIEW + "\n16,8,address,p:4\n" + ADDRESS); + Object projected = Pointer.storageBorrowedFieldRoot(storage, 0, type.getName(), "view", 0, 16, VIEW); + if (projected != storage || Pointer.loadBorrowedView(projected, 0, metadata) != bytes) + throw new AssertionError("borrowed field address materialized a carrier"); + projected = Pointer.storageBorrowedFieldRoot(storage, 0, type.getName(), "address", 16, 8, ADDRESS); + if (projected != storage || Pointer.loadBorrowedAddress(projected, 16, metadata) != words) + throw new AssertionError("stored pointer address materialized a carrier"); + Pointer.storeStorageLocation(storage, 0, second); + Pointer.storeBorrowedView(storage, 0, bytes, 1, 3); + Pointer.storeBorrowedAddress(storage, 16, words, 4, 4); + if (Pointer.loadBorrowedView(storage, 0, metadata) != bytes || metadata[0] != 1 || metadata[1] != 3 + || Pointer.loadBorrowedAddress(storage, 16, metadata) != words || metadata[0] != 4) + throw new AssertionError("typed borrowed fields detached after replacement"); + Field boundary = storage.getClass().getDeclaredField("boundary"); + boundary.setAccessible(true); + if (boundary.get(storage) != null) throw new AssertionError("typed borrows allocated Pointer"); + + Pointer escaped = Pointer.addressFromParts(storage, 0, 64); + Pointer raw = Pointer.fromStorageLocation(storage, 0); + // The metadata word is shared with byte aliases of the enclosing owner. + raw.byte_offset(8).retype(8, null).set(2L); + SliceView sliceValue = (SliceView) escaped.getObject(); + if (sliceValue.rustLength != 2) throw new AssertionError("byte alias did not update stored borrow"); + Pointer.storeBorrowedView(storage, 0, bytes, 2, 1); + if (((SliceView) escaped.getObject()).rustLength != 1 || raw.byte_offset(8).retype(8, null).getI64() != 1) + throw new AssertionError("stored borrow did not update byte alias"); + Pointer.storeStorageLocation(storage, 0, first); + Pointer.storeBorrowedView(storage, 0, bytes, 0, 4); + if (((SliceView) escaped.getObject()).rustLength != 4) + throw new AssertionError("escaped borrow retained replaced owner"); + } +} diff --git a/tests/integration/pointer_provenance/Main.java b/tests/integration/pointer_provenance/Main.java index 5ee2bf59..1c93a06c 100644 --- a/tests/integration/pointer_provenance/Main.java +++ b/tests/integration/pointer_provenance/Main.java @@ -6,8 +6,11 @@ public class Main { public static void main(String[] args) throws Exception { ArrayViews.check(); MemoryViews.check(); + RangeCodecs.check(); + TypedFields.check(); FieldProjections.check(); CyclicFieldViews.check(); + BorrowedFields.check(); CodecAdapters.check(); StructuralViews.check(); Field field = Pointer.class.getDeclaredField("EXPOSED_ADDRESSES"); diff --git a/tests/integration/pointer_provenance/RangeCodecs.java b/tests/integration/pointer_provenance/RangeCodecs.java new file mode 100644 index 00000000..98abd2cc --- /dev/null +++ b/tests/integration/pointer_provenance/RangeCodecs.java @@ -0,0 +1,35 @@ +import java.lang.reflect.Field; +import java.lang.reflect.Method; +import java.util.Arrays; +import org.rustlang.runtime.Pointer; +import pointer_provenance.Pixel; + +/** Generated array codecs must snapshot before touching overlapping storage. */ +public final class RangeCodecs { + public static void check() throws Exception { + Pointer seed = pointer_provenance.pointer_provenance.pixel_storage(); + try { + Field field = Pointer.class.getDeclaredField("viewCodecClassName"); + field.setAccessible(true); + String recipe = (String) field.get(seed); + String[] parts = recipe.split("#"); + Class owner = Class.forName(parts[0].replace('/', '.')); + // Require range access to prevent a return to allocating whole snapshots. + Method encode = owner.getMethod("w$" + parts[1], Pixel.class, byte[].class, int.class); + byte[] bytes = {11, 13, 17, 19, 23, 29}; + encode.invoke(null, new Pixel(bytes), bytes, 1); + byte[] expected = {11, 11, 13, 17, 19, 29}; + if (!Arrays.equals(bytes, expected)) { + throw new AssertionError("range codec overwrote an unread array element"); + } + bytes = new byte[] {31, 37, 41, 43, 47, 53}; + Pointer destination = Pointer.array(bytes, 0, 1).retype(4, recipe); + Pointer.storeStorageLocation(destination, 1, new Pixel(bytes)); + if (!Arrays.equals(bytes, new byte[] {31, 31, 37, 41, 43, 53})) { + throw new AssertionError("storage cleared the source before its range codec ran"); + } + } finally { + pointer_provenance.pointer_provenance.free_pixel(seed); + } + } +} diff --git a/tests/integration/pointer_provenance/TypedFields.java b/tests/integration/pointer_provenance/TypedFields.java new file mode 100644 index 00000000..1d120d48 --- /dev/null +++ b/tests/integration/pointer_provenance/TypedFields.java @@ -0,0 +1,118 @@ +import java.lang.reflect.Field; +import org.rustlang.runtime.Pointer; + +public final class TypedFields { + public static final class Nested { + public MemoryViews.Pair pair; + public float single; + public double wide; + public boolean flag; + public char small; + public int[] words; + public float[] floats; + } + + public static void check() throws Exception { + originViews(); + String codec = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;"; + Object owner = Pointer.storageAligned(new MemoryViews.Pair(17, 31), 8, codec, 4, + "0,4,first\n4,4,second"); + Object field = Pointer.storageFieldRoot(owner, 0, MemoryViews.Pair.class.getName(), "second", 4, 4, null); + if (field != owner || Pointer.storageFieldOffset(field, owner, 4) != 4) { + throw new AssertionError("typed field allocated an address carrier"); + } + Pointer.storeLocationBits(field, 4, 73, 4); + if (((MemoryViews.Pair) Pointer.loadStorageLocation(owner, 0, null)).second != 73) { + throw new AssertionError("field write missed authoritative storage"); + } + Pointer.storeStorageLocation(owner, 0, new MemoryViews.Pair(101, 0x12345678)); + Pointer.storeLocationBits(field, 5, 0xab, 1); + if (Pointer.loadLocationBits(field, 4, 4) != 0x1234ab78L) { + throw new AssertionError("field retained the replaced parent or lost neighboring bytes"); + } + Field boundary = owner.getClass().getDeclaredField("boundary"); + boundary.setAccessible(true); + if (boundary.get(owner) != null) throw new AssertionError("typed scalar access materialized Pointer"); + Pointer raw = Pointer.fromLocation(owner, 4, 4); + raw.set(113); + if (Pointer.loadLocationBits(field, 4, 4) != 113) throw new AssertionError("byte boundary split storage identity"); + Pointer.storeLocationBits(field, 4, 127, 4); + if (raw.getI32() != 127) throw new AssertionError("old alias missed a typed write"); + + Nested arrays = new Nested(); + arrays.words = new int[] {0x12345678, 17, 23}; + arrays.floats = new float[] {-0.0f, Float.intBitsToFloat(0x7fc01234)}; + Object arraysOwner = Pointer.storageAligned(arrays, 20, null, 4, + "0,4,words,3\n12,4,floats,2"); + Pointer.storeLocationBits(arraysOwner, 5, 0xab, 1); + if (arrays.words[1] != 0xab11 || Pointer.loadLocationBits(arraysOwner, 16, 4) != 0x7fc01234L) { + throw new AssertionError("inline array layout lost bits or index"); + } + arrays.words = new int[] {31, 37, 41}; + Pointer.storeLocationBits(arraysOwner, 8, 43, 4); + Pointer.storeLocationBits(arraysOwner, 12, 0x7fc05678L, 4); + if (arrays.words[2] != 43 || Float.floatToRawIntBits(arrays.floats[0]) != 0x7fc05678 + || Pointer.loadLocationBits(arraysOwner, 4, 4) != 37 + || boundary.get(arraysOwner) != null) { + throw new AssertionError("inline array replacement detached or materialized storage"); + } + Object bareArray = Pointer.storageAligned(new long[] {47, 53}, 16, null, 8, "0,8,,2"); + Pointer.storeLocationBits(bareArray, 8, 59, 8); + if (Pointer.loadLocationBits(bareArray, 8, 8) != 59 || boundary.get(bareArray) != null) { + throw new AssertionError("root inline array materialized storage"); + } + + Nested nested = new Nested(); + nested.pair = new MemoryViews.Pair(3, 5); + nested.single = Float.intBitsToFloat(0x7fc01234); + nested.wide = -0.0; + nested.flag = true; + nested.small = 60000; + Object aggregate = Pointer.storageAligned(nested, 32, null, 8, + "0,4,pair/first\n4,4,pair/second\n8,4,single\n16,8,wide\n24,1,flag\n26,2,small"); + if (Pointer.loadLocationBits(aggregate, 8, 4) != 0x7fc01234L + || Pointer.loadLocationBits(aggregate, 16, 8) != Long.MIN_VALUE + || Pointer.loadLocationBits(aggregate, 24, 1) != 1 + || Pointer.loadLocationBits(aggregate, 26, 2) != 60000) { + throw new AssertionError("typed scalar bit conversion changed its value"); + } + nested.pair = new MemoryViews.Pair(11, 13); + Pointer.storeLocationBits(aggregate, 4, 19, 4); + Pointer.storeLocationBits(aggregate, 8, 0x80000000L, 4); + Pointer.storeLocationBits(aggregate, 16, 0x7ff8000000004567L, 8); + Pointer.storeLocationBits(aggregate, 24, 0, 1); + Pointer.storeLocationBits(aggregate, 26, 50000, 2); + if (nested.pair.second != 19 || Float.floatToRawIntBits(nested.single) != 0x80000000 + || Double.doubleToRawLongBits(nested.wide) != 0x7ff8000000004567L + || nested.flag || nested.small != 50000 || boundary.get(aggregate) != null) { + throw new AssertionError("nested scalar store materialized or detached storage"); + } + } + + private static void originViews() { + String codec = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;"; + Pointer root = Pointer.cellAligned(new MemoryViews.Pair(17, 31), 8, codec, 4); + Pointer field = root.projectStructField(MemoryViews.Pair.class.getName(), "second", 4, 4, null); + Pointer bytes = field.retype(1, null); + long address = field.addr(); + for (int displacement = -4; displacement <= 3; displacement++) { + Pointer shifted = bytes.byte_offset(displacement); + Pointer restored = shifted.byte_offset(-displacement).retype(4, null); + if (shifted.addr() != address + displacement || restored.addr() != address + || bytes.addr() != address || field.addr() != root.addr() + 4) { + throw new AssertionError("deriving an offset view changed another view's origin"); + } + } + root.set(new MemoryViews.Pair(37, 41)); + bytes.byte_offset(1).set((byte) 0x55); + if (field.getI32() != 0x5529 || ((MemoryViews.Pair) root.getObject()).first != 37) { + throw new AssertionError("derived origin lost replacement or adjacent-field coherence"); + } + root = null; + field = null; + System.gc(); + if (bytes.getI8() != 41 || bytes.byte_offset(1).getI8() != 0x55 || bytes.addr() != address) { + throw new AssertionError("shared origin did not keep backing storage alive"); + } + } +} diff --git a/tests/integration/pointer_provenance/src/lib.rs b/tests/integration/pointer_provenance/src/lib.rs index 01033b60..a9df0e59 100644 --- a/tests/integration/pointer_provenance/src/lib.rs +++ b/tests/integration/pointer_provenance/src/lib.rs @@ -1,3 +1,7 @@ +#![feature(custom_inner_attributes)] +#![feature(register_tool)] +#![register_tool(jvm_codegen)] +#![jvm_codegen::export] pub fn address_bits(expose: bool) -> usize { let value = std::hint::black_box(17_u64); let pointer = std::hint::black_box(&value as *const u64); @@ -12,3 +16,21 @@ pub fn address_bits(expose: bool) -> usize { pub fn format_many(count: u32) -> usize { (0..count).map(|value| format!("{value}").len()).sum() } + +#[repr(C)] +#[derive(Clone, Copy)] +pub struct Pixel { + pub bytes: [u8; 4], +} + +pub fn pixel_storage() -> *mut Pixel { + Box::into_raw(Box::new(Pixel { bytes: [1, 2, 3, 4] })) +} + +pub unsafe fn free_pixel(pointer: *mut Pixel) { + unsafe { drop(Box::from_raw(pointer)); } +} + +pub fn first_word(words: [u32; 2]) -> u32 { + words[0] +} From 6dcd80e2661e443fa02f864f433aa93e01f3145f Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 04:51:41 +1000 Subject: [PATCH 06/61] load typed views directly --- runtime/src/Pointer.java | 236 +++++++++++++++++- .../pointer_provenance/ArrayViews.java | 56 +++++ .../pointer_provenance/CellArrayViews.java | 60 +++++ .../integration/pointer_provenance/Main.java | 3 + .../pointer_provenance/OwnedFields.java | 77 ++++++ .../PrimitiveArrayCodecs.java | 48 ++++ 6 files changed, 476 insertions(+), 4 deletions(-) create mode 100644 tests/integration/pointer_provenance/CellArrayViews.java create mode 100644 tests/integration/pointer_provenance/OwnedFields.java create mode 100644 tests/integration/pointer_provenance/PrimitiveArrayCodecs.java diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index bb989ed9..ad9acdac 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -4985,7 +4985,6 @@ public static Pointer constantCell( }); } if (value != null - && !value.getClass().isArray() && pointer.allocation != null && pointer.allocation.getClass().isArray() && !pointer.allocation.getClass().getComponentType().isPrimitive() @@ -5092,6 +5091,16 @@ public static Pointer constantArray( recordAlignment(value, checkedAlignment); return array(value, 0, checkedSize, elementCodecClassName); }); + if (Array.getLength(value) == 1 + && pointer.allocation != null + && value.getClass().getComponentType().isArray() + && value.getClass().getComponentType() == pointer.allocation.getClass()) { + // CTFE can share an array with its sole element, which can also be an array. + // Preserve the element value and the original allocation identity. + return cellAligned(pointer.allocation, checkedSize, + elementCodecClassName, checkedAlignment) + .inheritAddressOrigin(pointer, 0); + } return pointer.retype(checkedSize, elementCodecClassName); } @@ -9099,6 +9108,167 @@ private static FieldCell directBorrowedField(Object root, long offset, boolean v } /** Copying a value never establishes a mutable decoded-view binding. */ + public static Object loadStorageCopy(Object root, long offset, String target) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (offset == 0 && directStorage(storage)) { + return copyManagedValue(((Cell) storage).value); + } + root = storage.boundary(); + } + Pointer base = (Pointer) root; + if (base.allocation instanceof byte[] && base.rareState == null + && base.addressState == null && base.viewSize > 0 + && isGeneratedAggregateCodec(base.viewCodecClassName)) { + MemoryCodec plan = codecPlan(base.viewCodecClassName); + if (plan.decodeAt() != null && target != null + && matchesBinaryClassName(target, plan.encodeParameterType.getName())) { + byte[] bytes = (byte[]) base.allocation; + long absolute = Math.addExact(base.byteOffset, offset); + int start = Math.toIntExact(absolute); + int size = base.materializedViewSize(); + if (start < 0 || start > bytes.length - size) { + throw new IndexOutOfBoundsException("aggregate read exceeds byte-addressable Rust storage"); + } + base.flushMemoryViewsOverlapping(absolute, size); + return plan.decodeAt().decode(bytes, start); + } + } + return fromStorageLocation(base, offset).getObjectCopyAs(target); + } + + /** Owned field read without parent and projected address wrappers on byte storage. */ + public static Object loadStorageFieldCopy(Object root, long offset, String owner, + String field, long fieldOffset, long size, String codec, String target) { + if (size > 0 && size <= Integer.MAX_VALUE + && (root instanceof byte[] || (root instanceof Pointer + && ((Pointer) root).allocation instanceof byte[] + && ((Pointer) root).traitMetadataCarrier() == null))) { + return loadTypedStorageCopy(root, Math.addExact(offset, fieldOffset), + (int) size, codec, target); + } + return fromStorageLocation(root, offset) + .projectStructField(owner, field, fieldOffset, size, codec).getObjectCopyAs(target); + } + + /** Store an owned field through its exact layout without two projected wrappers. */ + public static void storeStorageField(Object root, long offset, String owner, String field, + long fieldOffset, long size, String codec, Object value) { + if (size > 0 && size <= Integer.MAX_VALUE + && (root instanceof byte[] + || (root instanceof Pointer && ((Pointer) root).allocation instanceof byte[] + && ((Pointer) root).traitMetadataCarrier() == null))) { + storeTypedStorage(root, Math.addExact(offset, fieldOffset), (int) size, codec, value); + return; + } + fromStorageLocation(root, offset).projectStructField(owner, field, fieldOffset, size, codec).set(value); + } + + /** Read a borrowed value without a temporary Pointer. Keep binding and commit behavior for decoded views. */ + public static Object loadTypedStorage(Object root, long offset, int size, String codec, String target) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (offset == 0 && size == storage.size + && java.util.Objects.equals(codec, storage.codec) && directStorage(storage)) { + Object value = ((Cell) storage).value; + if (target != null && !target.isEmpty() || value == null || !value.getClass().isArray()) return value; + } + } + if (root instanceof Pointer) { + Pointer pointer = (Pointer) root; + if (size > 0 && pointer.viewSize == size + && java.util.Objects.equals(codec, pointer.viewCodecClassName) + && pointer.rareState == null && pointer.addressState == null + && pointer.isDirectAllocationView() && !mayHaveStructuralView(pointer.allocation)) { + Object value = loadObjectLocation(pointer, offset, target); + if (target != null && !target.isEmpty() || value == null || !value.getClass().isArray()) return value; + } + } + return fromTypedStorageLocation(root, offset, size, codec).getObjectAs(target); + } + + /** The copy's layout is compiler metadata, independent of the backing view. */ + public static Object loadTypedStorageCopy(Object root, long offset, int size, + String codec, String target) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (offset == 0 && size == storage.size + && java.util.Objects.equals(codec, storage.codec) && directStorage(storage)) { + return copyManagedValue(((Cell) storage).value); + } + root = storage.boundary(); + } + if (root instanceof byte[] && size > 0 && isGeneratedAggregateCodec(codec) + && !mayBeInIdentityFilter(MEMORY_VIEW_FILTER, root) + && !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, root)) { + MemoryCodec plan = codecPlan(codec); + if (plan.decodeAt() != null && target != null + && matchesBinaryClassName(target, plan.encodeParameterType.getName())) { + byte[] bytes = (byte[]) root; + int start = Math.toIntExact(offset); + if (start < 0 || start > bytes.length - size) { + throw new IndexOutOfBoundsException("aggregate read exceeds byte-addressable Rust storage"); + } + return plan.decodeAt().decode(bytes, start); + } + } + Pointer base = root instanceof Pointer ? (Pointer) root : fromLocation(root, 0, 1); + if (base.allocation instanceof byte[] && base.rareState == null + && base.addressState == null && size > 0 && isGeneratedAggregateCodec(codec)) { + MemoryCodec plan = codecPlan(codec); + if (plan.decodeAt() != null && target != null + && matchesBinaryClassName(target, plan.encodeParameterType.getName())) { + byte[] bytes = (byte[]) base.allocation; + long absolute = Math.addExact(base.byteOffset, offset); + int start = Math.toIntExact(absolute); + if (start < 0 || start > bytes.length - size) { + throw new IndexOutOfBoundsException("aggregate read exceeds byte-addressable Rust storage"); + } + base.flushMemoryViewsOverlapping(absolute, size); + return plan.decodeAt().decode(bytes, start); + } + } + Pointer escaped = base.escapedFieldStorage(offset); + return (escaped == null ? base.byteOffsetRetype(offset, size, codec) + : escaped.retype(size, codec)).getObjectCopyAs(target); + } + + /** Write with compiler-owned layout metadata, materializing only at a general boundary. */ + public static void storeTypedStorage(Object root, long offset, int size, String codec, Object value) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (offset == 0 && size == storage.size + && java.util.Objects.equals(codec, storage.codec) && directStorage(storage)) { + Cell cell = storage; + cell.value = convertDirectValue(cell.value, value, size); + return; + } + } + if (root instanceof byte[] && size > 0 && isGeneratedAggregateCodec(codec) + && !mayBeInIdentityFilter(MEMORY_VIEW_FILTER, root) + && !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, root) + && !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, value)) { + MemoryCodec plan = codecPlan(codec); + CodecCalls.RangeEncoder encode = plan.encodeAt(); + if (encode != null) { + byte[] bytes = (byte[]) root; + int start = Math.toIntExact(offset); + if (start < 0 || start > bytes.length - size) { + throw new IndexOutOfBoundsException("aggregate store exceeds byte-addressable Rust storage"); + } + discardEncodedPointers(bytes, start, size); + encode.encode(value, bytes, start); + return; + } + } + if (root instanceof Pointer + && ((Pointer) root).storeAggregateRange(offset, size, codec, value)) { + return; + } + fromTypedStorageLocation(root, offset, size, codec).set(value); + } + + /** Write an exact byte window without a temporary address or encoded image. */ private boolean storeAggregateRange(long offset, long byteSize, String codec, Object value) { if (!(allocation instanceof byte[]) || rareState != null || addressState != null || byteSize <= 0 || !isGeneratedAggregateCodec(codec) @@ -9120,6 +9290,64 @@ private boolean storeAggregateRange(long offset, long byteSize, String codec, Ob return true; } + public static Object directStorageAggregate(Object root, long offset, Class type) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (offset == 0 && directStorage(storage)) { + Object value = ((Cell) storage).value; + return type.isInstance(value) ? value : null; + } + root = storage.boundary(); + } + if (root instanceof byte[]) return null; + Pointer base = (Pointer) root; + if (offset == 0) { + Object direct = base.directAggregate(type); + if (direct != null) return direct; + } + Object embedded = base.directCellArrayElement(offset, type); + if (embedded != null || offset == 0) return embedded; + if (base.rareState != null || base.addressState != null + || !(base.allocation instanceof Object[]) + || base.allocationElementSize <= 0 + || !base.isDirectAllocationView() + || mayHaveStructuralView(base.allocation)) return null; + long absolute = Math.addExact(base.byteOffset, offset); + if (absolute % base.allocationElementSize != 0) return null; + base.flushMemoryViewsOverlapping(absolute, base.allocationElementSize); + Object value = base.readElement(Math.toIntExact(absolute / base.allocationElementSize)); + return type.isInstance(value) ? value : null; + } + + /** Borrow an aggregate element directly from a Cell that owns a fixed array. */ + private Object directCellArrayElement(long offset, Class type) { + if (!(allocation instanceof Cell) || addressState != null + || traitObjectCarrier() != null || traitMetadataCarrier() != null + || zeroSizedSourceViewSize() >= 0 || mayHaveStructuralView(allocation) + || !isGeneratedAggregateCodec(allocationCodecClassName) + || !isGeneratedAggregateCodec(viewCodecClassName)) return null; + MemoryCodec plan = codecPlan(allocationCodecClassName); + int stride = plan.arrayElementSize; + if (stride <= 0 || viewSize != stride + || !java.util.Objects.equals(viewCodecClassName, plan.arrayElementCodec)) return null; + long absolute = Math.addExact(byteOffset, offset); + if (absolute < 0 || absolute % stride != 0) return null; + Object value = ((Cell) allocation).value; + if (!(value instanceof Object[]) || !plan.encodeParameterType.isInstance(value) + || (long) ((Object[]) value).length * stride != allocationElementSize + || absolute / stride >= ((Object[]) value).length) return null; + flushMemoryViewsOverlapping(absolute, stride); + // Read the current carrier again. A byte alias or whole-array assignment can replace it. + value = ((Cell) allocation).value; + if (!(value instanceof Object[]) || !plan.encodeParameterType.isInstance(value) + || (long) ((Object[]) value).length * stride != allocationElementSize + || absolute / stride >= ((Object[]) value).length) return null; + Object element = independentRepeatedArrayElement(value, (int) (absolute / stride)); + // Nested array values require their separate origin-registration rules. + return element != null && !element.getClass().isArray() && type.isInstance(element) ? element : null; + } + + /** Keep a scalar projection on its typed owner while its layout is exact. */ public static Object storageFieldRoot(Object root, long offset, String owner, String field, long fieldOffset, long size, String codec) { // Scalar components supply their own width. Byte storage needs no layout carrier. @@ -9708,7 +9936,7 @@ public Object directAggregate(Class type) { /** Reads an owned aggregate value without creating a live alias of byte storage. */ public Object getObjectCopyAs(String targetClassName) { - if (allocation instanceof byte[] && rareState == null && viewSize > 0 + if (viewSize > 0 && !isDirectAllocationView() && isGeneratedAggregateCodec(viewCodecClassName)) { MemoryCodec plan = codecPlan(viewCodecClassName); if (targetClassName != null @@ -9724,8 +9952,8 @@ && matchesBinaryClassName(targetClassName, plan.encodeParameterType.getName())) flushMemoryViewsOverlapping(byteOffset, size); return plan.decodeAt().decode(bytes, offset); } - // Pointer-bearing and external codecs still use an independent - // image carrying the source's reference/provenance metadata. + // Owned reads need an independent image with provenance metadata. + // A live decoded view could overwrite source padding when flushed. return decodeAggregate(viewCodecClassName, loadRange(materializedViewSize())); } } diff --git a/tests/integration/pointer_provenance/ArrayViews.java b/tests/integration/pointer_provenance/ArrayViews.java index c4f889c2..b46d7cc1 100644 --- a/tests/integration/pointer_provenance/ArrayViews.java +++ b/tests/integration/pointer_provenance/ArrayViews.java @@ -53,6 +53,62 @@ public static void check() { throw new AssertionError("array window lost pointer provenance"); } boundedAllocation(); + scalarComponents(); + promotedArrayElement(); + } + + private static void promotedArrayElement() { + String codec = "org/rustlang/runtime/ArrayMemoryCodec#array#[I#8"; + for (boolean arrayFirst : new boolean[] {false, true}) { + String identity = "promoted-array-element-" + arrayFirst; + int[] value = {17, 29}; + Pointer element; + Pointer array; + if (arrayFirst) { + array = Pointer.constantArray(identity, new int[][] {value}, 8, codec, 4); + element = Pointer.constantCell(identity, value.clone(), 8, codec, 4); + } else { + element = Pointer.constantCell(identity, value, 8, codec, 4); + array = Pointer.constantArray(identity, new int[][] {value.clone()}, 8, codec, 4); + } + if (!Arrays.equals((int[]) Pointer.sliceGetObject(array, 0), value) + || !Arrays.equals((int[]) element.getObjectAs("[I"), value) + || !array.sameAddress(element) + || array.retype(4, null).getI32() != 17 + || element.byte_offset(4).retype(4, null).getI32() != 29) { + throw new AssertionError("promoted array and its element have inconsistent views"); + } + } + } + + private static void scalarComponents() { + long[] words = {0x1122334455667788L, 0x99aabbccddeeff00L}; + if (Pointer.loadLocationBits(words, 6, 4) != 0xff001122L) { + throw new AssertionError("unaligned read lost the next primitive element"); + } + Pointer.storeLocationBits(words, 7, 0x3ff0, 2); + if (words[0] != 0xf022334455667788L || words[1] != 0x99aabbccddeeff3fL) { + throw new AssertionError("partial write lost neighboring primitive bytes"); + } + float[] singles = {Float.intBitsToFloat(0x7fc01234), -0.0f}; + if (Pointer.loadLocationBits(singles, 0, 4) != 0x7fc01234L + || Pointer.loadLocationBits(singles, 4, 4) != 0x80000000L) { + throw new AssertionError("primitive reads canonicalized float bits"); + } + Pointer.storeLocationBits(singles, 1, 0x9876, 2); + if (Float.floatToRawIntBits(singles[0]) != 0x7f987634) { + throw new AssertionError("partial float write converted its payload"); + } + double[] doubles = {Double.longBitsToDouble(0x7ff8000000001234L)}; + Pointer.storeLocationBits(doubles, 1, 0xabcd, 2); + if (Double.doubleToRawLongBits(doubles[0]) != 0x7ff8000000abcd34L) { + throw new AssertionError("partial double write converted its payload"); + } + char[] chars = {'\uffff', '\u8123'}; + Pointer.storeLocationBits(chars, 1, 0xabcd, 2); + if (chars[0] != '\ucdff' || chars[1] != '\u81ab') { + throw new AssertionError("partial char write lost unsigned bits"); + } } private static void equal(byte[] expected, Object actual) { diff --git a/tests/integration/pointer_provenance/CellArrayViews.java b/tests/integration/pointer_provenance/CellArrayViews.java new file mode 100644 index 00000000..208f16e6 --- /dev/null +++ b/tests/integration/pointer_provenance/CellArrayViews.java @@ -0,0 +1,60 @@ +import org.rustlang.runtime.Pointer; + +public final class CellArrayViews { + private static final String ELEMENT = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;#8"; + private static final String ARRAY = "CellArrayViews$Codec#pairs#[LMemoryViews$Pair;#16"; + + public static final class Codec { + public static int s$pairs() { return 8; } + public static String c$pairs() { return ELEMENT; } + public static byte[] e$pairs(MemoryViews.Pair[] values) { + byte[] bytes = new byte[16]; + Pointer.encodeArrayMemory(values, bytes, 0, 8, ELEMENT); + return bytes; + } + public static MemoryViews.Pair[] d$pairs(byte[] bytes) { + MemoryViews.Pair[] values = new MemoryViews.Pair[2]; + Pointer.decodeArrayMemory(bytes, 0, values, 8, ELEMENT); + return values; + } + } + + public static void check() { + MemoryViews.Pair[] values = {new MemoryViews.Pair(3, 5), new MemoryViews.Pair(7, 11)}; + Pointer whole = Pointer.cell(values, 16, ARRAY); + Pointer elements = whole.retype(8, ELEMENT); + if (direct(elements, 0) != values[0] || direct(elements, 8) != values[1]) + throw new AssertionError("cell array element decoded instead of borrowing its carrier"); + direct(elements, 8).second = 13; + Pointer.commitStorageLocation(elements, 8); + if (whole.byte_offset(12).retype(4, null).getI32() != 13) + throw new AssertionError("direct element field write was not visible through bytes"); + MemoryViews.Pair[] replacement = {new MemoryViews.Pair(17, 19), new MemoryViews.Pair(23, 29)}; + whole.set(replacement); + if (direct(elements, 0) != replacement[0] || direct(elements, 8) != replacement[1]) + throw new AssertionError("element lookup retained a replaced whole array"); + Pointer decoded = whole.byte_offset(8).retype(8, ELEMENT); + MemoryViews.Pair pending = (MemoryViews.Pair) decoded.getObject(); + pending.second = 31; + if (direct(elements, 8).second != 31) + throw new AssertionError("element lookup missed pending decoded writes"); + whole.byte_offset(8).retype(4, null).set(37); + MemoryViews.Pair refreshed = direct(elements, 8); + if (refreshed.first != 37 || refreshed.second != 31 || direct(elements, 0).first != 17) + throw new AssertionError("byte alias did not update the selected element"); + if (Pointer.directStorageAggregate(elements, 4, MemoryViews.Pair.class) != null + || Pointer.directStorageAggregate(whole, 0, MemoryViews.Pair.class) != null + || Pointer.directStorageAggregate(whole.retype(4, ELEMENT), 0, MemoryViews.Pair.class) != null + || Pointer.directStorageAggregate(whole.retype(8, null), 0, MemoryViews.Pair.class) != null + || Pointer.directStorageAggregate(elements, 0, String.class) != null) + throw new AssertionError("incompatible array layout accepted"); + Pointer shortArray = Pointer.cell(new MemoryViews.Pair[] {new MemoryViews.Pair()}, 16, ARRAY) + .retype(8, ELEMENT); + if (Pointer.directStorageAggregate(shortArray, 0, MemoryViews.Pair.class) != null) + throw new AssertionError("array extent mismatch accepted"); + } + + private static MemoryViews.Pair direct(Pointer root, long offset) { + return (MemoryViews.Pair) Pointer.directStorageAggregate(root, offset, MemoryViews.Pair.class); + } +} diff --git a/tests/integration/pointer_provenance/Main.java b/tests/integration/pointer_provenance/Main.java index 1c93a06c..8601bd8f 100644 --- a/tests/integration/pointer_provenance/Main.java +++ b/tests/integration/pointer_provenance/Main.java @@ -5,11 +5,14 @@ public class Main { public static void main(String[] args) throws Exception { ArrayViews.check(); + CellArrayViews.check(); MemoryViews.check(); RangeCodecs.check(); + PrimitiveArrayCodecs.check(); TypedFields.check(); FieldProjections.check(); CyclicFieldViews.check(); + OwnedFields.check(); BorrowedFields.check(); CodecAdapters.check(); StructuralViews.check(); diff --git a/tests/integration/pointer_provenance/OwnedFields.java b/tests/integration/pointer_provenance/OwnedFields.java new file mode 100644 index 00000000..e33f2b6c --- /dev/null +++ b/tests/integration/pointer_provenance/OwnedFields.java @@ -0,0 +1,77 @@ +import java.util.Arrays; +import org.rustlang.runtime.Pointer; + +/** Owned projected fields remain snapshots across managed and byte aliases. */ +public final class OwnedFields { + private static final String CODEC = "org/rustlang/runtime/ArrayMemoryCodec#array#[B#4"; + + public static final class Pixel { + public byte[] channels; + public Pixel(byte... channels) { this.channels = channels; } + } + + private static byte[] read(Object root, long offset) { + return (byte[]) Pointer.loadStorageFieldCopy(root, offset, + Pixel.class.getName(), "channels", 4, 4, CODEC, "[B"); + } + + private static void write(Object root, long offset, byte[] value) { + Pointer.storeStorageField(root, offset, Pixel.class.getName(), "channels", 4, 4, CODEC, value); + } + + private static void equal(byte[] actual, int... values) { + byte[] expected = new byte[values.length]; + for (int i = 0; i < values.length; i++) expected[i] = (byte) values[i]; + if (!Arrays.equals(actual, expected)) throw new AssertionError(Arrays.toString(actual)); + } + + public static void check() { + byte[] bytes = new byte[32]; + System.arraycopy(new byte[] {11, 13, 17, 19}, 0, bytes, 12, 4); + Pointer base = Pointer.array(bytes, 0, 1); + for (Object root : new Object[] {bytes, base}) { + byte[] snapshot = read(root, 8); + snapshot[0] = 23; + equal(read(root, 8), 11, 13, 17, 19); + bytes[13] = 29; + equal(snapshot, 23, 13, 17, 19); + bytes[13] = 13; + byte[] replacement = {73, 79, 83, 89}; + write(root, 8, replacement); + equal(read(root, 8), 73, 79, 83, 89); + replacement[0] = 97; + equal(read(root, 8), 73, 79, 83, 89); + if (bytes[11] != 0 || bytes[16] != 0) throw new AssertionError("store crossed field boundary"); + write(root, 8, new byte[] {11, 13, 17, 19}); + try { + read(root, 26); + throw new AssertionError("field copy accepted an invalid byte window"); + } catch (IndexOutOfBoundsException expected) { } + } + equal(read(base.byte_offset(4), 4), 11, 13, 17, 19); + byte[] live = (byte[]) base.byte_offset(12).retype(4, CODEC).getObject(); + live[2] = 31; + equal(read(base, 8), 11, 13, 31, 19); + equal(read(bytes, 8), 11, 13, 31, 19); + base.addr(); + equal(read(base, 8), 11, 13, 31, 19); + write(base, 8, new byte[] {11, 13, 17, 19}); + equal(read(bytes, 8), 11, 13, 17, 19); + + Pixel pixel = new Pixel((byte) 37, (byte) 41, (byte) 43, (byte) 47); + for (Object root : new Object[] {Pointer.cell(pixel, 8, null), + Pointer.storageAligned(pixel, 8, null, 4)}) { + byte[] snapshot = read(root, 0); + pixel.channels[0] = 53; + equal(snapshot, 37, 41, 43, 47); + equal(read(root, 0), 53, 41, 43, 47); + Pointer.storeStorageLocation(root, 0, new Pixel((byte) 59, (byte) 61, (byte) 67, (byte) 71)); + equal(read(root, 0), 59, 61, 67, 71); + equal(snapshot, 37, 41, 43, 47); + write(root, 0, new byte[] {73, 79, 83, 89}); + equal(read(root, 0), 73, 79, 83, 89); + equal(snapshot, 37, 41, 43, 47); + pixel.channels[0] = 37; + } + } +} diff --git a/tests/integration/pointer_provenance/PrimitiveArrayCodecs.java b/tests/integration/pointer_provenance/PrimitiveArrayCodecs.java new file mode 100644 index 00000000..b2ef5a39 --- /dev/null +++ b/tests/integration/pointer_provenance/PrimitiveArrayCodecs.java @@ -0,0 +1,48 @@ +import java.lang.reflect.Array; +import java.util.Arrays; +import org.rustlang.runtime.Pointer; + +/** Shared codecs preserve exact lengths, offsets and scalar bit patterns. */ +public final class PrimitiveArrayCodecs { + public static void check() { + roundTrip(new boolean[] {true, false, true}, 1); + roundTrip(new byte[] {-128, -1, 0, 127}, 1); + roundTrip(new short[] {Short.MIN_VALUE, -1, 0, Short.MAX_VALUE}, 2); + roundTrip(new char[] {0, 255, 256, 65535}, 2); + roundTrip(new int[] {Integer.MIN_VALUE, -1, 0, Integer.MAX_VALUE}, 4); + roundTrip(new long[] {Long.MIN_VALUE, -1, 0, Long.MAX_VALUE}, 8); + roundTrip(new float[] {-0.0f, Float.intBitsToFloat(0x7fc01234), 7.25f}, 4); + roundTrip(new double[] {-0.0, Double.longBitsToDouble(0x7ff8123456789abcL), 7.25}, 8); + roundTrip(new byte[0], 1); + roundTrip(new int[0], 4); + roundTrip(new int[] {197}, 4); + byte[] overlap = {3, 5, 7, 11, 13}; + Pointer.encodeArrayMemory(overlap, overlap, 0, 1, null); + if (!Arrays.equals(overlap, new byte[] {3, 5, 7, 11, 13})) + throw new AssertionError("in-place byte encoding changed data"); + } + + private static void roundTrip(Object source, int width) { + int size = Array.getLength(source) * width; + String recipe = "org/rustlang/runtime/ArrayMemoryCodec#array#" + + source.getClass().getName() + "#" + size; + byte[] storage = new byte[size + 10]; + Arrays.fill(storage, (byte) 37); + Pointer.storeTypedStorage(storage, 5, size, recipe, source); + Object decoded = Pointer.array(storage, 0, 1).byte_offset(5).retype(size, recipe) + .getObjectCopyAs(source.getClass().getName()); + if (decoded == source || decoded.getClass() != source.getClass() + || Array.getLength(decoded) != Array.getLength(source)) + throw new AssertionError("shared codec changed ownership, type or length"); + byte[] original = new byte[size], restored = new byte[size]; + Pointer.encodeArrayMemory(source, original, 0, width, null); + Pointer.encodeArrayMemory(decoded, restored, 0, width, null); + if (!Arrays.equals(original, restored) + || !Arrays.equals(original, Arrays.copyOfRange(storage, 5, 5 + size))) + throw new AssertionError("shared codec changed scalar bits"); + for (int i = 0; i < 5; i++) { + if (storage[i] != 37 || storage[size + 5 + i] != 37) + throw new AssertionError("shared codec wrote outside its range"); + } + } +} From 1f8fd0aa4ca049fed39312f7394939494ba0c428 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 07:27:06 +1000 Subject: [PATCH 07/61] observe component addresses directly --- runtime/src/Pointer.java | 143 +++++++++++++++--- .../pointer_provenance/AddressOrder.java | 56 +++++++ .../pointer_provenance/BorrowedLocals.java | 80 ++++++++++ .../pointer_provenance/LocationQueries.java | 70 +++++++++ .../integration/pointer_provenance/Main.java | 3 + 5 files changed, 334 insertions(+), 18 deletions(-) create mode 100644 tests/integration/pointer_provenance/AddressOrder.java create mode 100644 tests/integration/pointer_provenance/BorrowedLocals.java create mode 100644 tests/integration/pointer_provenance/LocationQueries.java diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index ad9acdac..a3c67c15 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -6642,17 +6642,24 @@ private Pointer wrappingByteOffset(long byteCount) { } public long align_offset(long alignment) { + return alignmentOffset(numericAddress(), viewSize, alignment); + } + + private static long alignmentOffset(long address, long stride, long alignment) { if (alignment <= 0 || (alignment & (alignment - 1)) != 0) { throw new IllegalArgumentException("Rust pointer alignment must be a power of two"); } - long current = numericAddress(); - long attempts = viewSize == 0 ? 1 : alignment; - for (long elements = 0; elements < attempts; elements++) { - if ((current + (long) elements * viewSize) % alignment == 0) { - return elements; - } - } - return -1; + if (stride < 0) throw new IllegalArgumentException("negative Rust pointee size"); + long divisor = stride == 0 ? alignment : Math.min(Long.lowestOneBit(stride), alignment); + if ((address & (divisor - 1)) != 0) return -1; + long mask = alignment / divisor - 1; + if (mask == 0) return 0; + // Solve address + n * stride = 0 modulo the power-of-two alignment. + // Divide by the gcd to get an odd stride with an inverse modulo 2^64. + long odd = stride / divisor; + long inverse = odd; + for (int i = 0; i < 6; i++) inverse *= 2 - odd * inverse; + return -(address >>> Long.numberOfTrailingZeros(divisor)) * inverse & mask; } public static long align_offset(Pointer pointer, long alignment) { @@ -6959,6 +6966,7 @@ public int compareAddress(Pointer other) { if (other != null && allocation != null && allocation == other.allocation + && allocationElementSize > 0 && other.allocationElementSize > 0 && addressOrigin() == null && other.addressOrigin() == null && byteOffset >= 0 @@ -7155,7 +7163,7 @@ public static Object asRefOption(Pointer pointer, String someClassName, String n private long numericAddress() { Pointer origin = addressOrigin(); if (origin != null) { - return Math.addExact(origin.numericAddress(), addressOriginOffset()); + return origin.numericAddress() + addressOriginOffset(); } if (allocation == null) { return exposedAddress; @@ -7167,11 +7175,12 @@ private long numericAddress() { Long cachedBase = cachedAllocationBase(ALLOCATION_BASE_CACHE, allocation); if (cachedBase != null) { - return Math.addExact(cachedBase.longValue(), byteOffset); + return cachedBase.longValue() + byteOffset; } synchronized (ALLOCATIONS) { AllocationInfo info = allocationInfo(allocation); - return Math.addExact(allocationBase(info), byteOffset); + // Address arithmetic can wrap. Memory accesses check their own ranges. + return allocationBase(info) + byteOffset; } } @@ -7181,6 +7190,10 @@ private static long numericAddress(Pointer pointer) { /** Must be called while holding {@link #ALLOCATIONS}. */ private long allocationBase(AllocationInfo info) { + return allocationBase(allocation, allocationElementSize, info); + } + + private static long allocationBase(Object allocation, int allocationElementSize, AllocationInfo info) { Long cached = cachedAllocationBase(ALLOCATION_BASE_CACHE, allocation); if (cached != null) { return cached.longValue(); @@ -7238,9 +7251,7 @@ public static long encodedAddress(Pointer pointer, Object owner) { } Pointer origin = pointer.addressOrigin(); if (origin != null) { - long address = Math.addExact( - encodedAddress(origin, owner), - pointer.addressOriginOffset()); + long address = encodedAddress(origin, owner) + pointer.addressOriginOffset(); pointer.setPublishedAddress(address); return address; } @@ -7251,8 +7262,7 @@ public static long encodedAddress(Pointer pointer, Object owner) { PUBLISHED_ALLOCATION_BASE_CACHE, pointer.allocation); long address; if (publishedBase != null) { - address = Math.addExact( - publishedBase.longValue(), pointer.byteOffset); + address = publishedBase.longValue() + pointer.byteOffset; } else { synchronized (ALLOCATIONS) { AllocationInfo info = allocationInfo(pointer.allocation); @@ -7261,7 +7271,7 @@ public static long encodedAddress(Pointer pointer, Object owner) { PUBLISHED_ALLOCATION_BASE_CACHE, pointer.allocation, base); - address = Math.addExact(base, pointer.byteOffset); + address = base + pointer.byteOffset; } } pointer.setPublishedAddress(address); @@ -7368,7 +7378,7 @@ private ExposedTarget exposedTargetAt(long displacement) { public long address() { Pointer origin = addressOrigin(); if (origin != null) { - long address = Math.addExact(origin.address(), addressOriginOffset()); + long address = origin.address() + addressOriginOffset(); setPublishedAddress(address); return address; } @@ -8868,6 +8878,37 @@ public static void atomicFence(int ordering) { } } + /** Thin-pointer comparison on storage components, preserving exposed addresses. */ + public static boolean sameLocation(Object left, long leftOffset, Object right, long rightOffset) { + if (left == right) return leftOffset == rightOffset; + if (left instanceof Pointer && right instanceof Pointer) { + Pointer a = (Pointer) left; + Pointer b = (Pointer) right; + if (a.allocation != null && a.allocation == b.allocation) { + return a.byteOffset + leftOffset == b.byteOffset + rightOffset; + } + } + return locationAddress(left) + leftOffset == locationAddress(right) + rightOffset; + } + + private static long locationAddress(Object root) { + return locationAddr(root, 0); + } + + /** Read an address word without materializing or exposing a Rust pointer. */ + public static long locationAddr(Object root, long offset) { + if (root == null) return offset; + if (root instanceof Pointer) return ((Pointer) root).numericAddress() + offset; + Object normalized = normalizeLocationOrigin(root); + if (normalized instanceof Pointer) return ((Pointer) normalized).numericAddress() + offset; + Long cached = cachedAllocationBase(ALLOCATION_BASE_CACHE, root); + if (cached != null) return cached.longValue() + offset; + int size = root instanceof Storage ? ((Storage) root).size : inferredArrayElementSize(root); + synchronized (ALLOCATIONS) { + return allocationBase(root, size, allocationInfo(root)) + offset; + } + } + private static Object normalizeLocationOrigin(Object root) { if (root == null || root instanceof Pointer || root instanceof Storage || !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, root)) return root; @@ -8881,6 +8922,72 @@ private static Object normalizeLocationOrigin(Object root) { } /** Pointer differences use allocation identity, never exposed address lookup. */ + public static long byteOffsetLocations(Object root, long offset, Object origin, long originOffset) { + root = normalizeLocationOrigin(root); + origin = normalizeLocationOrigin(origin); + Object allocation = root instanceof Pointer ? ((Pointer) root).provenanceAllocation() : root; + Object originAllocation = origin instanceof Pointer + ? ((Pointer) origin).provenanceAllocation() : origin; + if (allocation != originAllocation) { + throw new IllegalArgumentException("byte_offset_from requires pointers into one allocation"); + } + long start = root instanceof Pointer + ? Math.addExact(((Pointer) root).provenanceByteOffset(), offset) : offset; + long end = origin instanceof Pointer + ? Math.addExact(((Pointer) origin).provenanceByteOffset(), originOffset) : originOffset; + return Math.subtractExact(start, end); + } + + public static long offsetLocations(Object root, long offset, Object origin, long originOffset, long stride) { + if (stride == -1) stride = locationStride(root); + if (stride == 0) throw new ArithmeticException("offset_from is undefined for zero-sized pointees"); + long bytes = byteOffsetLocations(root, offset, origin, originOffset); + if (bytes % stride != 0) { + throw new ArithmeticException("pointer distance is not a whole number of elements"); + } + return bytes / stride; + } + + private static long unsignedLocationDistance(long distance) { + if (distance < 0) { + throw new ArithmeticException("offset_from_unsigned requires self at or after origin"); + } + return distance; + } + + public static long byteOffsetLocationsUnsigned(Object root, long offset, Object origin, long originOffset) { + return unsignedLocationDistance(byteOffsetLocations(root, offset, origin, originOffset)); + } + + public static long offsetLocationsUnsigned(Object root, long offset, Object origin, long originOffset, long stride) { + return unsignedLocationDistance(offsetLocations(root, offset, origin, originOffset, stride)); + } + + public static long alignLocation(Object root, long offset, long stride, long alignment) { + return alignmentOffset(locationAddr(root, offset), stride == -1 ? locationStride(root) : stride, alignment); + } + + /** Compare decomposed data addresses without creating temporary carriers. */ + public static int compareLocations(Object left, long leftOffset, Object right, long rightOffset) { + if (left == right && left != null && leftOffset >= 0 && rightOffset >= 0 + && (left instanceof Storage || (left.getClass().isArray() + && !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, left)))) { + return Long.compareUnsigned(leftOffset, rightOffset); + } + if (left instanceof Pointer && right instanceof Pointer) { + Pointer a = (Pointer) left, b = (Pointer) right; + long x = a.byteOffset + leftOffset, y = b.byteOffset + rightOffset; + if (a.allocation != null && a.allocation == b.allocation + && a.allocationElementSize > 0 && b.allocationElementSize > 0 + && a.addressOrigin() == null && b.addressOrigin() == null && x >= 0 && y >= 0) { + return Long.compareUnsigned(x, y); + } + } + // Other offsets and allocations require the full unsigned-address and provenance checks. + return Long.compareUnsigned(locationAddress(left) + leftOffset, locationAddress(right) + rightOffset); + } + + /** Element layout retained by a general storage root. */ public static long locationStride(Object root) { return root instanceof Storage ? ((Storage) root).size : ((Pointer) root).viewSize; } diff --git a/tests/integration/pointer_provenance/AddressOrder.java b/tests/integration/pointer_provenance/AddressOrder.java new file mode 100644 index 00000000..0896df68 --- /dev/null +++ b/tests/integration/pointer_provenance/AddressOrder.java @@ -0,0 +1,56 @@ +import org.rustlang.runtime.Pointer; + +/** Location comparisons preserve unsigned ordering without borrowing wrappers. */ +public final class AddressOrder { + public static final class Empty {} + public static final class Pair { public int first, second; } + + private static void compare(Pointer left, long x, Pointer right, long y) { + long a = Pointer.addr(left) + x, b = Pointer.addr(right) + y; + int expected = Long.compareUnsigned(a, b); + if (Pointer.compareLocations(left, x, right, y) != expected) { + throw new AssertionError("location ordering lost its address or origin"); + } + } + + public static void check() { + byte[] bytes = new byte[32]; + Pointer start = Pointer.array(bytes, 0, 1), end = start.add(16); + for (long x : new long[] {Long.MIN_VALUE, -32, -1, 0, 16, Long.MAX_VALUE}) { + for (long y : new long[] {Long.MIN_VALUE, -1, 0, 1, Long.MAX_VALUE}) { + compare(start, x, end, y); + } + encoded(Pointer.wrapping_byte_offset(start, x)); + } + if (Pointer.compareLocations(bytes, 7, bytes, 17) != -1 + || Pointer.compareLocations(null, -1, null, 0) != 1) { + throw new AssertionError("array or unprovenanced ordering"); + } + Object storage = Pointer.storageAligned(17, 4, null, 4); + if (Pointer.compareLocations(storage, 0, storage, 4) != -1) { + throw new AssertionError("typed storage ordering"); + } + // A typed ZST can have a dangling address. Positive offsets can wrap across unsigned zero. + Pointer zst = Pointer.withoutProvenance(-2L, 1) + .retype(0, "@zero-sized:" + Empty.class.getName()); + Pointer before = Pointer.wrapping_byte_offset(zst.retype(1), 1); + Pointer after = Pointer.wrapping_byte_offset(zst.retype(1), 3); + compare(before, 0, after, 0); + if (before.compareAddress(after) <= 0 || before.addr() != -1L || after.addr() != 1L) { + throw new AssertionError("dangling ZST address wrap"); + } + encoded(before); + encoded(after); + Pointer parent = Pointer.cell(new Pair(), 8, null); + Pointer field = parent.projectStructField(Pair.class.getName(), "second", 4, 4, null); + encoded(Pointer.wrapping_byte_offset(field, Long.MAX_VALUE)); + // Also cover the first publication, before any base-address cache hit. + encoded(Pointer.wrapping_byte_offset(Pointer.array(new byte[1], 0, 1), Long.MAX_VALUE)); + } + + private static void encoded(Pointer pointer) { + if (Pointer.encodedAddress(pointer, new byte[8]) != pointer.addr()) { + throw new AssertionError("encoding changed a wrapping address"); + } + } +} diff --git a/tests/integration/pointer_provenance/BorrowedLocals.java b/tests/integration/pointer_provenance/BorrowedLocals.java new file mode 100644 index 00000000..980c1626 --- /dev/null +++ b/tests/integration/pointer_provenance/BorrowedLocals.java @@ -0,0 +1,80 @@ +import java.lang.reflect.Field; +import org.rustlang.runtime.Pointer; +import org.rustlang.runtime.SliceView; + +/** Decomposed local borrows become ordinary coherent cells at byte boundaries. */ +public final class BorrowedLocals { + private static final String ADDRESS = "@raw-pointer\n4\n\n"; + private static final String VIEW = "@slice-pointer\norg/rustlang/runtime/SliceView\n1\n"; + + private static void unboxed(Object storage) throws Exception { + Field boundary = storage.getClass().getSuperclass().getDeclaredField("boundary"); + boundary.setAccessible(true); + if (boundary.get(storage) != null) throw new AssertionError("borrow allocated a boundary"); + } + + public static void check() throws Exception { + int[] words = {11, 13, 17}; + long[] metadata = new long[2]; + Object slot = Pointer.borrowedStorageAligned(null, 8, ADDRESS, 8, 0); + if (Pointer.loadBorrowedAddress(slot, 0, metadata) != null || metadata[0] != 0) + throw new AssertionError("initial null borrow changed"); + for (int i = 0; i < words.length; i++) { + Pointer.storeBorrowedAddress(slot, 0, words, i * 4, 4); + if (Pointer.loadBorrowedAddress(slot, 0, metadata) != words || metadata[0] != i * 4) + throw new AssertionError("local address components changed"); + } + unboxed(slot); + Pointer.storeBorrowedAddress(slot, 0, words, 4, 4); + Pointer escaped = Pointer.addressFromParts(slot, 0, 128); + Pointer value = (Pointer) escaped.getObject(); + value.set(29); + if (words[1] != 29) throw new AssertionError("stored borrow lost its pointee"); + Pointer.storeBorrowedAddress(slot, 0, words, 8, 4); + if (((Pointer) escaped.getObject()).getI32() != 17 || value.getI32() != 29) + throw new AssertionError("replacing a borrow changed its prior value"); + escaped.set(Pointer.array(words, 0, 4)); + Object root = Pointer.loadBorrowedAddress(slot, 0, metadata); + if (Pointer.loadLocationBits(root, metadata[0], 4) != 11) + throw new AssertionError("opaque store was not authoritative"); + Pointer.storeBorrowedAddress(slot, 0, null, 0, 4); + if (escaped.getObject() != null) throw new AssertionError("escaped null store changed"); + + byte[] bytes = {3, 5, 7, 11}; + Object viewSlot = Pointer.borrowedStorageAligned(null, 16, VIEW, 8, 1); + Pointer.storeBorrowedView(viewSlot, 0, bytes, 1, 3); + if (Pointer.loadBorrowedView(viewSlot, 0, metadata) != bytes || metadata[0] != 1 || metadata[1] != 3) + throw new AssertionError("local slice components changed"); + unboxed(viewSlot); + Pointer raw = Pointer.fromStorageLocation(viewSlot, 0); + SliceView saved = (SliceView) raw.getObject(); + raw.byte_offset(8).retype(8, null).set(2L); + if (((SliceView) raw.getObject()).rustLength != 2) + throw new AssertionError("byte store did not update materialized metadata"); + Pointer.storeBorrowedView(viewSlot, 0, bytes, 2, 1); + if (Pointer.loadBorrowedView(viewSlot, 0, metadata) != bytes || metadata[0] != 2 || metadata[1] != 1) + throw new AssertionError("view components ignored escaped storage"); + if (saved.offset != 1) throw new AssertionError("borrow replacement mutated prior view"); + + // Byte access must use the codec even before a Pointer escapes. + Object addressBytes = Pointer.borrowedStorageAligned(null, 8, ADDRESS, 8, 0); + Pointer.storeBorrowedAddress(addressBytes, 0, words, 8, 4); + long address = Pointer.loadLocationBits(addressBytes, 0, 8); + if (Pointer.fromEncodedAddress(address, 4, null, ADDRESS).getI32() != 17) + throw new AssertionError("deferred carrier lost exposed provenance"); + + Object self = Pointer.borrowedStorageAligned(null, 8, ADDRESS, 8, 0); + Pointer.storeBorrowedAddress(self, 0, self, 0, 0); + Pointer selfBoundary = Pointer.fromStorageLocation(self, 0); + if (!Pointer.sameLocation(selfBoundary.getObject(), 0, self, 0)) + throw new AssertionError("self-referential borrow lost its identity"); + Object other = Pointer.borrowedStorageAligned(null, 8, ADDRESS, 8, 0); + Object cycle = Pointer.borrowedStorageAligned(null, 8, ADDRESS, 8, 0); + Pointer.storeBorrowedAddress(other, 0, cycle, 0, 0); + Pointer.storeBorrowedAddress(cycle, 0, other, 0, 0); + Pointer left = Pointer.fromStorageLocation(other, 0); + Pointer right = Pointer.fromStorageLocation(cycle, 0); + if (left.getObject() != right || right.getObject() != left) + throw new AssertionError("mutually referential storage lost its identity"); + } +} diff --git a/tests/integration/pointer_provenance/LocationQueries.java b/tests/integration/pointer_provenance/LocationQueries.java new file mode 100644 index 00000000..323cc51a --- /dev/null +++ b/tests/integration/pointer_provenance/LocationQueries.java @@ -0,0 +1,70 @@ +import java.lang.reflect.Field; +import org.rustlang.runtime.Pointer; + +/** Address observations must preserve provenance without forcing boundary carriers. */ +public final class LocationQueries { + public static final class Pair { public int first, second; } + + public static void check() throws Exception { + Object storage = Pointer.storageAligned(17, 4, null, 64); + long address = Pointer.locationAddr(storage, 0); + if ((address & 63) != 0 || Pointer.locationAddr(storage, 3) != address + 3 + || Pointer.offsetLocations(storage, 12, storage, 0, 4) != 3 + || Pointer.alignLocation(storage, 1, 1, 16) != 15) { + throw new AssertionError("storage address query changed its layout"); + } + Field boundary = storage.getClass().getDeclaredField("boundary"); + boundary.setAccessible(true); + if (boundary.get(storage) != null) throw new AssertionError("address query materialized a pointer"); + Pointer materialized = Pointer.fromStorageLocation(storage, 0); + if (materialized.addr() != address || Pointer.byteOffsetLocations(storage, 3, materialized, 0) != 3) { + throw new AssertionError("materialization changed allocation identity"); + } + int[] values = new int[16]; + long base = Pointer.locationAddr(values, 0); + Pointer middle = Pointer.array(values, 4, 4); + if (base + 16 != middle.addr() || Pointer.offsetLocations(middle, 8, values, 4, 4) != 5 + || Pointer.byteOffsetLocations(values, 4, middle, 8) != -20 + || Pointer.locationAddr(values, Long.MAX_VALUE) != base + Long.MAX_VALUE) { + throw new AssertionError("array address query changed displacement"); + } + Pointer pair = Pointer.cell(new Pair(), 8, null); + Pointer second = pair.projectStructField(Pair.class.getName(), "second", 4, 4, null); + if (Pointer.locationAddr(second, 3) != pair.addr() + 7 + || Pointer.byteOffsetLocations(second, -4, pair, 0) != 0) { + throw new AssertionError("projected field lost its parent provenance"); + } + expect(IllegalArgumentException.class, () -> Pointer.byteOffsetLocations(values, 0, new int[16], 0)); + expect(ArithmeticException.class, () -> Pointer.offsetLocations(values, 3, values, 0, 4)); + expect(ArithmeticException.class, () -> Pointer.offsetLocations(values, 0, values, 0, 0)); + expect(ArithmeticException.class, () -> Pointer.byteOffsetLocationsUnsigned(values, 0, values, 1)); + expect(IllegalArgumentException.class, () -> Pointer.alignLocation(values, 0, 4, 3)); + for (long alignment = 1; alignment <= 256; alignment *= 2) { + for (long stride = 0; stride < 40; stride++) { + for (long bits : new long[] {0, 1, 3, 12, 255, -1, Long.MIN_VALUE, Long.MAX_VALUE}) { + long expected = -1; + for (long n = 0; n < alignment; n++) { + if ((bits + n * stride) % alignment == 0) { expected = n; break; } + } + long result = Pointer.alignLocation(null, bits, stride, alignment); + if (result != expected || Pointer.withoutProvenance(bits, stride).align_offset(alignment) != expected) { + throw new AssertionError("incorrect modular alignment solution"); + } + } + } + } + // The previous linear search cannot feasibly handle this alignment. + if (Pointer.alignLocation(null, 1, 3, 1L << 62) != 1537228672809129301L + || Pointer.alignLocation(null, 1, 2, 1L << 62) != -1) { + throw new AssertionError("large modular alignment solution"); + } + } + + private static void expect(Class type, Runnable action) { + try { action.run(); } catch (Throwable error) { + if (type.isInstance(error)) return; + throw error; + } + throw new AssertionError("missing " + type.getSimpleName()); + } +} diff --git a/tests/integration/pointer_provenance/Main.java b/tests/integration/pointer_provenance/Main.java index 8601bd8f..801b772d 100644 --- a/tests/integration/pointer_provenance/Main.java +++ b/tests/integration/pointer_provenance/Main.java @@ -12,8 +12,11 @@ public static void main(String[] args) throws Exception { TypedFields.check(); FieldProjections.check(); CyclicFieldViews.check(); + AddressOrder.check(); + LocationQueries.check(); OwnedFields.check(); BorrowedFields.check(); + BorrowedLocals.check(); CodecAdapters.check(); StructuralViews.check(); Field field = Pointer.class.getDeclaredField("EXPOSED_ADDRESSES"); From 27e23043237dd55279aefc4784b90f32e0b5acee Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 09:39:48 +1000 Subject: [PATCH 08/61] use component addresses for atomics --- runtime/src/Pointer.java | 200 ++++++++++++------ .../pointer_provenance/AtomicLocations.java | 79 +++++++ .../pointer_provenance/AtomicViews.java | 92 ++++++++ .../integration/pointer_provenance/Main.java | 3 + .../pointer_provenance/MetadataFilters.java | 77 +++++++ 5 files changed, 391 insertions(+), 60 deletions(-) create mode 100644 tests/integration/pointer_provenance/AtomicLocations.java create mode 100644 tests/integration/pointer_provenance/AtomicViews.java create mode 100644 tests/integration/pointer_provenance/MetadataFilters.java diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index a3c67c15..3904495c 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -2683,6 +2683,12 @@ private void flushMemoryViewsOverlapping(long offset, int size, boolean overwrit || !mayBeInIdentityFilter(MEMORY_VIEW_FILTER, allocation)) { return; } + synchronized (atomicStripe(this)) { + flushMemoryViewsOverlappingLocked(offset, size, overwrite); + } + } + + private void flushMemoryViewsOverlappingLocked(long offset, int size, boolean overwrite) { long observedEpoch = memoryViewEpoch(allocation); MemoryViewAbsenceCache absence = MEMORY_VIEW_ABSENCE.get(); if (absence.matches(allocation, observedEpoch)) { @@ -2726,6 +2732,12 @@ private void flushAllMemoryViews() { || !mayBeInIdentityFilter(MEMORY_VIEW_FILTER, allocation)) { return; } + synchronized (atomicStripe(this)) { + flushAllMemoryViewsLocked(); + } + } + + private void flushAllMemoryViewsLocked() { LongRangeMap pending; Map> stripe = stateStripe(MEMORY_VIEWS, allocation); @@ -8631,7 +8643,8 @@ private static Object[] createAtomicStripes() { return stripes; } - /** Maps aliases of a Rust atomic location to the same striped monitor. */ + /** Lock the allocation during atomic access and aggregate decoding or publication. + * A decoded view can overlap several atomic fields, so offset locks are insufficient. */ private static Object atomicStripe(Pointer pointer) { return ATOMIC_STRIPES[atomicStripeIndex(pointer)]; } @@ -8645,20 +8658,25 @@ static long atomicAddress(Pointer pointer) { } private static int atomicStripeIndex(Pointer pointer) { - Object identity = pointer.allocation; - int memberHash = 0; + return atomicStripeIndex(pointer.allocation, pointer.exposedAddress); + } + + private static Object atomicStripe(Object root, long offset) { + return root instanceof Pointer + ? ATOMIC_STRIPES[atomicStripeIndex(((Pointer) root).allocation, ((Pointer) root).exposedAddress + offset)] + : ATOMIC_STRIPES[atomicStripeIndex(root, offset)]; + } + + private static int atomicStripeIndex(Object identity, long exposedAddress) { if (identity instanceof ReceiverCell) { identity = ((ReceiverCell) identity).value; } else if (identity instanceof FieldCell) { FieldCell cell = (FieldCell) identity; identity = cell.owner(); - memberHash = cell.fieldNameHash; } long key = identity == null - ? pointer.exposedAddress - : ((long) System.identityHashCode(identity) << 32) - ^ ((long) memberHash << 1) - ^ pointer.byteOffset; + ? exposedAddress + : System.identityHashCode(identity); key ^= key >>> 33; key *= 0xff51afd7ed558ccdL; key ^= key >>> 33; @@ -8684,49 +8702,109 @@ private static long signExtendAtomic(long value, int byteCount) { return (value << shift) >> shift; } - private static void atomicStoreLocked(Pointer pointer, long value, int byteCount) { - pointer.prepareMemoryWrite(pointer.byteOffset, byteCount); - pointer.storeBytes(truncateAtomic(value, byteCount), byteCount); + public static long atomicLoad(Pointer pointer, int byteCount, int ordering) { + return atomicLoad(pointer, 0L, byteCount, ordering); + } + + public static void atomicStore(Pointer pointer, long value, int byteCount, int ordering) { + atomicStore(pointer, 0L, value, byteCount, ordering); } - private static long atomicLoadStriped(Pointer pointer, int byteCount) { - synchronized (atomicStripe(pointer)) { - long value = truncateAtomic(pointer.loadUnsigned(byteCount), byteCount); - return value; + public static long atomicExchange(Pointer pointer, long value, int byteCount, int ordering) { + return atomicExchange(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicAdd(Pointer pointer, long value, int byteCount, int ordering) { + return atomicAdd(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicSubtract(Pointer pointer, long value, int byteCount, int ordering) { + return atomicSubtract(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicAnd(Pointer pointer, long value, int byteCount, int ordering) { + return atomicAnd(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicNand(Pointer pointer, long value, int byteCount, int ordering) { + return atomicNand(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicOr(Pointer pointer, long value, int byteCount, int ordering) { + return atomicOr(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicXor(Pointer pointer, long value, int byteCount, int ordering) { + return atomicXor(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicMax(Pointer pointer, long value, int byteCount, int ordering) { + return atomicMax(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicMin(Pointer pointer, long value, int byteCount, int ordering) { + return atomicMin(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicUnsignedMax(Pointer pointer, long value, int byteCount, int ordering) { + return atomicUnsignedMax(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicUnsignedMin(Pointer pointer, long value, int byteCount, int ordering) { + return atomicUnsignedMin(pointer, 0L, value, byteCount, ordering); + } + + public static long atomicCompareExchange(Pointer pointer, long expected, + long value, + int byteCount, + int successOrdering, + int failureOrdering) { + return atomicCompareExchange(pointer, 0L, expected, value, byteCount, successOrdering, failureOrdering); + } + + private static void atomicStoreLocked(Object root, long offset, long value, int byteCount) { + storeLocationBits(root, offset, truncateAtomic(value, byteCount), byteCount); + } + + private static long atomicLoadStriped(Object root, long offset, int byteCount) { + synchronized (atomicStripe(root, offset)) { + return truncateAtomic(loadLocationBits(root, offset, byteCount), byteCount); } } - public static long atomicLoad(Pointer pointer, int byteCount, int ordering) { + public static long atomicLoad(Object root, long offset, int byteCount, int ordering) { + root = normalizeLocationOrigin(root); checkedAtomicByteCount(byteCount); if (isSequentiallyConsistent(ordering)) { synchronized (ATOMIC_SEQUENCE_LOCK) { - return atomicLoadStriped(pointer, byteCount); + return atomicLoadStriped(root, offset, byteCount); } } - return atomicLoadStriped(pointer, byteCount); + return atomicLoadStriped(root, offset, byteCount); } - private static void atomicStoreStriped(Pointer pointer, long value, int byteCount) { - synchronized (atomicStripe(pointer)) { - atomicStoreLocked(pointer, value, byteCount); + private static void atomicStoreStriped(Object root, long offset, long value, int byteCount) { + synchronized (atomicStripe(root, offset)) { + atomicStoreLocked(root, offset, value, byteCount); } } - public static void atomicStore(Pointer pointer, long value, int byteCount, int ordering) { + public static void atomicStore(Object root, long offset, long value, int byteCount, int ordering) { + root = normalizeLocationOrigin(root); checkedAtomicByteCount(byteCount); if (isSequentiallyConsistent(ordering)) { synchronized (ATOMIC_SEQUENCE_LOCK) { - atomicStoreStriped(pointer, value, byteCount); + atomicStoreStriped(root, offset, value, byteCount); } return; } - atomicStoreStriped(pointer, value, byteCount); + atomicStoreStriped(root, offset, value, byteCount); } private static long atomicRmwStriped( - Pointer pointer, long operand, int byteCount, int operation) { - synchronized (atomicStripe(pointer)) { - long oldValue = truncateAtomic(pointer.loadUnsigned(byteCount), byteCount); + Object root, long offset, long operand, int byteCount, int operation) { + synchronized (atomicStripe(root, offset)) { + long oldValue = truncateAtomic(loadLocationBits(root, offset, byteCount), byteCount); long right = truncateAtomic(operand, byteCount); long newValue; switch (operation) { @@ -8772,97 +8850,99 @@ private static long atomicRmwStriped( default: throw new IllegalArgumentException("unknown Rust atomic operation " + operation); } - atomicStoreLocked(pointer, newValue, byteCount); + atomicStoreLocked(root, offset, newValue, byteCount); return oldValue; } } private static long atomicRmw( - Pointer pointer, long operand, int byteCount, int operation, int ordering) { + Object root, long offset, long operand, int byteCount, int operation, int ordering) { + root = normalizeLocationOrigin(root); checkedAtomicByteCount(byteCount); if (isSequentiallyConsistent(ordering)) { synchronized (ATOMIC_SEQUENCE_LOCK) { - return atomicRmwStriped(pointer, operand, byteCount, operation); + return atomicRmwStriped(root, offset, operand, byteCount, operation); } } - return atomicRmwStriped(pointer, operand, byteCount, operation); + return atomicRmwStriped(root, offset, operand, byteCount, operation); } public static long atomicExchange( - Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 0, ordering); + Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 0, ordering); } - public static long atomicAdd(Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 1, ordering); + public static long atomicAdd(Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 1, ordering); } public static long atomicSubtract( - Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 2, ordering); + Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 2, ordering); } - public static long atomicAnd(Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 3, ordering); + public static long atomicAnd(Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 3, ordering); } - public static long atomicNand(Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 4, ordering); + public static long atomicNand(Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 4, ordering); } - public static long atomicOr(Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 5, ordering); + public static long atomicOr(Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 5, ordering); } - public static long atomicXor(Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 6, ordering); + public static long atomicXor(Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 6, ordering); } - public static long atomicMax(Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 7, ordering); + public static long atomicMax(Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 7, ordering); } - public static long atomicMin(Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 8, ordering); + public static long atomicMin(Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 8, ordering); } public static long atomicUnsignedMax( - Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 9, ordering); + Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 9, ordering); } public static long atomicUnsignedMin( - Pointer pointer, long value, int byteCount, int ordering) { - return atomicRmw(pointer, value, byteCount, 10, ordering); + Object root, long offset, long value, int byteCount, int ordering) { + return atomicRmw(root, offset, value, byteCount, 10, ordering); } private static long atomicCompareExchangeStriped( - Pointer pointer, long expected, long value, int byteCount) { - synchronized (atomicStripe(pointer)) { - long oldValue = truncateAtomic(pointer.loadUnsigned(byteCount), byteCount); + Object root, long offset, long expected, long value, int byteCount) { + synchronized (atomicStripe(root, offset)) { + long oldValue = truncateAtomic(loadLocationBits(root, offset, byteCount), byteCount); if (oldValue == truncateAtomic(expected, byteCount)) { - atomicStoreLocked(pointer, value, byteCount); + atomicStoreLocked(root, offset, value, byteCount); } return oldValue; } } public static long atomicCompareExchange( - Pointer pointer, + Object root, long offset, long expected, long value, int byteCount, int successOrdering, int failureOrdering) { + root = normalizeLocationOrigin(root); checkedAtomicByteCount(byteCount); boolean sequentiallyConsistent = isSequentiallyConsistent(successOrdering) | isSequentiallyConsistent(failureOrdering); if (sequentiallyConsistent) { synchronized (ATOMIC_SEQUENCE_LOCK) { - return atomicCompareExchangeStriped(pointer, expected, value, byteCount); + return atomicCompareExchangeStriped(root, offset, expected, value, byteCount); } } - return atomicCompareExchangeStriped(pointer, expected, value, byteCount); + return atomicCompareExchangeStriped(root, offset, expected, value, byteCount); } public static void atomicFence(int ordering) { diff --git a/tests/integration/pointer_provenance/AtomicLocations.java b/tests/integration/pointer_provenance/AtomicLocations.java new file mode 100644 index 00000000..5f105e4c --- /dev/null +++ b/tests/integration/pointer_provenance/AtomicLocations.java @@ -0,0 +1,79 @@ +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.atomic.AtomicReference; +import org.rustlang.runtime.Pointer; + +/** The carrier and component APIs must use the same allocation monitor. */ +public final class AtomicLocations { + public static final class Counter { public long value; } + + public static void check() throws Exception { + for (int width : new int[] {1, 2, 4, 8}) { + byte[] bytes = new byte[24]; + Pointer pointer = Pointer.array(bytes, 8, 1).retype(width, null); + Pointer.atomicStore(bytes, 8L, 17, width, 4); + if (Pointer.atomicAdd(pointer, 3, width, 0) != 17 + || Pointer.atomicSubtract(bytes, 8L, 2, width, 0) != 20 + || Pointer.atomicExchange(bytes, 8L, 7, width, 0) != 18 + || Pointer.atomicAnd(bytes, 8L, 6, width, 0) != 7 + || Pointer.atomicOr(bytes, 8L, 9, width, 0) != 6 + || Pointer.atomicXor(bytes, 8L, 3, width, 0) != 15 + || Pointer.atomicNand(bytes, 8L, 10, width, 0) != 12) { + throw new AssertionError("atomic component RMW"); + } + long mask = width == 8 ? -1L : (1L << (width * 8)) - 1; + if (Pointer.atomicLoad(pointer, width, 0) != (~8L & mask) + || Pointer.atomicCompareExchange(bytes, 8L, 1, 2, width, 0, 0) != (~8L & mask) + || Pointer.atomicCompareExchange(pointer, ~8L, 33, width, 4, 0) != (~8L & mask) + || Pointer.atomicLoad(bytes, 8L, width, 0) != 33) { + throw new AssertionError("atomic compare exchange or truncation"); + } + Pointer.atomicStore(pointer, -2, width, 0); + if (Pointer.atomicMax(bytes, 8L, 1, width, 0) != (-2L & mask) + || Pointer.atomicMin(bytes, 8L, -3, width, 0) != 1 + || Pointer.atomicUnsignedMin(bytes, 8L, 3, width, 0) != (-3L & mask) + || Pointer.atomicUnsignedMax(bytes, 8L, -4, width, 0) != 3 + || Pointer.atomicLoad(pointer, width, 0) != (-4L & mask)) { + throw new AssertionError("atomic signed comparison"); + } + } + byte[] bytes = new byte[24]; + mixed(bytes, 8, Pointer.array(bytes, 8, 1).retype(8, null)); + long[] words = new long[3]; + mixed(words, 8, Pointer.array(words, 1, 8)); + Object local = Pointer.storageAligned(0L, 8, null, 8); + mixed(local, 0, Pointer.fromLocation(local, 0, 8)); + Pointer field = Pointer.field(new Counter(), "value", 8, null); + mixed(field, 0, field.retype(8, null)); + } + + private static void mixed(Object root, long offset, Pointer pointer) throws Exception { + Pointer.atomicStore(root, offset, 0, 8, 0); + CountDownLatch start = new CountDownLatch(1); + AtomicReference failure = new AtomicReference<>(); + Thread[] workers = new Thread[4]; + for (int i = 0; i < workers.length; i++) { + final boolean component = (i & 1) != 0; + final int ordering = i < 2 ? 0 : 4; + workers[i] = new Thread(() -> { + try { + start.await(); + for (int n = 0; n < 10000; n++) { + if (component) Pointer.atomicAdd(root, offset, 1, 8, ordering); + else Pointer.atomicAdd(pointer, 1, 8, ordering); + } + } catch (Throwable error) { failure.compareAndSet(null, error); } + }); + workers[i].start(); + } + start.countDown(); + for (Thread worker : workers) { + worker.join(10000); + if (worker.isAlive()) throw new AssertionError("atomic worker stalled"); + } + if (failure.get() != null) throw new AssertionError(failure.get()); + if (Pointer.atomicLoad(root, offset, 8, 4) != 40000 + || Pointer.atomicLoad(pointer, 8, 0) != 40000) { + throw new AssertionError("component and carrier atomics did not synchronize"); + } + } +} diff --git a/tests/integration/pointer_provenance/AtomicViews.java b/tests/integration/pointer_provenance/AtomicViews.java new file mode 100644 index 00000000..8c636431 --- /dev/null +++ b/tests/integration/pointer_provenance/AtomicViews.java @@ -0,0 +1,92 @@ +import java.util.concurrent.CountDownLatch; +import java.util.concurrent.TimeUnit; +import java.util.concurrent.atomic.AtomicReference; +import org.rustlang.runtime.MemoryBytes; +import org.rustlang.runtime.Pointer; + +/** An aggregate decode must not publish bytes superseded by an interior atomic write. */ +public final class AtomicViews { + public static final class Pair { + public int payload; + public int state; + } + + public static final class Codec { + static volatile CountDownLatch decoding; + static volatile CountDownLatch resume; + + public static byte[] e$pair(Pair value) { + byte[] bytes = new byte[8]; + MemoryBytes.write(bytes, 0, 4, value.payload); + MemoryBytes.write(bytes, 4, 4, value.state); + return bytes; + } + + public static Pair d$pair(byte[] bytes) { + Pair value = new Pair(); + value.payload = (int) MemoryBytes.read(bytes, 0, 4); + value.state = (int) MemoryBytes.read(bytes, 4, 4); + CountDownLatch paused = decoding; + if (paused != null) { + paused.countDown(); + await(resume); + } + return value; + } + } + + private static void await(CountDownLatch latch) { + try { + if (!latch.await(10, TimeUnit.SECONDS)) throw new AssertionError("worker stalled"); + } catch (InterruptedException error) { + throw new AssertionError(error); + } + } + + public static void check() throws Exception { + check(false); + check(true); + } + + private static void check(boolean components) throws Exception { + byte[] storage = new byte[8]; + MemoryBytes.write(storage, 4, 4, 1); + Pointer whole = Pointer.array(storage, 0, 1).retype(8, "AtomicViews$Codec#pair#LAtomicViews$Pair;#8"); + Pointer atomic = Pointer.array(storage, 4, 1).retype(4, null); + AtomicReference decoded = new AtomicReference<>(); + AtomicReference failed = new AtomicReference<>(); + Codec.decoding = new CountDownLatch(1); + Codec.resume = new CountDownLatch(1); + Thread reader = new Thread(() -> { + try { decoded.set((Pair) whole.getObject()); } + catch (Throwable error) { failed.set(error); } + }); + Thread writer = new Thread(() -> { + try { + if (components) Pointer.atomicStore(storage, 4L, 0, 4, 0); + else Pointer.atomicStore(atomic, 0, 4, 0); + } + catch (Throwable error) { failed.set(error); } + }); + reader.start(); + await(Codec.decoding); + writer.start(); + // Wait until the writer finishes or blocks on the overlapping decode. + long deadline = System.nanoTime() + TimeUnit.SECONDS.toNanos(10); + while (writer.isAlive() && writer.getState() != Thread.State.BLOCKED) { + if (System.nanoTime() > deadline) throw new AssertionError("writer stalled"); + Thread.yield(); + } + Codec.resume.countDown(); + reader.join(10000); + writer.join(10000); + Codec.decoding = null; + if (reader.isAlive() || writer.isAlive()) throw new AssertionError("workers stalled"); + if (failed.get() != null) throw new AssertionError(failed.get()); + decoded.get().payload = 73; + // Flushing the aggregate must not restore stale state over the atomic store. + if (Pointer.atomicLoad(atomic, 4, 0) != 0) { + throw new AssertionError("decoded aggregate restored a completed atomic state"); + } + } +} diff --git a/tests/integration/pointer_provenance/Main.java b/tests/integration/pointer_provenance/Main.java index 801b772d..2ca970cb 100644 --- a/tests/integration/pointer_provenance/Main.java +++ b/tests/integration/pointer_provenance/Main.java @@ -7,6 +7,8 @@ public static void main(String[] args) throws Exception { ArrayViews.check(); CellArrayViews.check(); MemoryViews.check(); + AtomicViews.check(); + AtomicLocations.check(); RangeCodecs.check(); PrimitiveArrayCodecs.check(); TypedFields.check(); @@ -18,6 +20,7 @@ public static void main(String[] args) throws Exception { BorrowedFields.check(); BorrowedLocals.check(); CodecAdapters.check(); + MetadataFilters.check(); StructuralViews.check(); Field field = Pointer.class.getDeclaredField("EXPOSED_ADDRESSES"); field.setAccessible(true); diff --git a/tests/integration/pointer_provenance/MetadataFilters.java b/tests/integration/pointer_provenance/MetadataFilters.java new file mode 100644 index 00000000..3f1d673c --- /dev/null +++ b/tests/integration/pointer_provenance/MetadataFilters.java @@ -0,0 +1,77 @@ +import java.lang.reflect.Field; +import java.lang.reflect.Method; +import java.util.Map; +import java.util.concurrent.atomic.AtomicLong; +import java.util.concurrent.atomic.AtomicReference; +import org.rustlang.runtime.Pointer; + +/** A cache rebuild must not hide a newly published allocation dependency. */ +final class MetadataFilters { + private static Object field(Object owner, String name) throws Exception { + Class type = owner instanceof Class ? (Class) owner : owner.getClass(); + Field field = type.getDeclaredField(name); + field.setAccessible(true); + return field.get(owner instanceof Class ? null : owner); + } + + private static Method method(String name, Class... arguments) throws Exception { + Method method = Pointer.class.getDeclaredMethod(name, arguments); + method.setAccessible(true); + return method; + } + + @SuppressWarnings("unchecked") + static void check() throws Exception { + locationOriginFalsePositive(); + Object source = new byte[16], target = new byte[16], copied = new byte[16]; + Map[] stripes = (Map[]) field(Pointer.class, "ENCODED_REFERENCES"); + Method stripeIndex = method("stateStripeIndex", Object.class); + Map sourceStripe = stripes[(int) stripeIndex.invoke(null, source)]; + Method retain = method("retainEncodedReference", Object.class, Object.class); + AtomicReference failure = new AtomicReference<>(); + Thread writer = new Thread(() -> { + try { retain.invoke(null, source, target); } + catch (Throwable error) { failure.set(error); } + }); + synchronized (sourceStripe) { + writer.start(); + // Rebuild after the writer marks the old filter but before it publishes the entry. + long deadline = System.nanoTime() + 5_000_000_000L; + while (writer.getState() != Thread.State.BLOCKED) { + if (failure.get() != null || !writer.isAlive() || System.nanoTime() >= deadline) + throw new AssertionError("writer did not reach metadata publication", failure.get()); + Thread.onSpinWait(); + } + Object filter = field(Pointer.class, "ENCODED_REFERENCE_FILTER"); + ((AtomicLong) field(filter, "marks")).set(Long.MAX_VALUE); + method("maybeRebuildIdentityFilter", filter.getClass(), Map[].class).invoke(null, filter, stripes); + } + writer.join(5000); + if (writer.isAlive() || failure.get() != null) + throw new AssertionError("metadata publication failed", failure.get()); + method("transferEncodedReferences", Object.class, Object.class).invoke(null, source, copied); + Map copyStripe = stripes[(int) stripeIndex.invoke(null, copied)]; + synchronized (copyStripe) { + if (copyStripe.get(copied) != target) + throw new AssertionError("filter rebuild lost the copied allocation's strong dependency"); + } + method("discardEncodedReferences", Object.class).invoke(null, source); + method("discardEncodedReferences", Object.class).invoke(null, copied); + } + + private static void locationOriginFalsePositive() throws Exception { + long[] values = new long[] {17, 23}; + Object filter = field(Pointer.class, "MEMORY_VIEW_ORIGIN_FILTER"); + method("markIdentityFilter", filter.getClass(), Object.class).invoke(null, filter, values); + Object normalized = method("normalizeLocationOrigin", Object.class).invoke(null, values); + if (normalized != values) throw new AssertionError("false origin mark allocated an address"); + long address = Pointer.locationAddr(values, 8); + Pointer pointer = Pointer.array(values, 1, 8); + if (pointer.addr() != address || Pointer.byteOffsetLocations(values, 8, pointer, 0) != 0 + || Pointer.atomicAdd(values, 8L, 5, 8, 0) != 23 + || Pointer.atomicLoad(pointer, 8, 4) != 28) { + throw new AssertionError("false origin mark changed address or atomic identity"); + } + } + +} From 2d1ca2ac8b609041174db7e55f87e67da181b9ed Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 11:39:11 +1000 Subject: [PATCH 09/61] avoid slice address wrappers --- runtime/src/Pointer.java | 367 +++++++++++++++--- .../integration/pointer_provenance/Main.java | 2 + .../pointer_provenance/MemoryViews.java | 350 +++++++++++++++++ .../pointer_provenance/SliceAliases.java | 59 +++ .../pointer_provenance/SliceLocations.java | 98 +++++ 5 files changed, 821 insertions(+), 55 deletions(-) create mode 100644 tests/integration/pointer_provenance/SliceAliases.java create mode 100644 tests/integration/pointer_provenance/SliceLocations.java diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index 3904495c..4c1882c1 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -1145,9 +1145,13 @@ private MemoryViewOrigin(Pointer pointer) { } private boolean matches(Pointer pointer) { + return matches(pointer, pointer.byteOffset); + } + + private boolean matches(Pointer pointer, long offset) { return allocation.get() == pointer.allocation && allocationElementSize == pointer.allocationElementSize - && byteOffset == pointer.byteOffset + && byteOffset == offset && viewSize == pointer.viewSize && java.util.Objects.equals( allocationCodecClassName, pointer.allocationCodecClassName) @@ -1496,6 +1500,71 @@ public static void fillArray(Object array, Object value, boolean copyValue) { } } + // Primitive repeat initialization keeps values unboxed and preserves live aliases. + public static void fillArray(boolean[] array, boolean value) { + if (hasScalarWriteTracking(array)) { + for (int index = 0; index < array.length; index++) sliceSetBoolean(array, index, value); + } else { + Arrays.fill(array, value); + } + } + + public static void fillArray(byte[] array, byte value) { + if (hasScalarWriteTracking(array)) { + for (int index = 0; index < array.length; index++) sliceSetI8(array, index, value); + } else { + Arrays.fill(array, value); + } + } + + public static void fillArray(short[] array, short value) { + if (hasScalarWriteTracking(array)) { + for (int index = 0; index < array.length; index++) sliceSetI16(array, index, value); + } else { + Arrays.fill(array, value); + } + } + + public static void fillArray(char[] array, char value) { + if (hasScalarWriteTracking(array)) { + for (int index = 0; index < array.length; index++) sliceSetU16(array, index, value); + } else { + Arrays.fill(array, value); + } + } + + public static void fillArray(int[] array, int value) { + if (hasScalarWriteTracking(array)) { + for (int index = 0; index < array.length; index++) sliceSetI32(array, index, value); + } else { + Arrays.fill(array, value); + } + } + + public static void fillArray(long[] array, long value) { + if (hasScalarWriteTracking(array)) { + for (int index = 0; index < array.length; index++) sliceSetI64(array, index, value); + } else { + Arrays.fill(array, value); + } + } + + public static void fillArray(float[] array, float value) { + if (hasScalarWriteTracking(array)) { + for (int index = 0; index < array.length; index++) sliceSetF32(array, index, value); + } else { + Arrays.fill(array, value); + } + } + + public static void fillArray(double[] array, double value) { + if (hasScalarWriteTracking(array)) { + for (int index = 0; index < array.length; index++) sliceSetF64(array, index, value); + } else { + Arrays.fill(array, value); + } + } + private static Object independentRepeatedArrayElement(Object array, int index) { int length = Array.getLength(array); if (index < 0 || index >= length) { @@ -2847,6 +2916,30 @@ private Object decodedMemoryView() { } } + /** Reuse a decoded element only while its registered origin still matches. + * Keep the enclosing pointer bound to its own element. */ + private Object cachedSliceElement(long displacement) { + if (!(allocation instanceof byte[]) || viewSize <= 0 + || !isGeneratedAggregateCodec(viewCodecClassName) + || traitObjectCarrier() != null || traitMetadataCarrier() != null + || isDirectAllocationView()) return null; + long absolute = Math.addExact(byteOffset, displacement); + synchronized (atomicStripe(this)) { + Map> stripe = stateStripe(MEMORY_VIEWS, allocation); + synchronized (stripe) { + LongRangeMap views = stripe.get(allocation); + MemoryViewState cached = views == null ? null : views.get(absolute); + if (cached == null || !cached.active || cached.size != viewSize + || !cached.codecClassName.equals(viewCodecClassName) || cached.value == null) return null; + Map origins = stateStripe(MEMORY_VIEW_ORIGINS, cached.value); + synchronized (origins) { + MemoryViewOrigin origin = origins.get(cached.value); + return origin != null && origin.matches(this, absolute) ? cached.value : null; + } + } + } + } + private Object activeBoundMemoryViewValue(int materializedSize) { MemoryViewState bound = boundMemoryViewState(); if (bound != null @@ -9632,6 +9725,93 @@ public static Pointer fromTypedStorageLocation(Object root, long offset, int siz return result; } + private static Object directLocationSliceArray(Object root, int size) { + if (size <= 0) return null; + if (root instanceof Pointer) { + Pointer pointer = (Pointer) root; + if (pointer.rareState != null || pointer.allocationElementSize != size + || pointer.allocationCodecClassName != null) return null; + root = pointer.allocation; + } else if (mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, root)) { + return null; + } + return root != null && root.getClass().isArray() + && root.getClass().getComponentType().isPrimitive() + && inferredArrayElementSize(root) == size ? root : null; + } + + /** Re-form a scalar slice directly from its decomposed storage location. */ + public static Object locationSliceBacking(Object root, long offset, int size) { + if (root == null && offset == 0) return null; + Object array = directLocationSliceArray(root, size); + if (array != null) return array; + if (retainedSliceRoot(root, offset, size, null)) return root; + return fromLocation(root, offset, size).sliceBackingArray(); + } + + public static int locationSliceOffset(Object root, long offset, int size) { + if (root == null && offset == 0) return 0; + if (directLocationSliceArray(root, size) == null) { + if (retainedSliceRoot(root, offset, size, null)) return (int) (offset / size); + return fromLocation(root, offset, size).sliceElementOffset(); + } + if (root instanceof Pointer) offset += ((Pointer) root).byteOffset; + if (offset % size != 0) { + throw new IllegalStateException("slice data pointer is not element-aligned"); + } + return Math.toIntExact(offset / size); + } + + private static boolean retainedSliceRoot(Object root, long offset, int size, String codec) { + if (!(root instanceof Pointer) || size <= 0 || offset % size != 0 + || offset / size != (int) (offset / size)) return false; + Pointer pointer = (Pointer) root; + return pointer.viewSize == size && java.util.Objects.equals(pointer.viewCodecClassName, codec); + } + + /** Retain the element owner. The slice start is relative to its current displacement. */ + public static Object typedLocationSliceBacking(Object root, long offset, int size, String codec) { + if (retainedSliceRoot(root, offset, size, codec)) return root; + return fromTypedStorageLocation(root, offset, size, codec).sliceBackingArray(); + } + + public static int typedLocationSliceOffset(Object root, long offset, int size, String codec) { + if (retainedSliceRoot(root, offset, size, codec)) return (int) (offset / size); + return fromTypedStorageLocation(root, offset, size, codec).sliceElementOffset(); + } + + /** Normalize the aggregate slice root once. Keep byte offsets and slice length separate. */ + public static Object sliceAddressRoot(Object backing, int size, String codec) { + if (backing == null) return withoutProvenance(0, size, codec); + if (backing instanceof Pointer) { + return ((Pointer) backing).sliceStorageView(size, codec); + } + // Preserve the backing owner of decoded arrays with a registered origin. + return fromSliceParts(backing, 0, 0L, size, codec); + } + + /** Normalize unusual decoded/transparent storage once before scalar iteration. */ + public static Object scalarSliceRoot(Object backing, int size) { + if (backing == null) return null; + if (backing.getClass().isArray() + && backing.getClass().getComponentType().isPrimitive() + && !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, backing)) { + return backing; + } + if (backing instanceof Pointer) { + Pointer pointer = (Pointer) backing; + if (pointer.allocation == null + || (pointer.allocation.getClass().isArray() + && pointer.allocation.getClass().getComponentType().isPrimitive())) { + return pointer; + } + // Root normalization must not change slice metadata. + return pointer.sliceStorageView(size, null); + } + return fromSliceParts(backing, 0, 0L, size, null); + } + + /** Byte-backed places can access a known scalar offset without a field view. */ public static boolean hasByteStorage(Object root) { return root instanceof byte[] || (root instanceof Pointer && ((Pointer) root).allocation instanceof byte[]); @@ -10866,6 +11046,21 @@ private static Pointer slicePointer(Object backing, int index) { : null; } + private long sliceScalarBits(int index, int size, String name) { + Pointer storage = sliceElementView(); + storage.requireScalarViewSize(size, name); + return loadLocationBits(storage, Math.multiplyExact((long) index, size), size); + } + + private static boolean storeSliceScalar(Object backing, int index, long bits, int size) { + if (!(backing instanceof Pointer)) return false; + Pointer storage = ((Pointer) backing).sliceElementView(); + // Aggregate/encoded layouts retain set()'s conversion and ownership rules. + if (storage.viewSize != size || storage.viewCodecClassName != null) return false; + storeLocationBits(storage, Math.multiplyExact((long) index, size), bits, size); + return true; + } + /** * Finds a slice carried through transparent single-field Rust wrappers, * such as {@code UnsafeCell<[T]>}. @@ -10946,6 +11141,12 @@ private Pointer sliceStorageView(long elementSize, String elementCodecClassName) String allocationCodec = storage.allocationCodecClassName != null ? storage.allocationCodecClassName : elementCodecClassName; + if (checkedElementSize > 0 && storage.allocationElementSize == checkedElementSize + && storage.viewSize == checkedElementSize + && java.util.Objects.equals(storage.viewCodecClassName, elementCodecClassName) + && java.util.Objects.equals(storage.allocationCodecClassName, allocationCodec)) { + return storage; + } if (checkedElementSize == 0) { return new Pointer( storage.allocation, @@ -10995,24 +11196,31 @@ private static boolean hasScalarWriteTracking(Object root) { public static boolean sliceGetBoolean(Object backing, int index) { if (backing instanceof boolean[]) { - return ((boolean[]) backing)[index]; + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return ((boolean[]) backing)[index]; + } + return loadLocationBits(backing, index, 1) != 0; } - Pointer pointer = slicePointer(backing, index); - return pointer != null - ? pointer.getBoolean() + return backing instanceof Pointer + ? ((Pointer) backing).sliceScalarBits(index, 1, "bool") != 0 : ((Boolean) arrayGet(backing, index)).booleanValue(); } public static byte sliceGetI8(Object backing, int index) { if (backing instanceof byte[]) { - return ((byte[]) backing)[index]; + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return ((byte[]) backing)[index]; + } + return (byte) loadLocationBits(backing, index, 1); } if (backing instanceof boolean[]) { - return (byte) (((boolean[]) backing)[index] ? 1 : 0); + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return (byte) (((boolean[]) backing)[index] ? 1 : 0); + } + return (byte) loadLocationBits(backing, index, 1); } - Pointer pointer = slicePointer(backing, index); - if (pointer != null) { - return pointer.getI8(); + if (backing instanceof Pointer) { + return (byte) ((Pointer) backing).sliceScalarBits(index, 1, "i8/u8"); } Object value = arrayGet(backing, index); return value instanceof Boolean @@ -11031,24 +11239,31 @@ public static byte[] sliceToByteArray(Object backing, int offset, int length) { public static short sliceGetI16(Object backing, int index) { if (backing instanceof short[]) { - return ((short[]) backing)[index]; + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return ((short[]) backing)[index]; + } + return (short) loadLocationBits(backing, (long) index * 2, 2); } - Pointer pointer = slicePointer(backing, index); - return pointer != null - ? pointer.getI16() + return backing instanceof Pointer + ? (short) ((Pointer) backing).sliceScalarBits(index, 2, "i16/u16") : ((Number) arrayGet(backing, index)).shortValue(); } public static char sliceGetU16(Object backing, int index) { if (backing instanceof char[]) { - return ((char[]) backing)[index]; + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return ((char[]) backing)[index]; + } + return (char) loadLocationBits(backing, (long) index * 2, 2); } if (backing instanceof short[]) { - return (char) (((short[]) backing)[index] & 0xffff); + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return (char) (((short[]) backing)[index] & 0xffff); + } + return (char) loadLocationBits(backing, (long) index * 2, 2); } - Pointer pointer = slicePointer(backing, index); - if (pointer != null) { - return (char) (pointer.getI16() & 0xffff); + if (backing instanceof Pointer) { + return (char) ((Pointer) backing).sliceScalarBits(index, 2, "i16/u16"); } Object value = arrayGet(backing, index); return value instanceof Character @@ -11058,14 +11273,21 @@ public static char sliceGetU16(Object backing, int index) { public static int sliceGetI32(Object backing, int index) { if (backing instanceof int[]) { - return ((int[]) backing)[index]; + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return ((int[]) backing)[index]; + } + return (int) loadLocationBits(backing, (long) index * 4, 4); } if (backing instanceof char[]) { - return ((char[]) backing)[index]; + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return ((char[]) backing)[index]; + } + return (char) loadLocationBits(backing, (long) index * 2, 2); } - Pointer pointer = slicePointer(backing, index); - Object value = - pointer != null ? Integer.valueOf(pointer.getI32()) : arrayGet(backing, index); + if (backing instanceof Pointer) { + return (int) ((Pointer) backing).sliceScalarBits(index, 4, "i32/u32"); + } + Object value = arrayGet(backing, index); return value instanceof Character ? ((Character) value).charValue() : ((Number) value).intValue(); @@ -11073,31 +11295,37 @@ public static int sliceGetI32(Object backing, int index) { public static long sliceGetI64(Object backing, int index) { if (backing instanceof long[]) { - return ((long[]) backing)[index]; + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return ((long[]) backing)[index]; + } + return loadLocationBits(backing, (long) index * 8, 8); } - Pointer pointer = slicePointer(backing, index); - return pointer != null - ? pointer.getI64() + return backing instanceof Pointer + ? ((Pointer) backing).sliceScalarBits(index, 8, "i64/u64") : ((Number) arrayGet(backing, index)).longValue(); } public static float sliceGetF32(Object backing, int index) { if (backing instanceof float[]) { - return ((float[]) backing)[index]; + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return ((float[]) backing)[index]; + } + return Float.intBitsToFloat((int) loadLocationBits(backing, (long) index * 4, 4)); } - Pointer pointer = slicePointer(backing, index); - return pointer != null - ? pointer.getF32() + return backing instanceof Pointer + ? Float.intBitsToFloat((int) ((Pointer) backing).sliceScalarBits(index, 4, "f32")) : ((Number) arrayGet(backing, index)).floatValue(); } public static double sliceGetF64(Object backing, int index) { if (backing instanceof double[]) { - return ((double[]) backing)[index]; + if (!mayBeInIdentityFilter(MEMORY_VIEW_FILTER, backing)) { + return ((double[]) backing)[index]; + } + return Double.longBitsToDouble(loadLocationBits(backing, (long) index * 8, 8)); } - Pointer pointer = slicePointer(backing, index); - return pointer != null - ? pointer.getF64() + return backing instanceof Pointer + ? Double.longBitsToDouble(((Pointer) backing).sliceScalarBits(index, 8, "f64")) : ((Number) arrayGet(backing, index)).doubleValue(); } @@ -11105,61 +11333,84 @@ public static Object sliceGetObject(Object backing, int index) { if (backing instanceof Object[]) { return independentRepeatedArrayElement(backing, index); } - Pointer pointer = slicePointer(backing, index); - return pointer != null - ? pointer.getObject() - : independentRepeatedArrayElement(backing, index); + if (backing instanceof Pointer) { + Pointer storage = ((Pointer) backing).sliceElementView(); + long displacement = Math.multiplyExact((long) index, storage.viewSize); + Object decoded = storage.cachedSliceElement(displacement); + if (decoded != null) return decoded; + if (storage.rareState == null && storage.addressState == null + && storage.isDirectAllocationView() && !mayHaveStructuralView(storage.allocation)) { + Object value = loadObjectLocation(storage, displacement, null); + // Array origins belong to the derived location, not this base. + if (value == null || !value.getClass().isArray()) return value; + } + return storage.add(index).getObject(); + } + return independentRepeatedArrayElement(backing, index); } public static void sliceSetBoolean(Object backing, int index, boolean value) { if (backing instanceof boolean[]) { - ((boolean[]) backing)[index] = value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, index, value ? 1 : 0, 1); + else ((boolean[]) backing)[index] = value; return; } + if (storeSliceScalar(backing, index, value ? 1 : 0, 1)) return; sliceSetObject(backing, index, Boolean.valueOf(value)); } public static void sliceSetI8(Object backing, int index, byte value) { if (backing instanceof byte[]) { - ((byte[]) backing)[index] = value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, index, value, 1); + else ((byte[]) backing)[index] = value; return; } if (backing instanceof boolean[]) { - ((boolean[]) backing)[index] = value != 0; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, index, value != 0 ? 1 : 0, 1); + else ((boolean[]) backing)[index] = value != 0; return; } + if (storeSliceScalar(backing, index, value, 1)) return; sliceSetObject(backing, index, Byte.valueOf(value)); } public static void sliceSetI16(Object backing, int index, short value) { if (backing instanceof short[]) { - ((short[]) backing)[index] = value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, (long) index * 2, value, 2); + else ((short[]) backing)[index] = value; return; } + if (storeSliceScalar(backing, index, value, 2)) return; sliceSetObject(backing, index, Short.valueOf(value)); } public static void sliceSetU16(Object backing, int index, char value) { if (backing instanceof char[]) { - ((char[]) backing)[index] = value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, (long) index * 2, value, 2); + else ((char[]) backing)[index] = value; return; } if (backing instanceof short[]) { - ((short[]) backing)[index] = (short) value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, (long) index * 2, value, 2); + else ((short[]) backing)[index] = (short) value; return; } + if (storeSliceScalar(backing, index, value, 2)) return; sliceSetObject(backing, index, Character.valueOf(value)); } public static void sliceSetI32(Object backing, int index, int value) { if (backing instanceof int[]) { - ((int[]) backing)[index] = value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, (long) index * 4, value, 4); + else ((int[]) backing)[index] = value; return; } if (backing instanceof char[]) { - ((char[]) backing)[index] = (char) value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, (long) index * 2, value, 2); + else ((char[]) backing)[index] = (char) value; return; } + if (storeSliceScalar(backing, index, value, 4)) return; Pointer pointer = slicePointer(backing, index); if (pointer != null) { pointer.set(Integer.valueOf(value)); @@ -11167,13 +11418,13 @@ public static void sliceSetI32(Object backing, int index, int value) { } Class component = backing.getClass().getComponentType(); if (component == byte.class) { - Array.setByte(backing, index, (byte) value); + storeLocationBits(backing, index, value, 1); } else if (component == short.class) { - Array.setShort(backing, index, (short) value); + storeLocationBits(backing, (long) index * 2, value, 2); } else if (component == char.class) { - Array.setChar(backing, index, (char) value); + storeLocationBits(backing, (long) index * 2, value, 2); } else if (component == int.class) { - Array.setInt(backing, index, value); + storeLocationBits(backing, (long) index * 4, value, 4); } else { arraySet(backing, index, Integer.valueOf(value)); } @@ -11181,25 +11432,31 @@ public static void sliceSetI32(Object backing, int index, int value) { public static void sliceSetI64(Object backing, int index, long value) { if (backing instanceof long[]) { - ((long[]) backing)[index] = value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, (long) index * 8, value, 8); + else ((long[]) backing)[index] = value; return; } + if (storeSliceScalar(backing, index, value, 8)) return; sliceSetObject(backing, index, Long.valueOf(value)); } public static void sliceSetF32(Object backing, int index, float value) { if (backing instanceof float[]) { - ((float[]) backing)[index] = value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, (long) index * 4, Float.floatToRawIntBits(value), 4); + else ((float[]) backing)[index] = value; return; } + if (storeSliceScalar(backing, index, Float.floatToRawIntBits(value), 4)) return; sliceSetObject(backing, index, Float.valueOf(value)); } public static void sliceSetF64(Object backing, int index, double value) { if (backing instanceof double[]) { - ((double[]) backing)[index] = value; + if (hasScalarWriteTracking(backing)) storeLocationBits(backing, (long) index * 8, Double.doubleToRawLongBits(value), 8); + else ((double[]) backing)[index] = value; return; } + if (storeSliceScalar(backing, index, Double.doubleToRawLongBits(value), 8)) return; sliceSetObject(backing, index, Double.valueOf(value)); } diff --git a/tests/integration/pointer_provenance/Main.java b/tests/integration/pointer_provenance/Main.java index 2ca970cb..6101832a 100644 --- a/tests/integration/pointer_provenance/Main.java +++ b/tests/integration/pointer_provenance/Main.java @@ -6,6 +6,7 @@ public class Main { public static void main(String[] args) throws Exception { ArrayViews.check(); CellArrayViews.check(); + SliceAliases.check(); MemoryViews.check(); AtomicViews.check(); AtomicLocations.check(); @@ -16,6 +17,7 @@ public static void main(String[] args) throws Exception { CyclicFieldViews.check(); AddressOrder.check(); LocationQueries.check(); + SliceLocations.check(); OwnedFields.check(); BorrowedFields.check(); BorrowedLocals.check(); diff --git a/tests/integration/pointer_provenance/MemoryViews.java b/tests/integration/pointer_provenance/MemoryViews.java index 5e057e8a..6c40f7d4 100644 --- a/tests/integration/pointer_provenance/MemoryViews.java +++ b/tests/integration/pointer_provenance/MemoryViews.java @@ -59,6 +59,54 @@ private static Pair view(Pointer bytes) { } public static void check() { + Object nullableOwner = Pointer.storageAligned(new Pair(3, 5), 8, PAIR_CODEC, 8); + if (Pointer.nullableLocationTag(nullableOwner, 0) != 1 + || Pointer.nullableLocationTag(new int[] {7}, 0) != 1 + || Pointer.nullableLocationTag(null, 0) != 0 + || Pointer.nullableLocationTag(null, 8) != 1 + || Pointer.nullableTag(new Pointer(0L, 4)) != 0 + || Pointer.nullableTag(new Pointer(4L, 4)) != 1 + || Pointer.nullableLocationTag(new Pointer(4L, 4), -4) != 0) { + throw new AssertionError("invalid nullable pointer tag"); + } + Pair erasedOwner = new Pair(3, 5); + Pointer.field(erasedOwner, "erasedMarker", 0, null); + Pointer erasedRoot = Pointer.cell(erasedOwner, 8, PAIR_CODEC); + Pointer erased = erasedRoot.projectStructField(Pair.class.getName(), "erasedMarker", 8, 0, null); + if (!Pointer.sameLocation(erased, 0, erasedRoot, 8)) { + throw new AssertionError("erased field lost its containing allocation"); + } + try { + Pointer.field(erasedOwner, "missingValue", 4, null); + throw new AssertionError("a missing non-ZST field was accepted"); + } catch (IllegalArgumentException expected) { } + checkTypedStorage(); + checkOwnedRanges(); + checkTypedCopies(); + checkPointerTypedStores(); + checkSliceComponents(); + byte[] delayedBytes = new byte[8]; + Pointer delayed = Pointer.array(delayedBytes, 0, 1).retype(8, PAIR_CODEC); + Pointer commit = Pointer.fromStorageLocation(delayed, 0); + Pair delayedView = (Pair) Pointer.loadStorageLocation(delayed, 0, Pair.class.getName()); + delayedView.first = 139; + commit.commitMemoryView(); + if (MemoryBytes.read(delayedBytes, 0, 4) != 139) { + throw new AssertionError("component boundary detached a later view binding"); + } + Pair firstPair = new Pair(); + Pair secondPair = new Pair(); + Pointer pairArray = Pointer.array(new Pair[] { firstPair, secondPair }, 0, 8, PAIR_CODEC); + if (Pointer.locationStride(pairArray) != 8 + || Pointer.directStorageAggregate(pairArray, 8, Pair.class) != secondPair + || Pointer.loadStorageLocation(pairArray, 8, Pair.class.getName()) != secondPair) { + throw new AssertionError("aggregate component address lost its element layout"); + } + Pointer pairCell = Pointer.cell(firstPair, 8, PAIR_CODEC); + Pointer.storeStorageLocation(pairCell, 0, secondPair); + if (Pointer.directStorageAggregate(pairCell, 0, Pair.class) != secondPair) { + throw new AssertionError("aggregate component address retained a replaced object"); + } byte[] storage = new byte[16]; Pointer bytes = Pointer.array(storage, 0, 1); Pair original = view(bytes); @@ -114,8 +162,58 @@ public static void check() { throw new AssertionError("plain aggregate read used the wrong source window"); } + byte[] plain = new byte[24]; + if (Pointer.directStorageAggregate(plain, 8, Pair.class) != null) { + throw new AssertionError("byte storage unexpectedly supplied a managed aggregate"); + } + // Byte-backed fields need no Java member. Test both root representations. + for (Object root : new Object[] {plain, Pointer.array(plain, 0, 1)}) { + Pointer.storeScalarField(root, 8, "absent.Owner", "absentField", 4, 97, 4); + if (Pointer.loadScalarField(root, 8, "absent.Owner", "absentField", 4, 4) != 97 + || MemoryBytes.read(plain, 8, 4) != 0 || MemoryBytes.read(plain, 16, 4) != 0) { + throw new AssertionError("scalar field fallback lost its exact byte window"); + } + } + Pointer.storeTypedStorage(plain, 8, 8, PAIR_CODEC, new Pair(107, 109)); + if (MemoryBytes.read(plain, 8, 4) != 107 || MemoryBytes.read(plain, 12, 4) != 109 + || MemoryBytes.read(plain, 4, 4) != 0 || MemoryBytes.read(plain, 16, 4) != 0) { + throw new AssertionError("typed store changed the wrong byte window"); + } + // Existing decoded views and aliases require the general write protocol. + Pair live = view(bytes); + live.first = 103; + Pointer.storeTypedStorage(storage, 8, 8, PAIR_CODEC, new Pair(113, 127)); + if (view(bytes).first != 103 || view(bytes.byte_offset(8)).first != 113 + || view(bytes.byte_offset(8)).second != 127) { + throw new AssertionError("typed store bypassed a pending memory view"); + } + // A scalar field write must preserve pending view writes and the adjacent field. + Pair fieldPending = view(bytes); + fieldPending.first = 79; + fieldPending.second = 81; + Pointer projectedSecond = bytes.retype(8, PAIR_CODEC) + .projectStructField(Pair.class.getName(), "second", 4, 4, null); + if (Pointer.loadLocationBits(projectedSecond, 0, 4) != 81) { + throw new AssertionError("scalar field read ignored a pending view"); + } + if (Pointer.loadScalarField(bytes, 0, Pair.class.getName(), "second", 4, 4) != 81) { + throw new AssertionError("scalar field fallback ignored a pending view"); + } + Pointer.storeScalarField(bytes, 0, Pair.class.getName(), "second", 4, 82, 4); + if (view(bytes).first != 79 || view(bytes).second != 82) { + throw new AssertionError("scalar field write lost its sibling or view coherence"); + } Pair pending = view(bytes); pending.second = 83; + if (Pointer.loadLocationBits(storage, 4, 4) != 83) { + throw new AssertionError("component address ignored a pending typed view"); + } + view(bytes).second = 87; + Pointer.storeLocationBits(storage, 0, 89, 4); + if (view(bytes).first != 89 || view(bytes).second != 87) { + throw new AssertionError("component address lost an adjacent pending field"); + } + view(bytes).second = 83; byte[] image = new byte[storage.length]; Pointer.encodeArrayMemory(storage, image, 0, 1, null); if (MemoryBytes.read(image, 4, 4) != 83) { @@ -138,8 +236,260 @@ public static void check() { fields.second = 23; Pointer root = Pointer.cell(fields, 8, PAIR_CODEC); Pointer first = root.projectStructField(Pair.class.getName(), "first", 0, 4, null); + if (Pointer.loadLocationBits(first, 4, 4) != 23) { + throw new AssertionError("component address escaped the wrong field owner"); + } if (first.add(0).offset_from(first) != 0 || first.add(1).getI32() != 23) { throw new AssertionError("pointer arithmetic lost the containing allocation"); } + Pointer next = Pointer.fromLocation(first, 4, 4); + if (next.getI32() != 23 || next.offset_from(root.retype(4)) != 1) { + throw new AssertionError("materialized component address lost its field origin"); + } + Pointer back = Pointer.fromLocation(next, -4, 4); + if (back.getI32() != 17 || back.offset_from(first) != 0) { + throw new AssertionError("materialized component address lost backward provenance"); + } + if (!Pointer.sameLocation(first, 4, next, 0) + || !Pointer.sameLocation(next, -4, first, 0) + || Pointer.sameLocation(first, 0, next, 0) + || !Pointer.sameLocation(null, 73, Pointer.fromAddress(73, 4), 0)) { + throw new AssertionError("component equality lost an address or containing-field origin"); + } + int[] cells = {4, 8}; + Object scalarRoot = Pointer.scalarSliceRoot(cells, 4); + if (scalarRoot != cells || Pointer.loadLocationBits(scalarRoot, 4, 4) != 8) { + throw new AssertionError("scalar slice lost direct array storage"); + } + Pointer.storeLocationBits(scalarRoot, 4, 12, 4); + if (cells[1] != 12) { + throw new AssertionError("scalar slice write detached its allocation"); + } + if (Pointer.locationSliceBacking(cells, 4, 4) != cells + || Pointer.locationSliceOffset(cells, 4, 4) != 1 + || Pointer.locationSliceBacking(Pointer.array(cells, 1, 4), -4, 4) != cells + || Pointer.locationSliceOffset(Pointer.array(cells, 1, 4), -4, 4) != 0 + || Pointer.fromLocation(cells, 4, 4).getI32() != 12) { + throw new AssertionError("address/view round trip lost direct array storage"); + } + try { + Pointer.locationSliceOffset(cells, 1, 4); + throw new AssertionError("unaligned slice conversion succeeded"); + } catch (IllegalStateException expected) { } + Object fieldView = Pointer.locationSliceBacking(first, 4, 4); + int fieldStart = Pointer.locationSliceOffset(first, 4, 4); + if (Pointer.sliceGetI32(fieldView, fieldStart) != 23) { + throw new AssertionError("address/view conversion lost a projected field origin"); + } + if (!Pointer.sameLocation(cells, 4, Pointer.array(cells, 0, 4), 4) + || Pointer.sameLocation(cells, 0, cells, 4)) { + throw new AssertionError("component equality lost an array address"); + } + } + + private static void checkTypedStorage() { + Pair initial = new Pair(); + initial.first = 17; + Object owner = Pointer.storageAligned(initial, 8, PAIR_CODEC, 8); + if (owner instanceof Pointer || Pointer.locationStride(owner) != 8 + || Pointer.loadStorageLocation(owner, 0, Pair.class.getName()) != initial) { + throw new AssertionError("typed storage eagerly materialized an address"); + } + Pair replacement = new Pair(); + replacement.first = 23; + replacement.second = 29; + Pointer.storeStorageLocation(owner, 0, replacement); + Pair copy = (Pair) Pointer.loadStorageCopy(owner, 0, Pair.class.getName()); + copy.first = 101; + if (replacement.first != 23) throw new AssertionError("owned storage copy aliased its source"); + Pair direct = (Pair) Pointer.directStorageAggregate(owner, 0, Pair.class); + if (direct != replacement) throw new AssertionError("typed root retained its old value"); + direct.second = 31; + Pointer.commitStorageLocation(owner, 0); + try { + java.lang.reflect.Field boundary = owner.getClass().getDeclaredField("boundary"); + boundary.setAccessible(true); + if (boundary.get(owner) != null) { + throw new AssertionError("typed loads and stores allocated a Pointer"); + } + } catch (ReflectiveOperationException error) { + throw new AssertionError(error); + } + Pointer escaped = Pointer.fromStorageLocation(owner, 0); + if (escaped != Pointer.fromStorageLocation(owner, 0) + || !Pointer.sameLocation(owner, 0, escaped, 0)) { + throw new AssertionError("materialization changed typed allocation identity"); + } + Pointer first = escaped.projectStructField(Pair.class.getName(), "first", 0, 4, null); + Pair next = new Pair(); + next.first = 37; + next.second = 41; + Pointer.storeStorageLocation(owner, 0, next); + if (first.getI32() != 37) throw new AssertionError("field alias lost owner replacement"); + first.set(43); + if (((Pair) Pointer.loadStorageLocation(owner, 0, Pair.class.getName())).first != 43) { + throw new AssertionError("typed load lost an escaped field write"); + } + Pointer.storeLocationBits(owner, 4, 47, 4); + if (((Pair) Pointer.loadStorageLocation(owner, 0, Pair.class.getName())).second != 47 + || Pointer.loadLocationBits(owner, 0, 4) != 43) { + throw new AssertionError("typed and byte storage disagree"); + } + } + + private static void checkSliceComponents() { + Pair[] values = {new Pair(3, 5), new Pair(7, 11), new Pair(13, 17)}; + Pointer array = Pointer.array(values, 1, 8, PAIR_CODEC); + Object root = Pointer.sliceAddressRoot(array, 8, PAIR_CODEC); + if (root != array || Pointer.directStorageAggregate(root, 8, Pair.class) != values[2]) { + throw new AssertionError("normalized aggregate slice lost its root or counted the start twice"); + } + Object subview = Pointer.typedLocationSliceBacking(root, 8, 8, PAIR_CODEC); + int start = Pointer.typedLocationSliceOffset(root, 8, 8, PAIR_CODEC); + if (subview != root || start != 1 || Pointer.sliceGetObject(subview, start) != values[2]) { + throw new AssertionError("aggregate subview reconstructed or displaced its owner"); + } + values[2] = new Pair(19, 23); + if (Pointer.directStorageAggregate(root, 8, Pair.class) != values[2]) { + throw new AssertionError("component slice retained a replaced array element"); + } + byte[] bytes = new byte[24]; + Pointer base = Pointer.array(bytes, 0, 1).retype(8, PAIR_CODEC); + root = Pointer.sliceAddressRoot(base, 8, PAIR_CODEC); + if (root != base) throw new AssertionError("normalized byte slice allocated a new root"); + if (Pointer.typedLocationSliceBacking(root, 8, 8, PAIR_CODEC) != root + || Pointer.typedLocationSliceOffset(root, 8, 8, PAIR_CODEC) != 1) { + throw new AssertionError("byte subview reconstructed its root"); + } + Pair live = (Pair) base.byte_offset(8).getObject(); + live.second = 29; + Object field = Pointer.storageFieldRoot(root, 8, Pair.class.getName(), "second", 4, 4, null); + long offset = Pointer.storageFieldOffset(field, root, 12); + if (field != root || offset != 12 || Pointer.loadLocationBits(field, offset, 4) != 29) { + throw new AssertionError("scalar field components lost a pending decoded view"); + } + Pointer.storeLocationBits(field, offset, 31, 4); + Pair copy = (Pair) Pointer.loadTypedStorageCopy(root, 8, 8, PAIR_CODEC, Pair.class.getName()); + if (copy.second != 31 || MemoryBytes.read(bytes, 4, 4) != 0) { + throw new AssertionError("component write changed a neighboring element"); + } + Object owner = Pointer.storageAligned(new Pair(37, 41), 8, PAIR_CODEC, 8); + Object projected = Pointer.storageFieldRoot(owner, 0, Pair.class.getName(), "second", 4, 4, null); + long displacement = Pointer.storageFieldOffset(projected, owner, 4); + Pointer.storeStorageLocation(owner, 0, new Pair(43, 47)); + if (Pointer.loadLocationBits(projected, displacement, 4) != 47) { + throw new AssertionError("component field detached from a replaceable owner"); + } + } + + private static void checkPointerTypedStores() { + byte[] bytes = new byte[24]; + Arrays.fill(bytes, (byte) 11); + Pointer base = Pointer.array(bytes, 4, 1); + int before = PairCodec.rangeEncodes; + Pointer.storeTypedStorage(base, 4, 8, PAIR_CODEC, new Pair(17, 19)); + if (PairCodec.rangeEncodes != before + 1 + || MemoryBytes.read(bytes, 8, 4) != 17 || MemoryBytes.read(bytes, 12, 4) != 19 + || bytes[7] != 11 || bytes[16] != 11) { + throw new AssertionError("typed pointer store missed its destination window"); + } + Pointer start = base.byte_offset(-4); + Pair left = view(start); + Pair right = view(start.byte_offset(8)); + left.first = 23; + left.second = 29; + right.first = 31; + right.second = 37; + Pointer.storeTypedStorage(base, 0, 8, PAIR_CODEC, new Pair(41, 43)); + if (view(start).first != 23 || view(start).second != 41 + || view(start.byte_offset(8)).first != 43 + || view(start.byte_offset(8)).second != 37) { + throw new AssertionError("typed pointer store lost adjacent pending view fields"); + } + Pair source = view(start); + source.first = 47; + source.second = 53; + Pointer.storeTypedStorage(base, -4, 8, PAIR_CODEC, source); + if (view(start).first != 47 || view(start).second != 53) { + throw new AssertionError("same-origin typed store discarded the source view"); + } + source = view(start); + source.second = 59; + Pointer.storeTypedStorage(base, 12, 8, PAIR_CODEC, source); + if (view(start.byte_offset(16)).first != 47 || view(start.byte_offset(16)).second != 59 + || view(start).second != 59) { + throw new AssertionError("typed pointer store copied a stale live source"); + } + byte[] snapshot = bytes.clone(); + try { + Pointer.storeTypedStorage(base, 16, 8, PAIR_CODEC, new Pair(61, 67)); + throw new AssertionError("out-of-bounds typed pointer store succeeded"); + } catch (IndexOutOfBoundsException expected) { } + if (!Arrays.equals(bytes, snapshot)) { + throw new AssertionError("invalid typed pointer store changed the allocation"); + } + } + + private static void checkTypedCopies() { + byte[] bytes = new byte[24]; + MemoryBytes.write(bytes, 8, 4, 31); + MemoryBytes.write(bytes, 12, 4, 37); + for (Object root : new Object[] {bytes, Pointer.array(bytes, 0, 1)}) { + Pair value = (Pair) Pointer.loadTypedStorageCopy(root, 8, 8, PAIR_CODEC, Pair.class.getName()); + if (value.first != 31 || value.second != 37 || PairCodec.decodedStorage != bytes) { + throw new AssertionError("typed copy ignored its exact view layout or copied bytes"); + } + value.first = 99; + if (MemoryBytes.read(bytes, 8, 4) != 31) throw new AssertionError("copy became a live view"); + } + Pointer pointer = Pointer.array(bytes, 0, 1); + Pair live = (Pair) pointer.byte_offset(8).retype(8, PAIR_CODEC).getObject(); + live.second = 41; + Pair value = (Pair) Pointer.loadTypedStorageCopy(bytes, 8, 8, PAIR_CODEC, Pair.class.getName()); + if (value.second != 41) throw new AssertionError("typed copy missed pending mutations"); + Object storage = Pointer.storageAligned(new Pair(43, 47), 8, PAIR_CODEC, 8); + Pair copy = (Pair) Pointer.loadTypedStorageCopy(storage, 0, 8, PAIR_CODEC, Pair.class.getName()); + Pointer.storeStorageLocation(storage, 0, new Pair(53, 59)); + if (copy.first != 43) throw new AssertionError("owner replacement modified an owned copy"); + try { + Pointer.loadTypedStorageCopy(bytes, 20, 8, PAIR_CODEC, Pair.class.getName()); + throw new AssertionError("typed copy did not check its byte window"); + } catch (IndexOutOfBoundsException expected) { } + } + + private static void checkOwnedRanges() { + byte[] bytes = new byte[24]; + Arrays.fill(bytes, (byte) 11); + Pointer base = Pointer.array(bytes, 0, 1).retype(8, PAIR_CODEC); + Pair value = new Pair(); + value.first = 53; + value.second = 59; + int before = PairCodec.rangeEncodes; + Pointer.storeStorageLocation(base, 8, value); + if (PairCodec.rangeEncodes != before + 1 || bytes[7] != 11 || bytes[16] != 11) { + throw new AssertionError("range encoder failed to write directly to its destination window"); + } + Pair copy = (Pair) Pointer.loadStorageCopy(base, 8, Pair.class.getName()); + if (copy.first != 53 || copy.second != 59 || PairCodec.decodedStorage != bytes) { + throw new AssertionError("owned range load created a byte snapshot or read the wrong window"); + } + copy.first = 61; + if (MemoryBytes.read(bytes, 8, 4) != 53) { + throw new AssertionError("owned range load established a live alias"); + } + Pair bound = (Pair) base.byte_offset(8).getObjectAs(Pair.class.getName()); + bound.second = 67; + Pair fresh = (Pair) Pointer.loadStorageCopy(base, 8, Pair.class.getName()); + if (fresh.second != 67) throw new AssertionError("owned range load ignored a live byte view"); + Pointer.storeStorageLocation(base, 8, copy); + Pair reread = (Pair) base.byte_offset(8).getObjectAs(Pair.class.getName()); + if (reread.first != 61 || reread.second != 59) { + throw new AssertionError("range store retained a stale decoded view"); + } + try { + Pointer.storeStorageLocation(base, 20, copy); + throw new AssertionError("out-of-bounds range store succeeded"); + } catch (IndexOutOfBoundsException expected) { } + if (bytes[20] != 11) throw new AssertionError("invalid range store wrote partial bytes"); } } diff --git a/tests/integration/pointer_provenance/SliceAliases.java b/tests/integration/pointer_provenance/SliceAliases.java new file mode 100644 index 00000000..20d68a94 --- /dev/null +++ b/tests/integration/pointer_provenance/SliceAliases.java @@ -0,0 +1,59 @@ +import org.rustlang.runtime.Pointer; + +/** Scalar slices keep observing aggregate views created after decomposition. */ +public final class SliceAliases { + public static void check() { + Object[] arrays = { new byte[8], new short[4], new char[4], new int[2], + new long[1], new float[2], new double[1] }; + int[] sizes = {1, 2, 2, 4, 8, 4, 8}; + for (int index = 0; index < arrays.length; index++) { + Object array = arrays[index]; + int size = sizes[index]; + Pointer data = Pointer.array(array, 0, size); + Object backing = Pointer.locationSliceBacking(data, 0, size); + int start = Pointer.locationSliceOffset(data, 0, size); + Pointer aggregate = data.retype(8, "MemoryViews$PairCodec#pair#LMemoryViews$Pair;"); + MemoryViews.Pair view = (MemoryViews.Pair) aggregate.getObject(); + view.first = 0x44332211; + view.second = 0; + if (array instanceof int[] + && pointer_provenance.pointer_provenance.first_word((int[]) array) != 0x44332211) { + throw new AssertionError("stale generated array read"); + } + long mask = size == 8 ? -1L : (1L << (8 * size)) - 1; + long value; + if (array instanceof byte[]) value = Pointer.sliceGetI8(backing, start) & 255L; + else if (array instanceof short[]) value = Pointer.sliceGetI16(backing, start) & 65535L; + else if (array instanceof char[]) value = Pointer.sliceGetU16(backing, start); + else if (array instanceof int[]) value = Pointer.sliceGetI32(backing, start) & 0xffffffffL; + else if (array instanceof long[]) value = Pointer.sliceGetI64(backing, start); + else if (array instanceof float[]) value = Float.floatToRawIntBits(Pointer.sliceGetF32(backing, start)) & 0xffffffffL; + else value = Double.doubleToRawLongBits(Pointer.sliceGetF64(backing, start)); + if (value != (0x44332211L & mask)) throw new AssertionError("stale slice read " + array.getClass()); + + if (array instanceof byte[]) Pointer.sliceSetI8(backing, start, (byte) 0x55); + else if (array instanceof short[]) Pointer.sliceSetI16(backing, start, (short) 0x55); + else if (array instanceof char[]) Pointer.sliceSetU16(backing, start, (char) 0x55); + else if (array instanceof int[]) Pointer.sliceSetI32(backing, start, 0x55); + else if (array instanceof long[]) Pointer.sliceSetI64(backing, start, 0x55); + else if (array instanceof float[]) Pointer.sliceSetF32(backing, start, Float.intBitsToFloat(0x55)); + else Pointer.sliceSetF64(backing, start, Double.longBitsToDouble(0x55)); + MemoryViews.Pair result = (MemoryViews.Pair) aggregate.getObject(); + long expected = (0x44332211L & ~mask) | 0x55; + if ((result.first & 0xffffffffL) != expected) throw new AssertionError("lost slice write " + array.getClass()); + + if (array instanceof byte[]) Pointer.fillArray((byte[]) array, (byte) 0x33); + else if (array instanceof short[]) Pointer.fillArray((short[]) array, (short) 0x33); + else if (array instanceof char[]) Pointer.fillArray((char[]) array, (char) 0x33); + else if (array instanceof int[]) Pointer.fillArray((int[]) array, 0x33); + else if (array instanceof long[]) Pointer.fillArray((long[]) array, 0x33L); + else if (array instanceof float[]) Pointer.fillArray((float[]) array, Float.intBitsToFloat(0x33)); + else Pointer.fillArray((double[]) array, Double.longBitsToDouble(0x33)); + result = (MemoryViews.Pair) aggregate.getObject(); + long word = size == 1 ? 0x33333333L : size == 2 ? 0x00330033L : 0x33L; + if ((result.first & 0xffffffffL) != word || (result.second & 0xffffffffL) != (size == 8 ? 0 : word)) { + throw new AssertionError("lost array fill " + array.getClass()); + } + } + } +} diff --git a/tests/integration/pointer_provenance/SliceLocations.java b/tests/integration/pointer_provenance/SliceLocations.java new file mode 100644 index 00000000..4e2e1a7e --- /dev/null +++ b/tests/integration/pointer_provenance/SliceLocations.java @@ -0,0 +1,98 @@ +import org.rustlang.runtime.Pointer; +import org.rustlang.runtime.MemoryBytes; + +/** Scalar and aggregate slice reads retain aliases without temporary addresses. */ +public final class SliceLocations { + public static void check() throws Exception { + trackedArrayCarriers(); + String[] values = {"zero", "one", "two"}; + Pointer objects = Pointer.array(values, 1, 8); + if (Pointer.sliceGetObject(objects, 1) != values[2]) throw new AssertionError("object slice displacement"); + int[] ints = {17, 23, 29}; + Pointer words = Pointer.array(ints, 1, 4); + Pointer.sliceSetI32(words, 1, 37); + if (Pointer.sliceGetI32(words, -1) != 17 || ints[2] != 37) throw new AssertionError("scalar slice displacement"); + Pointer local = Pointer.withMetadata(Pointer.cell(43, 4, null), 19); + if (Pointer.scalarSliceRoot(local, 4) != local || local.metadata() != 19) { + throw new AssertionError("normalizing scalar storage changed its metadata"); + } + Object root = Pointer.locationSliceBacking(local, 0, 4); + int start = Pointer.locationSliceOffset(local, 0, 4); + Pointer.sliceSetI32(root, start, 47); + if (local.getI32() != 47) throw new AssertionError("scalar slice lost write-through storage"); + String codec = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;"; + byte[] bytes = new byte[16]; + Pointer pairs = Pointer.array(bytes, 0, 1).retype(8, codec); + MemoryViews.Pair first = (MemoryViews.Pair) Pointer.sliceGetObject(pairs, 0); + MemoryViews.Pair second = (MemoryViews.Pair) Pointer.sliceGetObject(pairs, 1); + first.first = 53; + second.second = 59; + for (int i = 0; i < 10; i++) { + if (Pointer.sliceGetObject(pairs, 0) != first || Pointer.sliceGetObject(pairs, 1) != second) { + throw new AssertionError("decoded slice identity changed"); + } + } + if (Pointer.loadLocationBits(bytes, 0, 4) != 53 || Pointer.loadLocationBits(bytes, 12, 4) != 59) { + throw new AssertionError("decoded slice lost its memory origin"); + } + Pointer.storeLocationBits(bytes, 8, 61, 4); + MemoryViews.Pair replaced = (MemoryViews.Pair) Pointer.sliceGetObject(pairs, 1); + if (replaced == second || replaced.first != 61 || replaced.second != 59) { + throw new AssertionError("decoded slice reused an invalidated view"); + } + Pointer scalar = pairs.retype(4, null); + Pointer.sliceSetI32(scalar, 3, 67); + if (Pointer.sliceGetI32(scalar, 2) != 61 || MemoryBytes.read(bytes, 12, 4) != 67) { + throw new AssertionError("scalar slice bypassed aggregate invalidation"); + } + int[][] rows = {new int[] {71, 73}, new int[] {79, 83}}; + Pointer row = Pointer.array(rows, 0, 8, "@array:i32:2"); + // Nested array reads must still register the original allocation. + Object inner = Pointer.sliceGetObject(row, 1); + if (inner != rows[1] || Pointer.array(inner, 0, 4).addr() != row.addr() + 8) { + throw new AssertionError("nested slice array lost origin"); + } + } + + private static void trackedArrayCarriers() throws Exception { + for (String filterName : new String[] { + "MEMORY_VIEW_FILTER", "ENCODED_POINTER_FILTER", "ENCODED_REFERENCE_FILTER"}) { + java.lang.reflect.Field field = Pointer.class.getDeclaredField(filterName); + field.setAccessible(true); + Object filter = field.get(null); + java.lang.reflect.Method mark = Pointer.class.getDeclaredMethod( + "markIdentityFilter", filter.getClass(), Object.class); + mark.setAccessible(true); + for (Object sample : new Object[] {new byte[4], new short[4], new int[4], new long[128]}) { + int width = sample instanceof byte[] ? 1 : sample instanceof short[] ? 2 + : sample instanceof int[] ? 4 : 8; + int size = java.lang.reflect.Array.getLength(sample) * width; + String codec = "org/rustlang/runtime/ArrayMemoryCodec#array#" + + sample.getClass().getName() + "#" + size; + byte[] bytes = new byte[size]; + Pointer owner = Pointer.fromTypedStorageLocation(bytes, 0, size, codec); + Object view = owner.getObject(); + // False filter matches must not change which array receives the write. + mark.invoke(null, filter, view); + Pointer.storeLocationBits(view, width, 17, width); + if (((Number) java.lang.reflect.Array.get(view, 1)).longValue() != 17 + || Pointer.loadLocationBits(view, width, width) != 17 + || Pointer.loadLocationBits(bytes, width, width) != 17) { + throw new AssertionError("tracked scalar array detached from its carrier: " + filterName); + } + if (view instanceof long[]) { + view = owner.getObject(); + mark.invoke(null, filter, view); + Pointer.sliceSetI64(view, 11, 23); + if (((long[]) view)[11] != 23 || Pointer.sliceGetI64(view, 11) != 23 + || Pointer.loadLocationBits(bytes, 88, 8) != 23) + throw new AssertionError("tracked slice array detached from its carrier"); + } + Pointer.storeLocationBits(bytes, width, 29, width); + Object refreshed = owner.getObject(); + if (refreshed == view || Pointer.loadLocationBits(refreshed, width, width) != 29) + throw new AssertionError("raw byte alias did not invalidate the decoded array"); + } + } + } +} From a0c681966aa901f83355561136294e9719e19a42 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 13:18:45 +1000 Subject: [PATCH 10/61] copy and fill storage directly --- runtime/src/Heap.java | 40 +++ runtime/src/Pointer.java | 262 +++++++++++++++--- runtime/src/RuntimeSupport.java | 12 +- .../pointer_provenance/HeapStorage.java | 62 +++++ .../pointer_provenance/IoCopies.java | 65 +++++ .../integration/pointer_provenance/Main.java | 4 + .../pointer_provenance/MemoryCopies.java | 170 ++++++++++++ .../pointer_provenance/MemoryFills.java | 94 +++++++ 8 files changed, 660 insertions(+), 49 deletions(-) create mode 100644 runtime/src/Heap.java create mode 100644 tests/integration/pointer_provenance/HeapStorage.java create mode 100644 tests/integration/pointer_provenance/IoCopies.java create mode 100644 tests/integration/pointer_provenance/MemoryCopies.java create mode 100644 tests/integration/pointer_provenance/MemoryFills.java diff --git a/runtime/src/Heap.java b/runtime/src/Heap.java new file mode 100644 index 00000000..839ec296 --- /dev/null +++ b/runtime/src/Heap.java @@ -0,0 +1,40 @@ +package org.rustlang.runtime; + +/** Keeps allocator storage separate from address offsets. */ +public final class Heap { + private Heap() {} + + public static Object allocate(long byteCount, long alignment) { + try { + int size = Pointer.checkedArrayLength(byteCount); + int checkedAlignment = Pointer.checkedAlignment(alignment); + byte[] bytes = new byte[size]; + Pointer.registerHeapAllocation(bytes, checkedAlignment); + return bytes; + } catch (IllegalArgumentException | ArithmeticException | OutOfMemoryError failure) { + return null; + } + } + + public static Object reallocate(Object root, long offset, long oldByteCount, + long alignment, long newByteCount) { + if (root == null) throw new NullPointerException("Rust realloc requires a non-null pointer"); + byte[] destination = null; + try { + int oldSize = Pointer.checkedArrayLength(oldByteCount); + int newSize = Pointer.checkedArrayLength(newByteCount); + destination = (byte[]) allocate(newSize, alignment); + if (destination == null) return null; + Pointer.copyHeapAllocation(root, offset, destination, Math.min(oldSize, newSize)); + deallocate(root, offset); + return destination; + } catch (IllegalArgumentException | ArithmeticException | OutOfMemoryError failure) { + if (destination != null) Pointer.releaseHeapAllocation(destination); + return null; + } + } + + public static void deallocate(Object root, long offset) { + Pointer.releaseHeapAllocation(root); + } +} diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index 4c1882c1..89026f75 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -4966,20 +4966,25 @@ public static Pointer array(Object array, int elementOffset) { } public static Pointer allocateBytes(long byteCount, long alignment) { - try { - int size = checkedArrayLength(byteCount); - int checkedAlignment = checkedAlignment(alignment); - byte[] bytes = new byte[size]; - recordAlignment(bytes, checkedAlignment); - Pointer pointer = new Pointer(bytes, 1, 0, 1, null); - synchronized (ALLOCATIONS) { - ALLOCATOR_OWNED_ALLOCATIONS.put(bytes, Boolean.TRUE); - } - return pointer; - } catch (IllegalArgumentException | ArithmeticException | OutOfMemoryError failure) { - // GlobalAlloc reports allocation failure with a null pointer. In - // particular, Rust's usize range is much larger than a JVM array. - return null; + return fromLocation(Heap.allocate(byteCount, alignment), 0, 1); + } + + static void registerHeapAllocation(Object allocation, int alignment) { + recordAlignment(allocation, alignment); + synchronized (ALLOCATIONS) { + ALLOCATOR_OWNED_ALLOCATIONS.put(allocation, Boolean.TRUE); + } + } + + static void copyHeapAllocation(Object root, long offset, byte[] destination, int size) { + if (root instanceof byte[] + && !mayBeInIdentityFilter(MEMORY_VIEW_FILTER, root) + && !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, root) + && !mayBeInIdentityFilter(ENCODED_POINTER_FILTER, root) + && !mayBeInIdentityFilter(ENCODED_REFERENCE_FILTER, root)) { + System.arraycopy(root, Math.toIntExact(offset), destination, 0, size); + } else { + copy(fromLocation(root, offset, 1), array(destination, 0, 1), size); } } @@ -5210,34 +5215,18 @@ public static Pointer constantArray( } public static Pointer reallocateBytes( - Pointer source, - long oldByteCount, - long alignment, - long newByteCount) { - if (source == null) { - throw new NullPointerException("Rust realloc requires a non-null pointer"); - } - try { - int oldSize = checkedArrayLength(oldByteCount); - int newSize = checkedArrayLength(newByteCount); - Pointer destination = allocateBytes(newSize, alignment); - if (destination == null) { - return null; - } - copy(source, destination, Math.min(oldSize, newSize)); - deallocateBytes(source); - return destination; - } catch (IllegalArgumentException | ArithmeticException | OutOfMemoryError failure) { - return null; - } + Pointer source, long oldByteCount, long alignment, long newByteCount) { + return fromLocation(Heap.reallocate(source, 0, oldByteCount, alignment, newByteCount), 0, 1); } /** Releases a byte allocation after Rust's allocator has ended its lifetime. */ public static void deallocateBytes(Pointer pointer) { - if (pointer == null || pointer.allocation == null) { - return; - } - Object allocation = pointer.allocation; + Heap.deallocate(pointer, 0); + } + + static void releaseHeapAllocation(Object allocation) { + if (allocation instanceof Pointer) allocation = ((Pointer) allocation).allocation; + if (allocation == null) return; synchronized (ALLOCATIONS) { ALLOCATOR_OWNED_ALLOCATIONS.remove(allocation); AllocationInfo info = ALLOCATIONS.remove(allocation); @@ -5280,7 +5269,7 @@ public static void deallocateBytes(Pointer pointer) { discardEncodedReferences(allocation); } - private static int checkedAlignment(long alignment) { + static int checkedAlignment(long alignment) { if (alignment <= 0 || alignment > Integer.MAX_VALUE || (alignment & (alignment - 1L)) != 0) { @@ -10699,6 +10688,100 @@ public void set(Object value) { storeBytes(incomingBits(value, materializedSize), materializedSize); } + /** Copy exact locations without constructing their boundary wrappers. */ + public static void copyStorage(Object source, long sourceOffset, int sourceSize, String sourceCodec, + Object destination, long destinationOffset, int destinationSize, String destinationCodec, + long byteCount, boolean nonoverlapping) { + int count = checkedArrayLength(byteCount); + if (count == 0) return; + Object from = plainCopyArray(source); + Object to = plainCopyArray(destination); + if (from != null && to != null && from.getClass() == to.getClass()) { + long fromOffset = sourceOffset + (source instanceof Pointer ? ((Pointer) source).byteOffset : 0); + long toOffset = destinationOffset + (destination instanceof Pointer ? ((Pointer) destination).byteOffset : 0); + int width = inferredArrayElementSize(from); + if (fromOffset >= 0 && toOffset >= 0 && fromOffset % width == 0 + && toOffset % width == 0 && count % width == 0) { + if (nonoverlapping && from == to + && fromOffset < Math.addExact(toOffset, count) + && toOffset < Math.addExact(fromOffset, count)) { + throw new IllegalArgumentException("copy_nonoverlapping regions overlap"); + } + System.arraycopy(from, Math.toIntExact(fromOffset / width), + to, Math.toIntExact(toOffset / width), count / width); + return; + } + } + if (count <= 8 && tryCopyScalarLocations(source, sourceOffset, + destination, destinationOffset, count, nonoverlapping)) return; + Pointer fromPointer = fromTypedStorageLocation(source, sourceOffset, sourceSize, sourceCodec); + Pointer toPointer = fromTypedStorageLocation(destination, destinationOffset, destinationSize, destinationCodec); + if (nonoverlapping) copyNonOverlapping(fromPointer, toPointer, count); + else copy(fromPointer, toPointer, count); + } + + /** Read all source bytes before writing so overlapping scalar copies preserve their input. */ + private static boolean tryCopyScalarLocations(Object source, long sourceOffset, + Object destination, long destinationOffset, int count, boolean nonoverlapping) { + Object from = scalarCopyAllocation(source, sourceOffset, count); + Object to = scalarCopyAllocation(destination, destinationOffset, count); + if (from == null || to == null || hasEncodedCopyState(from) || hasEncodedCopyState(to)) return false; + long fromOffset = source instanceof Pointer + ? Math.addExact(((Pointer) source).byteOffset, sourceOffset) : sourceOffset; + long toOffset = destination instanceof Pointer + ? Math.addExact(((Pointer) destination).byteOffset, destinationOffset) : destinationOffset; + if (from == to) { + if (nonoverlapping && fromOffset < Math.addExact(toOffset, count) + && toOffset < Math.addExact(fromOffset, count)) { + throw new IllegalArgumentException("copy_nonoverlapping regions overlap"); + } + if (fromOffset == toOffset) return true; + } + long bits = loadLocationBits(source, sourceOffset, count); + // Flushing can publish encoded references. The general copy must transfer their provenance and GC roots. + if (hasEncodedCopyState(from) || hasEncodedCopyState(to)) return false; + storeLocationBits(destination, destinationOffset, bits, count); + return true; + } + + private static boolean hasEncodedCopyState(Object allocation) { + return mayBeInIdentityFilter(ENCODED_POINTER_FILTER, allocation) + || mayBeInIdentityFilter(ENCODED_REFERENCE_FILTER, allocation); + } + + private static Object scalarCopyAllocation(Object root, long offset, int count) { + if (root instanceof Storage) { + Storage storage = (Storage) root; + if (!directStorage(storage)) return null; + StorageLayout layout = storage.layout(((Cell) storage).value); + return layout != null && layout.at(offset, count) != null ? storage : null; + } + if (root instanceof Pointer) { + Pointer pointer = (Pointer) root; + // Projected fields can redirect arithmetic to a different owner. + if (pointer.addressState != null || pointer.allocationElementSize <= 0) return null; + root = pointer.allocation; + } else if (mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, root)) { + return null; + } + return root != null && root.getClass().isArray() + && root.getClass().getComponentType().isPrimitive() ? root : null; + } + + private static Object plainCopyArray(Object root) { + if (root instanceof Pointer) { + Pointer pointer = (Pointer) root; + if (pointer.allocation == null || pointer.rareState != null || pointer.addressState != null + || pointer.allocationCodecClassName != null + || pointer.allocationElementSize != inferredArrayElementSize(pointer.allocation)) return null; + root = pointer.allocation; + } + return root != null && root.getClass().isArray() + && root.getClass().getComponentType().isPrimitive() + && !hasScalarWriteTracking(root) + && !mayBeInIdentityFilter(MEMORY_VIEW_ORIGIN_FILTER, root) ? root : null; + } + public static void copy(Pointer source, Pointer destination, int byteCount) { if (tryCopyScalarRange( source, @@ -11006,15 +11089,102 @@ private static int rustIntegerCarrierValue(Object value) { } public static void writeBytes(Pointer destination, int value, int byteCount) { - byte[] bytes = new byte[byteCount]; - for (int index = 0; index < byteCount; index++) { - bytes[index] = (byte) value; - } - destination.storeRange(bytes); + writeBytes(destination, 0L, value, (long) byteCount); } public static void writeBytes(Pointer destination, int value, long byteCount) { - writeBytes(destination, value, checkedArrayLength(byteCount)); + writeBytes(destination, 0L, value, byteCount); + } + + /** Fill an exact byte range without allocating an address or a byte image. */ + public static void writeBytes(Object root, long offset, int value, long byteCount) { + int count = checkedArrayLength(byteCount); + if (count == 0) return; + long bits = (value & 255L) * 0x0101010101010101L; + if (root instanceof Storage && fillScalarStorage((Storage) root, offset, count, bits)) return; + root = normalizeLocationOrigin(root); + Pointer pointer = root instanceof Pointer ? (Pointer) root : null; + Object allocation = pointer == null ? root : pointer.allocation; + if (allocation != null && allocation.getClass().isArray() + && allocation.getClass().getComponentType().isPrimitive()) { + int width = inferredArrayElementSize(allocation); + if (pointer == null || pointer.allocationElementSize == width) { + long start = pointer == null ? offset : Math.addExact(pointer.byteOffset, offset); + long capacity = (long) Array.getLength(allocation) * width; + if (start < 0 || start > capacity - count) + throw new IndexOutOfBoundsException("byte fill exceeds primitive array storage"); + if (pointer == null && hasScalarWriteTracking(allocation)) + pointer = new Pointer(allocation, width, 0, width, null); + if (pointer != null) { + pointer.prepareMemoryWrite(start, count); + discardEncodedPointers(allocation, start, count); + } + fillPrimitiveArray(allocation, start, count, width, bits); + // Partial fills can leave other encoded references alive. + if (start == 0 && count == capacity) discardEncodedReferences(allocation); + return; + } + } + pointer = root instanceof Storage || root instanceof Pointer + ? fromStorageLocation(root, offset) : fromLocation(root, offset, 1); + byte[] bytes = new byte[count]; + Arrays.fill(bytes, (byte) value); + pointer.storeRange(bytes); + } + + private static boolean fillScalarStorage(Storage storage, long offset, int count, long bits) { + if (!directStorage(storage) || offset < 0 || offset > (long) storage.size - count) return false; + Object value = ((Cell) storage).value; + if (storage.size <= 8 && isPrimitiveScalarCarrier(value)) { + long mask = atomicMask(count) << (offset * 8); + long updated = (valueBits(value, storage.size) & ~mask) | ((bits << (offset * 8)) & mask); + ((Cell) storage).value = carrierFromBits(value, updated, storage.size); + return true; + } + StorageLayout layout = storage.layout(value); + if (layout == null) return false; + long end = offset + count; + // Validate all leaves first. Padding needs the general byte path without prior partial writes. + for (long cursor = offset; cursor < end; ) { + StorageLayout.Leaf leaf = layout.at(cursor, 1); + if (leaf == null) return false; + cursor += Math.min(end - cursor, leaf.size - (cursor - leaf.offset) % leaf.size); + } + for (long cursor = offset; cursor < end; ) { + StorageLayout.Leaf leaf = layout.at(cursor, 1); + int chunk = (int) Math.min(end - cursor, leaf.size - (cursor - leaf.offset) % leaf.size); + leaf.write(value, cursor, bits, chunk); + cursor += chunk; + } + return true; + } + + private static void fillPrimitiveArray(Object array, long offset, int count, int width, long bits) { + long end = offset + count; + int first = Math.toIntExact(offset / width), within = (int) (offset % width); + if (within != 0) { + int chunk = (int) Math.min(end - offset, width - within); + long mask = atomicMask(chunk) << (within * 8); + storePrimitiveArrayBits(array, first, + (primitiveArrayBits(array, first) & ~mask) | ((bits << (within * 8)) & mask)); + offset += chunk; + first++; + } + int last = Math.toIntExact(end / width); + if (first < last) { + if (array instanceof byte[]) Arrays.fill((byte[]) array, first, last, (byte) bits); + else if (array instanceof boolean[]) Arrays.fill((boolean[]) array, first, last, bits != 0); + else if (array instanceof short[]) Arrays.fill((short[]) array, first, last, (short) bits); + else if (array instanceof char[]) Arrays.fill((char[]) array, first, last, (char) bits); + else if (array instanceof int[]) Arrays.fill((int[]) array, first, last, (int) bits); + else if (array instanceof long[]) Arrays.fill((long[]) array, first, last, bits); + else if (array instanceof float[]) Arrays.fill((float[]) array, first, last, Float.intBitsToFloat((int) bits)); + else Arrays.fill((double[]) array, first, last, Double.longBitsToDouble(bits)); + } + if (offset < end && end % width != 0) { + long mask = atomicMask((int) (end % width)); + storePrimitiveArrayBits(array, last, (primitiveArrayBits(array, last) & ~mask) | (bits & mask)); + } } public static void writeElements(Pointer destination, int value, long elementCount) { @@ -11025,7 +11195,7 @@ private static int checkedElementByteCount(Pointer pointer, long elementCount) { return checkedArrayLength(Math.multiplyExact(elementCount, (long) pointer.viewSize)); } - private static int checkedArrayLength(long length) { + static int checkedArrayLength(long length) { if (length < 0 || length > Integer.MAX_VALUE) { throw new IllegalArgumentException("Rust memory operation exceeds JVM array limits"); } diff --git a/runtime/src/RuntimeSupport.java b/runtime/src/RuntimeSupport.java index e184014d..f5fbdb53 100644 --- a/runtime/src/RuntimeSupport.java +++ b/runtime/src/RuntimeSupport.java @@ -231,9 +231,7 @@ public static long readStdin(Pointer destination, long length) { if (read < 0) { return 0; } - for (int index = 0; index < read; index++) { - destination.add(index).set(Byte.valueOf(copy[index])); - } + copyBytes(copy, read, destination); return read; } @@ -256,6 +254,10 @@ private static String utf8(Pointer bytes, long length) { static byte[] copyFromPointer(Pointer source, long length) { int checkedLength = Math.toIntExact(length); byte[] copy = new byte[checkedLength]; + if (checkedLength != 0 && Pointer.locationStride(source) == 1) { + Pointer.copy(source, Pointer.array(copy, 0, 1), checkedLength); + return copy; + } for (int index = 0; index < checkedLength; index++) { copy[index] = source.add(index).getI8(); } @@ -270,6 +272,10 @@ static void copyBytes(byte[] source, int length, Pointer destination) { if (length < 0 || length > source.length) { throw new IndexOutOfBoundsException("invalid byte copy length " + length); } + if (length != 0 && Pointer.locationStride(destination) == 1) { + Pointer.copy(Pointer.array(source, 0, 1), destination, length); + return; + } for (int index = 0; index < length; index++) { destination.add(index).set(Byte.valueOf(source[index])); } diff --git a/tests/integration/pointer_provenance/HeapStorage.java b/tests/integration/pointer_provenance/HeapStorage.java new file mode 100644 index 00000000..35f62a43 --- /dev/null +++ b/tests/integration/pointer_provenance/HeapStorage.java @@ -0,0 +1,62 @@ +import java.lang.reflect.Field; +import java.util.Map; +import org.rustlang.runtime.Heap; +import org.rustlang.runtime.Pointer; + +public final class HeapStorage { + public static void check() throws Exception { + if (Heap.allocate(-1, 8) != null || Heap.allocate(16, 3) != null) { + throw new AssertionError("invalid allocation did not report failure"); + } + Object root = Heap.allocate(32, 512); + if (!(root instanceof byte[]) || Pointer.loadLocationBits(root, 0, 8) != 0) { + throw new AssertionError("allocator did not return zeroed storage directly"); + } + Pointer alias = Pointer.fromLocation(root, 8, 8); + if ((alias.addr() - 8) % 512 != 0 || !Pointer.sameLocation(root, 8, alias, 0)) { + throw new AssertionError("component allocation lost alignment or identity"); + } + Pointer.storeLocationBits(root, 8, 0x123456789abcdefL, 8); + if (alias.getI64() != 0x123456789abcdefL) throw new AssertionError("alias detached"); + if (Heap.reallocate(root, 0, 32, 512, -1) != null || alias.getI64() != 0x123456789abcdefL) { + throw new AssertionError("failed realloc changed the old allocation"); + } + Object grown = Heap.reallocate(root, 0, 32, 512, 64); + if (!(grown instanceof byte[]) || Pointer.loadLocationBits(grown, 8, 8) != 0x123456789abcdefL + || Pointer.loadLocationBits(grown, 56, 8) != 0) { + throw new AssertionError("growing allocator storage changed its contents"); + } + Object shrunk = Heap.reallocate(grown, 0, 64, 512, 16); + if (Pointer.loadLocationBits(shrunk, 8, 8) != 0x123456789abcdefL) { + throw new AssertionError("shrinking allocator storage changed its prefix"); + } + long exposed = Pointer.fromLocation(shrunk, 0, 1).expose_provenance(); + Field field = Pointer.class.getDeclaredField("EXPOSED_ADDRESSES"); + field.setAccessible(true); + Map addresses = (Map) field.get(null); + if (!addresses.containsKey(exposed)) throw new AssertionError("address was not exposed"); + Heap.deallocate(shrunk, 0); + if (addresses.containsKey(exposed)) throw new AssertionError("freed provenance was retained"); + + // Reallocation must flush decoded writes and retain encoded references. + String codec = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;"; + root = Heap.allocate(8, 8); + Pointer object = Pointer.fromLocation(root, 0, 1).retype(8, codec); + object.set(new MemoryViews.Pair(3, 5)); + MemoryViews.Pair view = (MemoryViews.Pair) object.getObject(); + view.second = 19; + grown = Heap.reallocate(root, 0, 8, 8, 16); + if (Pointer.loadLocationBits(grown, 4, 4) != 19) throw new AssertionError("realloc missed a live view"); + Heap.deallocate(grown, 0); + + root = Heap.allocate(8, 8); + Pointer target = Pointer.cell(73L, 8, null); + Pointer.fromLocation(root, 0, 8).retype(8, "@raw-pointer").set(target); + grown = Heap.reallocate(root, 0, 8, 8, 16); + Pointer stored = (Pointer) Pointer.fromLocation(grown, 0, 8).retype(8, "@raw-pointer").getObject(); + if (!stored.samePointer(target) || stored.getI64() != 73) { + throw new AssertionError("realloc lost an encoded reference"); + } + Heap.deallocate(grown, 0); + } +} diff --git a/tests/integration/pointer_provenance/IoCopies.java b/tests/integration/pointer_provenance/IoCopies.java new file mode 100644 index 00000000..2e5547f7 --- /dev/null +++ b/tests/integration/pointer_provenance/IoCopies.java @@ -0,0 +1,65 @@ +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.InputStream; +import java.io.PrintStream; +import java.util.Arrays; +import org.rustlang.runtime.Pointer; +import org.rustlang.runtime.RuntimeSupport; + +/** Host I/O shares the runtime's bulk-copy alias and byte-window semantics. */ +public final class IoCopies { + private static final String CODEC = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;"; + + public static void check() { + PrintStream previousOut = System.out; + InputStream previousIn = System.in; + ByteArrayOutputStream output = new ByteArrayOutputStream(); + System.setOut(new PrintStream(output)); + try { + RuntimeSupport.writeStdout(null, 0); + if (RuntimeSupport.readStdin(null, 0) != 0) throw new AssertionError("empty read"); + byte[] bytes = {11, 13, 17, 19, 23, 29}; + RuntimeSupport.writeStdout(Pointer.array(bytes, 1, 1), 4); + if (!Arrays.equals(output.toByteArray(), new byte[] {13, 17, 19, 23})) + throw new AssertionError("offset output"); + output.reset(); + try { + RuntimeSupport.writeStdout(Pointer.array(new int[] {31, 37, 255}, 0, 4), 3); + throw new AssertionError("wide output view accepted"); + } catch (IllegalStateException expected) { } + System.setIn(new ByteArrayInputStream(new byte[] {41, 43, 47})); + if (RuntimeSupport.readStdin(Pointer.array(bytes, 1, 1), 4) != 3 + || !Arrays.equals(bytes, new byte[] {11, 41, 43, 47, 23, 29})) + throw new AssertionError("short read changed neighboring bytes"); + int[] strided = {0, 0, 0}; + System.setIn(new ByteArrayInputStream(new byte[] {53, 59, -1})); + RuntimeSupport.readStdin(Pointer.array(strided, 0, 4), 3); + if (!Arrays.equals(strided, new int[] {53, 59, -1})) throw new AssertionError("strided input"); + + byte[] storage = new byte[16]; + Pointer owner = Pointer.array(storage, 4, 1).retype(8, CODEC); + MemoryViews.Pair view = (MemoryViews.Pair) owner.getObject(); + view.first = 61; view.second = 67; + output.reset(); + RuntimeSupport.writeStdout(owner.retype(1, null), 8); + if (!Arrays.equals(output.toByteArray(), new byte[] {61, 0, 0, 0, 67, 0, 0, 0})) + throw new AssertionError("output missed pending managed writes"); + view = (MemoryViews.Pair) owner.getObject(); + view.second = 71; + owner.commitMemoryView(); + System.setIn(new ByteArrayInputStream(new byte[] {73, 0, 0, 0})); + RuntimeSupport.readStdin(owner.retype(1, null), 4); + view = (MemoryViews.Pair) owner.getObject(); + if (view.first != 73 || view.second != 71 || storage[3] != 0 || storage[12] != 0) + throw new AssertionError("input lost alias coherence or neighboring bytes: " + + view.first + "," + view.second + " " + Arrays.toString(storage)); + try { + RuntimeSupport.writeStdout(Pointer.array(bytes, 4, 1), 4); + throw new AssertionError("invalid output range accepted"); + } catch (IndexOutOfBoundsException expected) { } + } finally { + System.setOut(previousOut); + System.setIn(previousIn); + } + } +} diff --git a/tests/integration/pointer_provenance/Main.java b/tests/integration/pointer_provenance/Main.java index 6101832a..eb2c41b1 100644 --- a/tests/integration/pointer_provenance/Main.java +++ b/tests/integration/pointer_provenance/Main.java @@ -19,8 +19,12 @@ public static void main(String[] args) throws Exception { LocationQueries.check(); SliceLocations.check(); OwnedFields.check(); + IoCopies.check(); + MemoryCopies.check(); + MemoryFills.check(); BorrowedFields.check(); BorrowedLocals.check(); + HeapStorage.check(); CodecAdapters.check(); MetadataFilters.check(); StructuralViews.check(); diff --git a/tests/integration/pointer_provenance/MemoryCopies.java b/tests/integration/pointer_provenance/MemoryCopies.java new file mode 100644 index 00000000..2b26d617 --- /dev/null +++ b/tests/integration/pointer_provenance/MemoryCopies.java @@ -0,0 +1,170 @@ +import java.util.Arrays; +import org.rustlang.runtime.Pointer; + +public final class MemoryCopies { + private static final String PAIR = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;"; + private static final String ADDRESS = "@raw-pointer\n4\n\n"; + + public static final class ReferenceBox { public Pointer value = Pointer.withoutProvenance(0L, 4); } + public static final class ReferenceCodec { + public static byte[] e$reference(ReferenceBox value) { + byte[] bytes = new byte[8]; + Pointer.fromTypedStorageLocation(bytes, 0, 8, ADDRESS).set(value.value); + return bytes; + } + public static ReferenceBox d$reference(byte[] bytes) { + ReferenceBox result = new ReferenceBox(); + result.value = (Pointer) Pointer.fromTypedStorageLocation(bytes, 0, 8, ADDRESS).getObject(); + return result; + } + } + + public static void check() throws Exception { + zeroLengthCopies(); + scalarLocations(); + encodedScalarLocations(); + for (boolean boxed : new boolean[] {false, true}) { + byte[] bytes = {2, 3, 5, 7, 11, 13, 17, 19}; + Object root = boxed ? Pointer.array(bytes, 1, 1) : bytes; + long base = boxed ? -1 : 0; + Pointer.copyStorage(root, base, 1, null, root, base + 2, 1, null, 6, false); + if (!Arrays.equals(bytes, new byte[] {2, 3, 2, 3, 5, 7, 11, 13})) + throw new AssertionError("forward overlap"); + Pointer.copyStorage(root, base + 2, 1, null, root, base, 1, null, 6, false); + if (!Arrays.equals(bytes, new byte[] {2, 3, 5, 7, 11, 13, 11, 13})) + throw new AssertionError("backward overlap"); + try { + Pointer.copyStorage(root, base, 1, null, root, base + 1, 1, null, 6, true); + throw new AssertionError("nonoverlap constraint lost"); + } catch (IllegalArgumentException expected) { } + Pointer.copyStorage(root, base + 8, 1, null, root, base + 8, 1, null, 0, true); + try { + Pointer.copyStorage(root, base, 1, null, root, base + 7, 1, null, 2, false); + throw new AssertionError("bounds not checked"); + } catch (IndexOutOfBoundsException expected) { } + } + int[] source = {0x01020304, 0x05060708, 0x11121314}; + int[] actual = {0, 0, 0}, expected = {0, 0, 0}; + Pointer.copyStorage(source, 4, 4, null, actual, 0, 4, null, 8, true); + if (!Arrays.equals(actual, new int[] {source[1], source[2], 0})) + throw new AssertionError("aligned primitive array copy"); + Pointer.copy(Pointer.array(source, 0, 4).byte_offset(1), + Pointer.array(expected, 0, 4).byte_offset(2), 7); + Arrays.fill(actual, 0); + Pointer.copyStorage(source, 1, 4, null, actual, 2, 4, null, 7, false); + if (!Arrays.equals(actual, expected)) throw new AssertionError("unaligned copy"); + + byte[] from = new byte[16], to = new byte[16]; + Pointer owner = Pointer.array(from, 4, 1).retype(8, PAIR); + MemoryViews.Pair pair = (MemoryViews.Pair) owner.getObject(); + pair.first = 29; pair.second = 31; + Pointer.copyStorage(owner, 0, 8, PAIR, to, 4, 8, PAIR, 8, false); + MemoryViews.Pair copied = (MemoryViews.Pair) Pointer.array(to, 4, 1).retype(8, PAIR).getObject(); + if (copied.first != 29 || copied.second != 31) throw new AssertionError("pending source writes"); + copied.first = 37; + Pointer.copyStorage(owner, 4, 4, null, to, 8, 4, null, 4, false); + copied = (MemoryViews.Pair) Pointer.array(to, 4, 1).retype(8, PAIR).getObject(); + if (copied.first != 37 || copied.second != 31) throw new AssertionError("destination aliases"); + + int[] pointee = {41, 43}; + Pointer.array(from, 0, 1).retype(8, ADDRESS).set(Pointer.array(pointee, 1, 4)); + Pointer.copyStorage(from, 0, 8, ADDRESS, to, 0, 8, ADDRESS, 8, true); + Pointer restored = (Pointer) Pointer.array(to, 0, 1).retype(8, ADDRESS).getObject(); + restored.set(47); + if (pointee[1] != 47) throw new AssertionError("encoded pointer provenance lost"); + } + + private static void scalarLocations() throws Exception { + for (int count = 1; count <= 8; count++) { + long[] source = {0x123456789abcdef0L, 0x1122334455667788L}; + byte[] bytes = new byte[16]; + Pointer.copyStorage(source, 1, 8, null, bytes, 3, 1, null, count, true); + if (Pointer.loadLocationBits(source, 1, count) != Pointer.loadLocationBits(bytes, 3, count)) + throw new AssertionError("cross-representation scalar copy: " + count); + } + MemoryViews.Pair value = new MemoryViews.Pair(17, 19); + Object storage = Pointer.storageAligned(value, 8, PAIR, 4, + "0,4,first\n4,4,second"); + byte[] bytes = new byte[8]; + Pointer.copyStorage(storage, 0, 4, null, bytes, 0, 1, null, 4, true); + Pointer.copyStorage(bytes, 0, 1, null, storage, 4, 4, null, 4, true); + java.lang.reflect.Field boundary = storage.getClass().getDeclaredField("boundary"); + boundary.setAccessible(true); + if (value.second != 17 || boundary.get(storage) != null) + throw new AssertionError("scalar copy materialized typed storage"); + + Pointer aggregate = Pointer.fromTypedStorageLocation(bytes, 0, 8, PAIR); + MemoryViews.Pair view = (MemoryViews.Pair) aggregate.getObject(); + view.first = 23; view.second = 29; + Pointer.copyStorage(aggregate, 0, 4, null, aggregate, 4, 4, null, 4, false); + view = (MemoryViews.Pair) aggregate.getObject(); + if (view.first != 23 || view.second != 23) + throw new AssertionError("decoded scalar copy lost pending writes or invalidation"); + // Distinct carriers into the same tracked allocation still overlap. + Pointer alias = Pointer.fromLocation(bytes, 1, 1); + try { + Pointer.copyStorage(aggregate, 0, 1, null, alias, 0, 1, null, 7, true); + throw new AssertionError("tracked byte aliases lost overlap validation"); + } catch (IllegalArgumentException expected) { } + byte[] expected = bytes.clone(); + System.arraycopy(expected, 0, expected, 1, 7); + Pointer.copyStorage(aggregate, 0, 1, null, alias, 0, 1, null, 7, false); + if (!Arrays.equals(bytes, expected)) throw new AssertionError("tracked overlapping scalar copy"); + } + + private static void encodedScalarLocations() throws Exception { + byte[] source = new byte[16], destination = new byte[8]; + int[] pointee = {41}; + String codec = "MemoryCopies$ReferenceCodec#reference#LMemoryCopies$ReferenceBox;#8"; + Pointer whole = Pointer.fromTypedStorageLocation(source, 0, 8, codec); + ReferenceBox decoded = (ReferenceBox) whole.getObject(); + decoded.value = Pointer.array(pointee, 0, 4); + // Flushing first publishes this provenance. A metadata check before the read cannot detect it. + Pointer.copyStorage(whole, 0, 1, null, destination, 0, 1, null, 8, true); + Pointer copied = (Pointer) Pointer.fromTypedStorageLocation(destination, 0, 8, ADDRESS).getObject(); + copied.set(43); + if (pointee[0] != 43) throw new AssertionError("newly encoded copy lost pointer provenance"); + Pointer.copyStorage(source, 0, 1, null, source, 4, 1, null, 8, false); + copied = (Pointer) Pointer.fromTypedStorageLocation(source, 4, 8, ADDRESS).getObject(); + copied.set(47); + if (pointee[0] != 47) throw new AssertionError("overlapping encoded copy lost provenance"); + + byte[] referenceSource = new byte[8], referenceTarget = new byte[8]; + Object dependency = new Object(); + java.lang.reflect.Method retain = Pointer.class.getDeclaredMethod("retainEncodedReference", Object.class, Object.class); + retain.setAccessible(true); + retain.invoke(null, referenceSource, dependency); + Pointer.copyStorage(referenceSource, 0, 1, null, referenceTarget, 0, 1, null, 8, true); + java.lang.reflect.Method index = Pointer.class.getDeclaredMethod("stateStripeIndex", Object.class); + index.setAccessible(true); + java.lang.reflect.Field references = Pointer.class.getDeclaredField("ENCODED_REFERENCES"); + references.setAccessible(true); + java.util.Map stripe = ((java.util.Map[]) references.get(null))[(int) index.invoke(null, referenceTarget)]; + if (stripe.get(referenceTarget) != dependency) throw new AssertionError("scalar copy lost GC dependency"); + } + + private static void zeroLengthCopies() throws Exception { + Object source = Pointer.storageAligned(new MemoryViews.Pair(5, 7), 8, PAIR, 4); + Object destination = Pointer.storageAligned(new MemoryViews.Pair(11, 13), 8, PAIR, 4); + Pointer.copyStorage(source, 0, 8, PAIR, destination, 0, 8, PAIR, 0, true); + java.lang.reflect.Field boundary = source.getClass().getDeclaredField("boundary"); + boundary.setAccessible(true); + if (boundary.get(source) != null || boundary.get(destination) != null) { + throw new AssertionError("zero-byte copy materialized storage boundaries"); + } + // Zero-sized copies need no offset arithmetic, even with dangling or maximal addresses. + Pointer.copyStorage(null, Long.MAX_VALUE, 0, "@zero-sized:unused", null, -1L, + 0, "@zero-sized:unused", 0, false); + Pointer dangling = Pointer.withoutProvenance(-1L, 0); + Pointer.copyElements(dangling, dangling, Long.MAX_VALUE); + for (long count : new long[] {-1L, (long) Integer.MAX_VALUE + 1, Long.MAX_VALUE}) { + try { + Pointer.copyStorage(source, 0, 8, PAIR, destination, 0, 8, PAIR, count, false); + throw new AssertionError("invalid byte count accepted"); + } catch (IllegalArgumentException expected) { } + } + if (boundary.get(source) != null || boundary.get(destination) != null) { + throw new AssertionError("invalid copy count materialized storage boundaries"); + } + } +} diff --git a/tests/integration/pointer_provenance/MemoryFills.java b/tests/integration/pointer_provenance/MemoryFills.java new file mode 100644 index 00000000..6d188443 --- /dev/null +++ b/tests/integration/pointer_provenance/MemoryFills.java @@ -0,0 +1,94 @@ +import java.lang.reflect.Field; +import java.lang.reflect.Method; +import java.util.Arrays; +import java.util.Map; +import org.rustlang.runtime.Pointer; + +public final class MemoryFills { + private static final String PAIR = "MemoryViews$PairCodec#pair#LMemoryViews$Pair;#8"; + + public static void check() throws Exception { + for (Object array : new Object[] {new byte[5], new boolean[5], new short[5], new char[5], + new int[5], new long[5], new float[5], new double[5]}) { + int width = array instanceof byte[] || array instanceof boolean[] ? 1 + : array instanceof short[] || array instanceof char[] ? 2 + : array instanceof int[] || array instanceof float[] ? 4 : 8; + int fill = array instanceof boolean[] ? 1 : 0xa5; + for (boolean carrier : new boolean[] {false, true}) { + for (int count : new int[] {1, width + 2, 3 * width}) { + byte[] initial = new byte[5 * width]; + Arrays.fill(initial, (byte) (array instanceof boolean[] ? 0 : 0x33)); + Pointer.decodeArrayMemory(initial, 0, array, width, null); + byte[] expected = initial.clone(); + Arrays.fill(expected, width - 1, width - 1 + count, (byte) fill); + if (carrier) Pointer.writeBytes(Pointer.array(array, 1, width).byte_offset(-1), fill, count); + else Pointer.writeBytes(array, width - 1, fill, count); + byte[] actual = new byte[initial.length]; + Pointer.encodeArrayMemory(array, actual, 0, width, null); + if (!Arrays.equals(actual, expected)) throw new AssertionError("unaligned fill: " + array.getClass()); + } + } + } + MemoryViews.Pair value = new MemoryViews.Pair(0x11223344, 0x55667788); + Object storage = Pointer.storageAligned(value, 8, PAIR, 4, "0,4,first\n4,4,second"); + Pointer.writeBytes(storage, 2, 0xa5, 4); + if (value.first != 0xa5a53344 || value.second != 0x5566a5a5) throw new AssertionError("scalar leaf fill"); + Field boundary = storage.getClass().getDeclaredField("boundary"); + boundary.setAccessible(true); + if (boundary.get(storage) != null) throw new AssertionError("scalar fill materialized a boundary"); + Object scalar = Pointer.storage(Integer.valueOf(0x11223344), 4, null); + Pointer.writeBytes(scalar, 1, 0xff, 2); + if (!Pointer.loadStorageLocation(scalar, 0, null).equals(0x11ffff44) || boundary.get(scalar) != null) + throw new AssertionError("scalar cell fill materialized or changed neighboring bytes"); + Pointer.writeBytes(storage, Long.MAX_VALUE, 1, 0); + Pointer.writeBytes(null, Long.MAX_VALUE, 1, 0); + Pointer.writeBytes((Pointer) null, 1, 0L); + Pointer.writeElements(Pointer.withoutProvenance(-1L, 0), 1, Long.MAX_VALUE); + for (long count : new long[] {-1, (long) Integer.MAX_VALUE + 1}) { + try { Pointer.writeBytes(storage, 0, 0, count); throw new AssertionError("invalid fill count"); } + catch (IllegalArgumentException expected) { } + } + if (boundary.get(storage) != null) throw new AssertionError("empty or invalid fill materialized storage"); + + Pointer aggregate = Pointer.cell(new MemoryViews.Pair(0x11223344, 0x55667788), 8, PAIR); + Pointer byteView = aggregate.retype(1, "org/rustlang/runtime/ArrayMemoryCodec#array#[B#1"); + Pointer.writeBytes(byteView, 1L, 0xa5, 4L); + MemoryViews.Pair filled = (MemoryViews.Pair) aggregate.getObject(); + if (filled.first != 0xa5a5a544 || filled.second != 0x556677a5) + throw new AssertionError("fill used the transient view codec instead of the allocation codec"); + + byte[] bytes = new byte[16]; + Pointer owner = Pointer.fromTypedStorageLocation(bytes, 4, 8, PAIR); + MemoryViews.Pair view = (MemoryViews.Pair) owner.getObject(); + view.first = 17; view.second = 19; + Pointer.writeBytes(bytes, 4, 0x11, 4); + view = (MemoryViews.Pair) owner.getObject(); + if (view.first != 0x11111111 || view.second != 19 || bytes[3] != 0 || bytes[12] != 0) + throw new AssertionError("fill lost decoded writes or neighboring bytes"); + byte[] before = bytes.clone(); + try { Pointer.writeBytes(bytes, 15, 0, 2); throw new AssertionError("invalid fill range"); } + catch (IndexOutOfBoundsException expected) { } + if (!Arrays.equals(bytes, before)) throw new AssertionError("invalid fill modified storage"); + encodedReferences(); + } + + private static void encodedReferences() throws Exception { + byte[] bytes = new byte[16]; + String codec = "@raw-pointer\n4\n\n"; + int[] values = {17, 19}; + Pointer.fromTypedStorageLocation(bytes, 0, 8, codec).set(Pointer.array(values, 0, 4)); + Pointer.fromTypedStorageLocation(bytes, 8, 8, codec).set(Pointer.array(values, 1, 4)); + Pointer.writeBytes(bytes, 0, 0, 8); + if (!Pointer.is_null((Pointer) Pointer.fromTypedStorageLocation(bytes, 0, 8, codec).getObject())) + throw new AssertionError("fill retained overwritten pointer metadata"); + ((Pointer) Pointer.fromTypedStorageLocation(bytes, 8, 8, codec).getObject()).set(23); + if (values[1] != 23) throw new AssertionError("partial fill erased neighboring provenance"); + Pointer.writeBytes(bytes, 0, 0, 16); + Field references = Pointer.class.getDeclaredField("ENCODED_REFERENCES"); + references.setAccessible(true); + Method index = Pointer.class.getDeclaredMethod("stateStripeIndex", Object.class); + index.setAccessible(true); + Map stripe = ((Map[]) references.get(null))[(int) index.invoke(null, (Object) bytes)]; + if (stripe.get(bytes) != null) throw new AssertionError("complete fill retained overwritten GC dependencies"); + } +} From de7972c9c978e987a0734335fe0fb5148642a9a4 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 15:27:47 +1000 Subject: [PATCH 11/61] share function call adapters --- runtime/src/FunctionCallSite.java | 63 +++++++++++++++++++++++ runtime/src/FunctionPointers.java | 85 +++++++++++++++++++++++++++++-- runtime/src/Pointer.java | 39 ++++++++++---- 3 files changed, 175 insertions(+), 12 deletions(-) create mode 100644 runtime/src/FunctionCallSite.java diff --git a/runtime/src/FunctionCallSite.java b/runtime/src/FunctionCallSite.java new file mode 100644 index 00000000..5eb946e2 --- /dev/null +++ b/runtime/src/FunctionCallSite.java @@ -0,0 +1,63 @@ +package org.rustlang.runtime; + +import java.lang.invoke.CallSite; +import java.lang.invoke.MethodHandle; +import java.lang.invoke.MethodHandles; +import java.lang.invoke.MethodType; +import java.lang.invoke.MutableCallSite; +import java.lang.invoke.WrongMethodTypeException; +import java.util.ArrayList; + +/** A few executed targets become direct calls without a class per cold target. */ +public final class FunctionCallSite extends MutableCallSite { + private static final int LIMIT = 4; + private static final MethodHandle SAME, LINK; + static { + try { + MethodHandles.Lookup lookup = MethodHandles.lookup(); + SAME = lookup.findStatic(FunctionCallSite.class, "same", + MethodType.methodType(boolean.class, MethodHandle.class, MethodHandle.class)); + LINK = lookup.findVirtual(FunctionCallSite.class, "link", + MethodType.methodType(Object.class, MethodHandle.class, Object[].class)); + } catch (ReflectiveOperationException error) { + throw new ExceptionInInitializerError(error); + } + } + private final ArrayList targets = new ArrayList<>(LIMIT); + private final MethodHandle linker, generic; + + private FunctionCallSite(MethodType type) { + super(type); + generic = MethodHandles.exactInvoker(type.dropParameterTypes(0, 1)); + linker = LINK.bindTo(this).asCollector(Object[].class, type.parameterCount() - 1).asType(type); + setTarget(linker); + } + + public static CallSite bootstrap(MethodHandles.Lookup lookup, String name, MethodType type) { + return new FunctionCallSite(type); + } + + private static boolean same(MethodHandle expected, MethodHandle actual) { + return expected == actual; + } + + private Object link(MethodHandle implementation, Object[] arguments) throws Throwable { + if (!implementation.type().equals(type().dropParameterTypes(0, 1))) { + throw new WrongMethodTypeException("Rust function ABI does not match its target"); + } + synchronized (this) { + if (targets.size() < LIMIT && !targets.contains(implementation)) { + targets.add(implementation); + MethodHandle chain = targets.size() == LIMIT ? generic : linker; + for (MethodHandle target : targets) { + chain = MethodHandles.guardWithTest(SAME.bindTo(target), + MethodHandles.dropArguments(target, 0, MethodHandle.class), chain); + } + setTarget(chain); + MutableCallSite.syncAll(new MutableCallSite[] {this}); + } + } + // Cached calls use the exact ABI and do not allocate an argument array. + return implementation.invokeWithArguments(arguments); + } +} diff --git a/runtime/src/FunctionPointers.java b/runtime/src/FunctionPointers.java index 66c7e1b2..89aa4291 100644 --- a/runtime/src/FunctionPointers.java +++ b/runtime/src/FunctionPointers.java @@ -23,6 +23,13 @@ protected Map computeValue(Class owner) { private FunctionPointers() {} + private static final ClassValue> CONSTANTS = + new ClassValue>() { + protected Map computeValue(Class signature) { + return new IdentityHashMap<>(); + } + }; + private static Object target(Class owner, String method) { Map methods = TARGETS.get(owner); synchronized (methods) { @@ -37,6 +44,10 @@ static Object identity(Object value) { return identity; } } + if (value instanceof MethodHandle) { + MethodHandleInfo info = MethodHandles.lookup().revealDirect((MethodHandle) value); + return target(info.getDeclaringClass(), info.getName() + ":" + info.getMethodType().toMethodDescriptorString()); + } if (value instanceof StaticFunctionPointer) { String name = ((StaticFunctionPointer) value).functionPointerIdentity(); int separator = name.indexOf("::"); @@ -51,18 +62,86 @@ static Object identity(Object value) { return value; } + /** Share the invocation class per exact ABI instead of per code target. */ + private static final ClassValue FACTORIES = new ClassValue() { + protected Factory computeValue(Class signature) { return new Factory(); } + }; + + private static final class Factory { + private MethodHandle factory; + private boolean initialized; + + synchronized MethodHandle get(MethodHandles.Lookup lookup, Class signature, + MethodType type) throws Throwable { + if (!initialized) { + if (signature.getName().startsWith("org.rustlang.runtime.FnPtr_")) { + try { + MethodHandle bridge = lookup.findStatic(signature, "$rust$invoke", + type.insertParameterTypes(0, MethodHandle.class)); + factory = LambdaMetafactory.metafactory(lookup, "call", + MethodType.methodType(signature, MethodHandle.class), type, + bridge, type).getTarget().asType( + MethodType.methodType(Object.class, MethodHandle.class)); + } catch (NoSuchMethodException absent) { + // Foreign/older interfaces retain their own Java SAM contract. + } + } + initialized = true; + } + return factory; + } + } + + public static Object bind(MethodHandles.Lookup lookup, Class signature, MethodHandle implementation) { + try { + return constant(lookup, signature, implementation); + } catch (RuntimeException | Error failure) { + throw failure; + } catch (Throwable failure) { + throw new IllegalStateException("could not link constant Rust function", failure); + } + } + + private static Object constant(MethodHandles.Lookup lookup, Class signature, + MethodHandle implementation) throws Throwable { + MethodType type = implementation.type(); + MethodHandleInfo info = lookup.revealDirect(implementation); + Object identity = target(info.getDeclaringClass(), + info.getName() + ":" + info.getMethodType().toMethodDescriptorString()); + Map constants = CONSTANTS.get(signature); + synchronized (constants) { + Object function = constants.get(identity); + if (function != null) return function; + MethodHandle factory = FACTORIES.get(signature).get(lookup, signature, type); + if (factory != null) { + function = (Object) factory.invokeExact(implementation); + } else { + CallSite site = LambdaMetafactory.metafactory(lookup, "call", + MethodType.methodType(signature), type, implementation, type); + function = site.getTarget().invoke(); + } + synchronized (IDENTITIES) { IDENTITIES.put(function, identity); } + constants.put(identity, function); + return function; + } + } + public static CallSite metafactory(MethodHandles.Lookup lookup, String name, MethodType factoryType, MethodType interfaceType, MethodHandle implementation, MethodType instantiatedType) throws Throwable { + if (!name.equals("call") || factoryType.parameterCount() != 0) { + return LambdaMetafactory.metafactory(lookup, name, factoryType, + interfaceType, implementation, instantiatedType); + } + // Executed reifications use direct JVM calls. Constant vtable entries use + // shared adapters in bind() to avoid hidden classes for unused targets. CallSite site = LambdaMetafactory.metafactory(lookup, name, factoryType, interfaceType, implementation, instantiatedType); Object function = site.getTarget().invoke(); MethodHandleInfo info = lookup.revealDirect(implementation); Object identity = target(info.getDeclaringClass(), info.getName() + ":" + info.getMethodType().toMethodDescriptorString()); - synchronized (IDENTITIES) { - IDENTITIES.put(function, identity); - } + synchronized (IDENTITIES) { IDENTITIES.put(function, identity); } return new ConstantCallSite(MethodHandles.constant(factoryType.returnType(), function)); } } diff --git a/runtime/src/Pointer.java b/runtime/src/Pointer.java index 89026f75..0e660822 100644 --- a/runtime/src/Pointer.java +++ b/runtime/src/Pointer.java @@ -148,11 +148,28 @@ static Object invokeRustFunction(Object function, Object... arguments) { } Method target = null; for (Method method : function.getClass().getMethods()) { - if (method.getName().equals("call") - && method.getParameterTypes().length == arguments.length) { - target = method; - break; - } + if (!method.getName().equals("call") || Modifier.isStatic(method.getModifiers())) continue; + Class[] params = method.getParameterTypes(); + if (params.length == arguments.length) { target = method; break; } + Object[] expanded = new Object[params.length]; + int next = 0; + for (Object argument : arguments) { + if (next >= params.length) { next = -1; break; } + if ((argument == null || argument instanceof Pointer) + && params[next] == Object.class && next + 1 < params.length + && params[next + 1] == long.class) { + expanded[next++] = argument; + expanded[next++] = 0L; + } else if ((argument == null || isSliceViewCarrierType(argument.getClass())) + && params[next] == Object.class && next + 2 < params.length + && params[next + 1] == int.class && params[next + 2] == long.class) { + SliceView view = (SliceView) argument; + expanded[next++] = view == null ? null : view.array; + expanded[next++] = view == null ? 0 : view.offset; + expanded[next++] = view == null ? 0L : view.rustLength; + } else { expanded[next++] = argument; } + } + if (next == params.length) { target = method; arguments = expanded; break; } } if (target == null) { throw new IllegalArgumentException( @@ -5667,10 +5684,14 @@ private static Pointer canonicalFunctionPointer(Object value) { Map cells = FUNCTION_POINTER_CELLS; Object identity = FunctionPointers.identity(value); if (value instanceof FunctionPointerAdapter) { - identity = canonicalFunctionPointer( - ((FunctionPointerAdapter) value).functionPointerTarget()); - cells = FUNCTION_POINTER_ADAPTER_CELLS.computeIfAbsent( - value.getClass(), key -> new IdentityHashMap<>()); + Object target = ((FunctionPointerAdapter) value).functionPointerTarget(); + if (target instanceof MethodHandle) { + identity = FunctionPointers.identity(target); + } else { + identity = canonicalFunctionPointer(target); + cells = FUNCTION_POINTER_ADAPTER_CELLS.computeIfAbsent( + value.getClass(), key -> new IdentityHashMap<>()); + } } Pointer existing = cells.get(identity); if (existing != null) { From 6766f723df8b125d1c46f64d88ac377a757b608f Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 17:59:08 +1000 Subject: [PATCH 12/61] load compact constant byte blocks --- runtime/src/ConstantData.java | 133 ++++++++++++++++++ runtime/src/MemoryBytes.java | 4 + .../pointer_provenance/ConstantBlocks.java | 116 +++++++++++++++ .../integration/pointer_provenance/Main.java | 1 + 4 files changed, 254 insertions(+) create mode 100644 runtime/src/ConstantData.java create mode 100644 tests/integration/pointer_provenance/ConstantBlocks.java diff --git a/runtime/src/ConstantData.java b/runtime/src/ConstantData.java new file mode 100644 index 00000000..1788f615 --- /dev/null +++ b/runtime/src/ConstantData.java @@ -0,0 +1,133 @@ +package org.rustlang.runtime; + +import java.io.EOFException; +import java.io.IOException; +import java.io.InputStream; +import java.nio.ByteBuffer; +import java.nio.ByteOrder; +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.Map; + +/** Lazy binary constants. Payload sharing never shares the destination allocation. */ +public final class ConstantData { + private ConstantData() {} + + private static final int MAX_BYTES = 4 * 1024 * 1024; + private static final int MAX_ENTRIES = 256; + // Entries retain only the per-owner map, never its Class or ClassLoader. + private static final ClassValue> OWNERS = new ClassValue>() { + protected Map computeValue(Class owner) { return new HashMap<>(); } + }; + private static final Map BLOCKS = new LinkedHashMap(16, 0.75f, true); + private static int retainedBytes; + + private static final class Block { + final Map owner; + final String name; + final ByteBuffer data; + Block(Map owner, String name, byte[] bytes) { + this.owner = owner; + this.name = name; + data = ByteBuffer.wrap(bytes).asReadOnlyBuffer(); + } + } + + private static ByteBuffer block(String name, Class owner, int size) { + if (size < 0 || size > 65536 || !name.startsWith("META-INF/rust-data/") + || name.indexOf("..") >= 0) throw new IllegalArgumentException("constant block"); + Map owned = OWNERS.get(owner); + synchronized (BLOCKS) { + Block cached = owned.get(name); + if (cached != null) { + if (cached.data.capacity() != size) throw new IllegalStateException("constant block size changed"); + BLOCKS.get(cached); + return cached.data; + } + } + // Resource I/O does not serialize unrelated constructor threads. + Block loaded = new Block(owned, name, readBlock(name, owner, size)); + synchronized (BLOCKS) { + Block existing = owned.get(name); + if (existing != null) return existing.data; + owned.put(name, loaded); + BLOCKS.put(loaded, Boolean.TRUE); + retainedBytes += size; + while (retainedBytes > MAX_BYTES || BLOCKS.size() > MAX_ENTRIES) { + java.util.Iterator oldest = BLOCKS.keySet().iterator(); + Block evicted = oldest.next(); + retainedBytes -= evicted.data.capacity(); + evicted.owner.remove(evicted.name); + oldest.remove(); + } + } + return loaded.data; + } + + /** Independent cursor over immutable, bounded, loader-local constant data. */ + public static ByteBuffer buffer(String name, Class owner, int size) { + return block(name, owner, size).duplicate().order(ByteOrder.LITTLE_ENDIAN); + } + + private static byte[] readBlock(String name, Class owner, int size) { + byte[] bytes = new byte[size]; + try (InputStream input = owner.getResourceAsStream("/" + name)) { + if (input == null) throw new IOException("missing Rust constant " + name); + for (int i = 0; i < size;) { + int read = input.read(bytes, i, size - i); + if (read < 0) throw new EOFException("truncated Rust constant"); + if (read == 0) { + int value = input.read(); + if (value < 0) throw new EOFException("truncated Rust constant"); + bytes[i++] = (byte) value; + } else i += read; + } + if (input.read() != -1) throw new IOException("oversized Rust constant"); + } catch (IOException error) { + throw new IllegalStateException("could not load Rust constant " + name, error); + } + return bytes; + } + + public static void readArray(Object target, ByteBuffer input) { + readArray(target, 0, java.lang.reflect.Array.getLength(target), input); + } + + public static void fill(Object target, int start, int count, String name, Class owner) { + int width; + if (target instanceof byte[] || target instanceof boolean[]) width = 1; + else if (target instanceof short[] || target instanceof char[]) width = 2; + else if (target instanceof int[] || target instanceof float[]) width = 4; + else if (target instanceof long[] || target instanceof double[]) width = 8; + else throw new IllegalArgumentException("not a primitive constant array"); + check(start, count, java.lang.reflect.Array.getLength(target)); + ByteBuffer input = buffer(name, owner, Math.multiplyExact(count, width)); + readArray(target, start, count, input); + } + + private static void readArray(Object target, int start, int count, ByteBuffer input) { + int end = start + count; + if (target instanceof byte[]) input.get((byte[]) target, start, count); + else if (target instanceof short[]) { + short[] a = (short[]) target; for (int i = start; i < end; i++) a[i] = input.getShort(); + } else if (target instanceof char[]) { + char[] a = (char[]) target; for (int i = start; i < end; i++) a[i] = input.getChar(); + } else if (target instanceof int[]) { + int[] a = (int[]) target; for (int i = start; i < end; i++) a[i] = input.getInt(); + } else if (target instanceof long[]) { + long[] a = (long[]) target; for (int i = start; i < end; i++) a[i] = input.getLong(); + } else if (target instanceof float[]) { + float[] a = (float[]) target; for (int i = start; i < end; i++) a[i] = input.getFloat(); + } else if (target instanceof double[]) { + double[] a = (double[]) target; for (int i = start; i < end; i++) a[i] = input.getDouble(); + } else if (target instanceof boolean[]) { + boolean[] a = (boolean[]) target; for (int i = start; i < end; i++) a[i] = input.get() != 0; + } else throw new IllegalArgumentException("not a primitive constant array"); + } + + private static void check(int start, int count, int length) { + if (start < 0 || count < 0 || start > length - count) { + throw new IndexOutOfBoundsException("constant array range"); + } + } +} diff --git a/runtime/src/MemoryBytes.java b/runtime/src/MemoryBytes.java index bbfd6a21..5024f8cf 100644 --- a/runtime/src/MemoryBytes.java +++ b/runtime/src/MemoryBytes.java @@ -11,6 +11,10 @@ public static void fillConstant(byte[] target, int offset, String chunk) { } } + public static void clear(byte[] bytes, int offset, int size) { + java.util.Arrays.fill(bytes, offset, offset + size, (byte) 0); + } + public static void write(byte[] bytes, int offset, int size, long value) { for (int index = 0; index < size; index++) { bytes[offset + index] = (byte) (value >>> (index * 8)); diff --git a/tests/integration/pointer_provenance/ConstantBlocks.java b/tests/integration/pointer_provenance/ConstantBlocks.java new file mode 100644 index 00000000..0d3cab00 --- /dev/null +++ b/tests/integration/pointer_provenance/ConstantBlocks.java @@ -0,0 +1,116 @@ +import java.io.ByteArrayInputStream; +import java.io.ByteArrayOutputStream; +import java.io.InputStream; +import java.nio.ByteBuffer; +import java.nio.ByteOrder; +import java.nio.ReadOnlyBufferException; +import java.util.Arrays; +import java.util.concurrent.atomic.AtomicInteger; +import java.util.concurrent.atomic.AtomicReference; +import org.rustlang.runtime.ConstantData; + +public final class ConstantBlocks { + public static final class Anchor { } + private static final String PREFIX = "META-INF/rust-data/"; + + private static final class Loader extends ClassLoader { + final int marker; + final AtomicInteger opens = new AtomicInteger(); + Loader(int marker) { super(ConstantBlocks.class.getClassLoader()); this.marker = marker; } + Class anchor(byte[] code) { return defineClass(null, code, 0, code.length); } + public InputStream getResourceAsStream(String name) { + if (!name.startsWith(PREFIX)) return super.getResourceAsStream(name); + opens.incrementAndGet(); + int size = Integer.parseInt(name.substring(name.lastIndexOf('-') + 1)); + if (name.contains("truncated")) size--; + if (name.contains("oversized")) size++; + byte[] bytes = new byte[size]; + for (int i = 0; i < size; i++) bytes[i] = (byte) (marker + i); + return new ByteArrayInputStream(bytes); + } + } + + public static void check() throws Exception { + byte[] code; + try (InputStream input = ConstantBlocks.class.getResourceAsStream("/ConstantBlocks$Anchor.class")) { + ByteArrayOutputStream output = new ByteArrayOutputStream(); + byte[] scratch = new byte[4096]; + for (int n; (n = input.read(scratch)) >= 0;) output.write(scratch, 0, n); + code = output.toByteArray(); + } + Loader first = new Loader(17), second = new Loader(43); + Class a = first.anchor(code), b = second.anchor(code); + String name = PREFIX + "cursor-8"; + ByteBuffer left = ConstantData.buffer(name, a, 8); + ByteBuffer right = ConstantData.buffer(name, a, 8); + if (left.getInt() != 0x14131211 || right.position() != 0 + || right.getInt() != 0x14131211 || first.opens.get() != 1) { + throw new AssertionError("constant blocks reread data or shared cursor state"); + } + try { + right.put(0, (byte) 0); + throw new AssertionError("shared constant data is mutable"); + } catch (ReadOnlyBufferException expected) { } + if (ConstantData.buffer(name, b, 8).getInt() != 0x2e2d2c2b) { + throw new AssertionError("constant data crossed loader boundaries"); + } + byte[] target = new byte[12]; + Arrays.fill(target, (byte) 9); + ConstantData.fill(target, 2, 8, name, a); + target[3] = 0; + ConstantData.fill(target, 2, 8, name, a); + if (target[1] != 9 || target[2] != 17 || target[3] != 18 || target[10] != 9) { + throw new AssertionError("constant destination range is shared or incorrect"); + } + long[] words = new long[3]; + ConstantData.fill(words, 1, 1, name, a); + if (words[0] != 0 || words[1] != 0x1817161514131211L || words[2] != 0) { + throw new AssertionError("constant byte order or scalar range"); + } + float[] floats = new float[2]; + ByteBuffer bits = ByteBuffer.allocate(8).order(ByteOrder.LITTLE_ENDIAN); + bits.putInt(0x7fc01234).putInt(0x80000000).flip(); + ConstantData.readArray(floats, bits); + if (Float.floatToRawIntBits(floats[0]) != 0x7fc01234 + || Float.floatToRawIntBits(floats[1]) != 0x80000000 || bits.position() != 8) { + throw new AssertionError("constant floating-point bits or cursor"); + } + for (String invalid : new String[] {"truncated-8", "oversized-8"}) { + try { + ConstantData.buffer(PREFIX + invalid, a, 8); + throw new AssertionError("malformed constant block accepted"); + } catch (IllegalStateException expected) { } + } + try { + ConstantData.buffer(name, a, 4); + throw new AssertionError("cached constant ignored its expected size"); + } catch (IllegalStateException expected) { } + int reads = first.opens.get(); + AtomicReference failure = new AtomicReference<>(); + Thread[] threads = new Thread[4]; + for (int i = 0; i < threads.length; i++) { + threads[i] = new Thread(() -> { + try { + for (int n = 0; n < 200; n++) { + if (ConstantData.buffer(name, a, 8).getLong() != 0x1817161514131211L) { + throw new AssertionError("concurrent constant cursor"); + } + } + } catch (Throwable error) { failure.set(error); } + }); + threads[i].start(); + } + for (Thread thread : threads) thread.join(); + if (failure.get() != null || first.opens.get() != reads) { + throw new AssertionError("constant cache concurrency", failure.get()); + } + for (int i = 0; i < 80; i++) ConstantData.buffer(PREFIX + "large" + i + "-65536", a, 65536); + reads = first.opens.get(); + ConstantData.buffer(name, a, 8); + if (first.opens.get() == reads) throw new AssertionError("constant byte cache is unbounded"); + for (int i = 0; i < 300; i++) ConstantData.buffer(PREFIX + "tiny" + i + "-1", a, 1); + reads = first.opens.get(); + ConstantData.buffer(name, a, 8); + if (first.opens.get() == reads) throw new AssertionError("constant entry cache is unbounded"); + } +} diff --git a/tests/integration/pointer_provenance/Main.java b/tests/integration/pointer_provenance/Main.java index eb2c41b1..dacd07a1 100644 --- a/tests/integration/pointer_provenance/Main.java +++ b/tests/integration/pointer_provenance/Main.java @@ -26,6 +26,7 @@ public static void main(String[] args) throws Exception { BorrowedLocals.check(); HeapStorage.check(); CodecAdapters.check(); + ConstantBlocks.check(); MetadataFilters.check(); StructuralViews.check(); Field field = Pointer.class.getDeclaredField("EXPOSED_ADDRESSES"); From 07fc163943af64162d2f52941db917f3a76bc90b Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 20:33:23 +1000 Subject: [PATCH 13/61] track exact classfile dependencies --- compiler-core/src/classfile/constant_pool.rs | 38 +- compiler-core/src/classfile/mod.rs | 1 + compiler-core/src/classfile/names.rs | 75 +++ compiler-core/src/classfile/resources.rs | 32 ++ compiler-core/src/classfile/summary.rs | 444 ++++++++++++++++-- compiler-core/src/classfile/summary/code.rs | 112 +++++ .../src/classfile/summary/forwarder.rs | 216 +++++++++ java-linker/src/tests.rs | 6 +- 8 files changed, 873 insertions(+), 51 deletions(-) create mode 100644 compiler-core/src/classfile/resources.rs create mode 100644 compiler-core/src/classfile/summary/code.rs create mode 100644 compiler-core/src/classfile/summary/forwarder.rs diff --git a/compiler-core/src/classfile/constant_pool.rs b/compiler-core/src/classfile/constant_pool.rs index b423cb7a..cd3c562e 100644 --- a/compiler-core/src/classfile/constant_pool.rs +++ b/compiler-core/src/classfile/constant_pool.rs @@ -9,6 +9,8 @@ pub struct InternedConstantPool { pool: ConstantPool<'static>, constants: HashMap, strings: HashMap, + resource_anchor: Option, + resources: Vec, } impl Default for InternedConstantPool { @@ -17,6 +19,8 @@ impl Default for InternedConstantPool { pool: ConstantPool::default(), constants: HashMap::default(), strings: HashMap::default(), + resource_anchor: None, + resources: Vec::new(), } } } @@ -30,7 +34,30 @@ impl Deref for InternedConstantPool { } impl InternedConstantPool { + pub fn set_resource_anchor(&mut self, owner: u16) { + self.resource_anchor = Some(owner); + } + + pub fn resource_anchor(&self) -> Option { + self.resource_anchor + } + + pub fn add_resource(&mut self, bytes: Vec) -> jvm::Result { + let resource = super::resources::Resource::new(bytes); + let name = self.add_name_string(&resource.name)?; + self.resources.push(resource); + Ok(name) + } + + pub fn take_resources(&mut self) -> Vec { + std::mem::take(&mut self.resources) + } + pub fn into_inner(self) -> ConstantPool<'static> { + assert!( + self.resources.is_empty(), + "binary constants were not emitted" + ); self.pool } @@ -84,20 +111,15 @@ impl InternedConstantPool { pub fn add_string>(&mut self, value: S) -> jvm::Result { let value = value.as_ref(); - let string_index = if value.starts_with(super::names::STRING_TAG) { - self.add_utf8(format!("{}{value}", super::names::LITERAL_STRING))? - } else { - self.add_utf8(value)? - }; + // Literal text must never become a reflection or class-liveness root. + // The final linker removes this tag without interpreting its payload. + let string_index = self.add_utf8(format!("{}{value}", super::names::LITERAL_STRING))?; self.add(Constant::String(string_index)) } /// Compiler-generated reflection names/descriptors need the same namespace /// relocation as class references. Rust string literals use `add_string`. pub fn add_name_string(&mut self, value: &str) -> jvm::Result { - if !value.contains(super::names::CRATE_MARKER) { - return self.add_string(value); - } let index = self.add_utf8(format!("{}{value}", super::names::NAME_STRING))?; self.add(Constant::String(index)) } diff --git a/compiler-core/src/classfile/mod.rs b/compiler-core/src/classfile/mod.rs index d3cf6cc2..8e32dc3e 100644 --- a/compiler-core/src/classfile/mod.rs +++ b/compiler-core/src/classfile/mod.rs @@ -5,6 +5,7 @@ pub mod encode; pub mod key; pub mod names; pub mod registry; +pub mod resources; pub mod summary; pub use ristretto_classfile::byte_reader::ByteReader; pub use ristretto_classfile::*; diff --git a/compiler-core/src/classfile/names.rs b/compiler-core/src/classfile/names.rs index 0a444fe8..17ae2bbe 100644 --- a/compiler-core/src/classfile/names.rs +++ b/compiler-core/src/classfile/names.rs @@ -5,6 +5,81 @@ pub const STRING_TAG: &str = "\u{1}rustc-jvm:"; pub const NAME_STRING: &str = "\u{1}rustc-jvm:name:"; pub const LITERAL_STRING: &str = "\u{1}rustc-jvm:literal:"; +/// Include the temporary linker tag in the JVM's modified UTF-8 limit. NUL +/// expands to two bytes and supplementary characters to a surrogate pair. +pub fn literal_fits(value: &str) -> bool { + let limit = u16::MAX as usize - LITERAL_STRING.len(); + value.len() <= limit / 2 + || value + .chars() + .map(|c| match c { + '\0' => 2, + c if c as u32 > 0xffff => 6, + c => c.len_utf8(), + }) + .sum::() + <= limit +} + +pub fn codec_owner(name: &str) -> bool { + name.rsplit('/') + .next() + .is_some_and(|name| name.starts_with("Codecs_")) +} + +pub fn codec_method_key(name: &str) -> Option<&str> { + let (operation, key) = name.split_once('$')?; + (matches!(operation, "e" | "w" | "d" | "a" | "b" | "s" | "c") + && key.len() == 16 + && key.bytes().all(|byte| byte.is_ascii_hexdigit())) + .then_some(key) +} + +/// Parse packed memory codecs in pointer and view recipes. +/// Declaration, emission and linking share this dependency grammar. +pub fn codec_recipes(value: &str) -> impl Iterator { + value.lines().filter_map(|line| { + let line = line.strip_prefix(NAME_STRING).unwrap_or(line); + let mut parts = line.split('#'); + let (Some(owner), Some(key), Some(descriptor)) = (parts.next(), parts.next(), parts.next()) + else { + return None; + }; + let valid_size = parts.next().is_none_or(|size| size.parse::().is_ok()); + (valid_size + && parts.next().is_none() + && codec_owner(owner) + && key.len() == 16 + && key.bytes().all(|byte| byte.is_ascii_hexdigit()) + && !descriptor.is_empty()) + .then_some((line, key)) + }) +} + +#[cfg(test)] +mod codec_tests { + #[test] + fn nested_recipes_retain_every_exact_identity() { + let first = "pkg/Codecs_12#1234567890abcdef#Lpkg/Value;"; + let second = "pkg/Codecs_ab#abcdef1234567890#[B#8"; + let nested = format!( + "{}@slice-pointer\nview\n4\n{first}\n@raw-pointer\n{second}", + super::NAME_STRING + ); + assert_eq!( + super::codec_recipes(&nested) + .map(|(text, _)| text) + .collect::>(), + [first, second] + ); + assert!( + super::codec_recipes("pkg/Value#1234567890abcdef#I\npkg/Codecs_12#oops#I") + .next() + .is_none() + ); + } +} + pub fn is_crate_marker(bytes: &[u8]) -> bool { bytes.len() == CRATE_MARKER_LEN && bytes.starts_with(CRATE_MARKER.as_bytes()) diff --git a/compiler-core/src/classfile/resources.rs b/compiler-core/src/classfile/resources.rs new file mode 100644 index 00000000..668265df --- /dev/null +++ b/compiler-core/src/classfile/resources.rs @@ -0,0 +1,32 @@ +//! Binary constant payloads, independent of Rust allocation identity. +pub const BUNDLE_PREFIX: &str = "@resource:"; +pub const DIRECTORY: &str = "META-INF/rust-data/"; +pub const BLOCK_BYTES: usize = 64 * 1024; + +#[derive(Clone, Debug)] +pub struct Resource { + pub name: String, + pub bytes: Vec, +} + +impl Resource { + pub fn new(bytes: Vec) -> Self { + // Two byte hashes make names deterministic across hosts. + // The linker compares complete payloads and rejects hash collisions. + let mut a = 0xcbf29ce484222325u64; + let mut b = 0x84222325cbf29ce4u64; + for &byte in &bytes { + a = (a ^ u64::from(byte)).wrapping_mul(0x100000001b3); + b = (b ^ u64::from(byte)).wrapping_mul(0x9e3779b185ebca87); + } + Self { + name: format!("{DIRECTORY}{a:016x}{b:016x}"), + bytes, + } + } +} + +pub fn valid_name(name: &str) -> bool { + name.strip_prefix(DIRECTORY) + .is_some_and(|hash| hash.len() == 32 && hash.bytes().all(|b| b.is_ascii_hexdigit())) +} diff --git a/compiler-core/src/classfile/summary.rs b/compiler-core/src/classfile/summary.rs index 37757bd1..8e127c0e 100644 --- a/compiler-core/src/classfile/summary.rs +++ b/compiler-core/src/classfile/summary.rs @@ -1,12 +1,71 @@ -//! Read class identity and entry-point metadata without decoding method bodies. -//! Constant strings borrow the input; only the returned class name is allocated. +//! Borrowed dependency summaries for private holders and proven carrier helpers. +//! Ordinary classes retain their complete constant-pool dependencies. use super::JavaStr; use std::io; +mod code; +mod forwarder; #[derive(Debug, PartialEq, Eq)] pub struct Summary { pub name: String, pub has_main: bool, + pub private: bool, + pub opaque_reflection: bool, + pub method_demands: bool, + pub carrier: Option>, +} + +pub const PRIVATE_ATTRIBUTE: &str = "RustJvmPrivate"; +pub const CARRIER_ATTRIBUTE: &str = "RustJvmCarrier"; + +/// These namespaces have no observable Java method ownership. +/// Eligible classes must also have the private marker and no instance state. +pub fn method_owner(name: &str) -> bool { + name.contains("/mono/Mono_") + || name + .rsplit('/') + .next() + .is_some_and(|n| n.starts_with("Codecs_")) +} + +/// Static enum helpers do not participate in virtual dispatch. +/// Exact calls and runtime reflection names determine which helpers are live. +pub fn enum_helper(name: &str) -> bool { + matches!( + name, + "eq" | "variantIndex" + | "is_some" + | "is_none" + | "_unionDiscriminant" + | "_fromUnionDiscriminant" + | "_writeUnionStorage" + | "_readUnionStorage" + ) || name.starts_with("_rust_drop_fields$") +} + +#[derive(Clone, Copy, PartialEq, Eq)] +pub struct MethodKey<'a> { + pub owner: &'a [u8], + pub name: &'a [u8], + pub descriptor: &'a [u8], +} + +#[derive(Clone, Copy)] +pub enum Dependency<'a> { + /// Definition metadata is separate from dependency edges. + Definition { + public: bool, + forward: Option>, + /// Count definition bytes before constant-pool sharing. + /// Count each repeated fragment once and exclude unreachable methods. + bytes: usize, + }, + Class(&'a [u8]), + Text(&'a [u8]), + String(&'a [u8]), + Method(MethodKey<'a>), + /// A field or non-private method name is part of a nominal Java surface. + FixedMemberName(&'a [u8]), } #[derive(Clone, Copy)] @@ -14,16 +73,26 @@ enum Constant<'a> { Other, Utf8(&'a [u8]), Class(u16), + Text(u16), + String(u16), + Member(u16, u16, bool), + NameAndType(u16, u16), + Handle(u16), + Dynamic(u16, u16), +} +struct Method<'a> { + flags: u16, + name: &'a [u8], + descriptor: &'a [u8], + code: Option<&'a [u8]>, + supported: bool, } - struct Reader<'a> { bytes: &'a [u8], } - fn invalid() -> io::Error { io::Error::new(io::ErrorKind::InvalidData, "invalid JVM class metadata") } - impl<'a> Reader<'a> { fn take(&mut self, count: usize) -> io::Result<&'a [u8]> { let (head, tail) = self.bytes.split_at_checked(count).ok_or_else(invalid)?; @@ -39,29 +108,118 @@ impl<'a> Reader<'a> { fn u32(&mut self) -> io::Result { Ok(u32::from_be_bytes(self.take(4)?.try_into().unwrap())) } - fn attributes(&mut self) -> io::Result<()> { - for _ in 0..self.u16()? { - self.u16()?; - let count = usize::try_from(self.u32()?).map_err(|_| invalid())?; - self.take(count)?; +} +struct Pool<'a> { + constants: Vec>, + bootstrap: Vec>, +} +impl<'a> Pool<'a> { + fn utf8(&self, index: u16) -> io::Result<&'a [u8]> { + match self.constants.get(index as usize) { + Some(Constant::Utf8(bytes)) => Ok(bytes), + _ => Err(invalid()), + } + } + fn class(&self, index: u16) -> io::Result<&'a [u8]> { + match self.constants.get(index as usize) { + Some(Constant::Class(index)) => self.utf8(*index), + _ => Err(invalid()), + } + } + fn member(&self, owner: u16, index: u16) -> io::Result> { + let Some(Constant::NameAndType(name, descriptor)) = self.constants.get(index as usize) + else { + return Err(invalid()); + }; + Ok(MethodKey { + owner: self.class(owner)?, + name: self.utf8(*name)?, + descriptor: self.utf8(*descriptor)?, + }) + } + fn visit( + &self, + index: u16, + stamp: u32, + seen: &mut [u32], + visit: &mut impl FnMut(Dependency<'a>), + ) -> io::Result<()> { + let entry = seen.get_mut(index as usize).ok_or_else(invalid)?; + if *entry == stamp { + return Ok(()); + } + *entry = stamp; + match self.constants.get(index as usize).ok_or_else(invalid)? { + Constant::Class(index) => visit(Dependency::Class(self.utf8(*index)?)), + Constant::Utf8(bytes) => visit(Dependency::Text(bytes)), + Constant::Text(index) => visit(Dependency::Text(self.utf8(*index)?)), + Constant::String(index) => visit(Dependency::String(self.utf8(*index)?)), + Constant::Member(owner, member, method) => { + let key = self.member(*owner, *member)?; + visit(Dependency::Class(key.owner)); + visit(Dependency::Text(key.descriptor)); + if *method { + visit(Dependency::Method(key)); + } else { + visit(Dependency::FixedMemberName(key.name)); + } + } + Constant::NameAndType(_, descriptor) => { + visit(Dependency::Text(self.utf8(*descriptor)?)) + } + Constant::Handle(index) => self.visit(*index, stamp, seen, visit)?, + Constant::Dynamic(bootstrap, signature) => { + if let Some(Constant::NameAndType(name, _)) = + self.constants.get(*signature as usize) + { + visit(Dependency::FixedMemberName(self.utf8(*name)?)); + } + self.visit(*signature, stamp, seen, visit)?; + for &index in self + .bootstrap + .get(*bootstrap as usize) + .ok_or_else(invalid)? + { + self.visit(index, stamp, seen, visit)?; + } + } + Constant::Other => {} } Ok(()) } } pub fn read(bytes: &[u8]) -> io::Result { + read_dependencies(bytes, |_| {}) +} +pub fn read_dependencies( + bytes: &[u8], + mut visit: impl FnMut(Dependency<'_>), +) -> io::Result { + read_demands(bytes, |_, dependency| visit(dependency)) +} + +/// Limit method scopes to private, stateless compiler namespaces. +/// Keep other fragments intact. Borrowed data remains local to this call. +pub fn read_demands<'a>( + bytes: &'a [u8], + mut visit: impl FnMut(Option>, Dependency<'a>), +) -> io::Result { let mut r = Reader { bytes }; if r.take(4)? != b"\xca\xfe\xba\xbe" { return Err(invalid()); } - r.take(4)?; // minor/major version - let count = usize::from(r.u16()?); - let mut constants = vec![Constant::Other; count]; + r.take(4)?; + let count = r.u16()? as usize; + let mut pool = Pool { + constants: vec![Constant::Other; count], + bootstrap: Vec::new(), + }; let mut index = 1; while index < count { - constants[index] = match r.u8()? { + pool.constants[index] = match r.u8()? { 1 => { - let len = usize::from(r.u16()?); + let len = r.u16()? as usize; Constant::Utf8(r.take(len)?) } 7 => Constant::Class(r.u16()?), @@ -77,52 +235,254 @@ pub fn read(bytes: &[u8]) -> io::Result { } Constant::Other } - 8 | 16 | 19 | 20 => { + 8 => Constant::String(r.u16()?), + 16 => Constant::Text(r.u16()?), + 19 | 20 => { r.take(2)?; Constant::Other } - 9 | 10 | 11 | 12 | 17 | 18 => { - r.take(4)?; - Constant::Other - } + tag @ 9..=11 => Constant::Member(r.u16()?, r.u16()?, tag != 9), + 12 => Constant::NameAndType(r.u16()?, r.u16()?), + 17 | 18 => Constant::Dynamic(r.u16()?, r.u16()?), 15 => { - r.take(3)?; - Constant::Other + r.u8()?; + Constant::Handle(r.u16()?) } _ => return Err(invalid()), }; index += 1; } - let utf8 = |index: u16| match constants.get(usize::from(index)) { - Some(Constant::Utf8(bytes)) => Ok(*bytes), - _ => Err(invalid()), - }; - r.u16()?; // access flags - let Some(Constant::Class(name_index)) = constants.get(usize::from(r.u16()?)) else { - return Err(invalid()); - }; - let name = JavaStr::from_mutf8(utf8(*name_index)?) + let flags = r.u16()?; + let owner = pool.class(r.u16()?)?; + let name = JavaStr::from_mutf8(owner) .map_err(|_| invalid())? .to_rust_string(); - r.u16()?; // superclass - let interface_count = usize::from(r.u16()?); - r.take(interface_count * 2)?; - for _ in 0..r.u16()? { - r.take(6)?; // flags, name, descriptor - r.attributes()?; + let superclass = r.u16()?; + let interfaces = (0..r.u16()?) + .map(|_| r.u16()) + .collect::>>()?; + let fields = r.u16()?; + let mut field_descriptors = Vec::with_capacity(fields as usize); + let mut plain_fields = true; + for _ in 0..fields { + plain_fields &= r.u16()? & 0x0008 == 0; + visit(None, Dependency::FixedMemberName(pool.utf8(r.u16()?)?)); + field_descriptors.push(pool.utf8(r.u16()?)?); + let attributes = r.u16()?; + plain_fields &= attributes == 0; + for _ in 0..attributes { + r.u16()?; + let length = r.u32()? as usize; + r.take(length)?; + } } + let mut methods = Vec::new(); let mut has_main = false; + let mut opaque_reflection = false; for _ in 0..r.u16()? { let flags = r.u16()?; - let name = utf8(r.u16()?)?; - let descriptor = utf8(r.u16()?)?; + opaque_reflection |= flags & 0x0100 != 0; + let name = pool.utf8(r.u16()?)?; + let descriptor = pool.utf8(r.u16()?)?; has_main |= flags & 0x0009 == 0x0009 && name == b"main" && descriptor == b"([Ljava/lang/String;)V"; - r.attributes()?; + let mut method = Method { + flags, + name, + descriptor, + code: None, + supported: true, + }; + for _ in 0..r.u16()? { + let attribute = pool.utf8(r.u16()?)?; + let len = r.u32()? as usize; + let data = r.take(len)?; + match attribute { + b"Code" => method.code = Some(data), + b"MethodParameters" => {} + _ => method.supported = false, + } + } + methods.push(method); + } + let mut private = false; + let mut carrier = None; + let mut supported = true; + let mut metadata_classes = Vec::new(); + for _ in 0..r.u16()? { + let attribute = pool.utf8(r.u16()?)?; + let len = r.u32()? as usize; + let payload = r.take(len)?; + match attribute { + b"RustJvmPrivate" if payload.is_empty() => private = true, + b"RustJvmCarrier" => carrier = Some(payload.to_vec()), + b"SourceFile" => {} + b"InnerClasses" => { + let mut nested = Reader { bytes: payload }; + for _ in 0..nested.u16()? { + for _ in 0..2 { + let index = nested.u16()?; + if index != 0 { + metadata_classes.push(pool.class(index)?); + } + } + nested.take(4)?; // simple name and access flags + } + if !nested.bytes.is_empty() { + return Err(invalid()); + } + } + b"BootstrapMethods" => { + let mut b = Reader { bytes: payload }; + for _ in 0..b.u16()? { + let handle = b.u16()?; + let mut arguments = vec![handle]; + for _ in 0..b.u16()? { + arguments.push(b.u16()?); + } + pool.bootstrap.push(arguments); + } + if !b.bytes.is_empty() { + return Err(invalid()); + } + } + _ => supported = false, + } } - r.attributes()?; if !r.bytes.is_empty() { return Err(invalid()); } - Ok(Summary { name, has_main }) + for constant in &pool.constants { + let Constant::Member(owner, member, true) = *constant else { + continue; + }; + let key = pool.member(owner, member)?; + // A symbolic owner may be a ClassLoader subclass, so do not require + // the exact platform owner for class-loading operations. + opaque_reflection |= matches!( + key.name, + b"forName" + | b"loadClass" + | b"findClass" + | b"defineClass" + | b"getResource" + | b"getResources" + | b"getResourceAsStream" + ) || (matches!( + key.owner, + b"java/lang/Class" | b"java/lang/invoke/MethodHandles$Lookup" + ) && matches!( + key.name, + b"findStatic" + | b"findVirtual" + | b"findConstructor" + | b"findGetter" + | b"findSetter" + | b"findStaticGetter" + | b"findStaticSetter" + | b"getMethod" + | b"getDeclaredMethod" + | b"getField" + | b"getDeclaredField" + | b"getConstructor" + | b"getDeclaredConstructor" + )); + } + let static_owner = method_owner(&name) + && flags & 0x0200 == 0 + && interfaces.is_empty() + && methods + .iter() + .all(|m| m.flags & 0x0008 != 0 && m.flags & 0x0520 == 0 && m.name != b""); + let private_interface = flags & 0x0200 != 0 && !method_owner(&name); + let private_carrier = carrier.is_some() && plain_fields; + let method_demands = private + && supported + && (((static_owner || private_interface) && fields == 0) || private_carrier) + && superclass != 0 + && pool.class(superclass)? == b"java/lang/Object" + && methods.iter().all(|m| m.supported); + let mut seen = vec![0; count]; + if !method_demands || !static_owner { + for method in &methods { + visit(None, Dependency::FixedMemberName(method.name)); + } + } + if method_demands { + visit(None, Dependency::Class(b"java/lang/Object")); + for name in metadata_classes { + visit(None, Dependency::Class(name)); + } + for &interface in &interfaces { + visit(None, Dependency::Class(pool.class(interface)?)); + } + for descriptor in field_descriptors { + visit(None, Dependency::Text(descriptor)); + } + for (index, method) in methods.iter().enumerate() { + // Pure carrier recipes prove this final equality helper. + // Its virtual calls name the exact owner. Copy and interface dispatch + // still require the containing class. + let carrier_eq = private_carrier + && method.flags & 0x0010 != 0 + && method.name == b"eq" + && method + .descriptor + .strip_prefix(b"(L") + .and_then(|s| s.strip_suffix(b";)Z")) + == Some(owner); + let scope = ((method.flags & 0x0008 != 0 + && method.flags & 0x0520 == 0 + && method.name != b"") + || carrier_eq) + .then_some(MethodKey { + owner, + name: method.name, + descriptor: method.descriptor, + }); + if scope.is_some() { + let forward = (static_owner && method.flags & 0x0001 != 0) + .then(|| { + method + .code + .and_then(|bytes| forwarder::target(bytes, method.descriptor, &pool)) + }) + .flatten(); + visit( + scope, + Dependency::Definition { + public: method.flags & 0x0001 != 0, + forward, + bytes: method.code.map_or(0, <[u8]>::len) + + method.name.len() + + method.descriptor.len() + + 32, + }, + ); + } + visit(scope, Dependency::Text(method.descriptor)); // defines even an empty body + let mut dependency = |d| visit(scope, d); + if let Some(bytes) = method.code { + let stamp = index as u32 + 1; + code::constants(bytes, |i| pool.visit(i, stamp, &mut seen, &mut dependency))?; + code::stack_maps( + bytes, + |i| pool.utf8(i), + |i| pool.visit(i, stamp, &mut seen, &mut dependency), + )?; + } + } + } else { + for index in 1..count { + pool.visit(index as u16, 1, &mut seen, &mut |d| visit(None, d))?; + } + } + Ok(Summary { + name, + has_main, + private, + opaque_reflection, + method_demands, + carrier: private.then_some(carrier).flatten(), + }) } diff --git a/compiler-core/src/classfile/summary/code.rs b/compiler-core/src/classfile/summary/code.rs new file mode 100644 index 00000000..4c038118 --- /dev/null +++ b/compiler-core/src/classfile/summary/code.rs @@ -0,0 +1,112 @@ +//! Visit constant indexes without retaining an instruction or stack-map graph. +use super::*; +use crate::classfile::{ByteReader, attributes::Instruction}; + +pub(super) fn constants( + bytes: &[u8], + mut visit: impl FnMut(u16) -> io::Result<()>, +) -> io::Result<()> { + let mut r = Reader { bytes }; + r.take(4)?; // stack/local bounds + let length = r.u32()? as usize; + let mut code = ByteReader::new(r.take(length)?); + while code.remaining() != 0 { + use Instruction::*; + let instruction = Instruction::from_bytes(&mut code).map_err(|_| invalid())?; + let index = match instruction { + Ldc(i) => u16::from(i), + Ldc_w(i) + | Ldc2_w(i) + | Getstatic(i) + | Putstatic(i) + | Getfield(i) + | Putfield(i) + | Invokevirtual(i) + | Invokespecial(i) + | Invokestatic(i) + | Invokeinterface(i, _) + | Invokedynamic(i) + | New(i) + | Anewarray(i) + | Checkcast(i) + | Instanceof(i) + | Multianewarray(i, _) => i, + _ => continue, + }; + visit(index)?; + } + for _ in 0..r.u16()? { + r.take(6)?; + let catch = r.u16()?; + if catch != 0 { + visit(catch)?; + } + } + Ok(()) +} + +pub(super) fn stack_maps<'a>( + bytes: &[u8], + utf8: impl Fn(u16) -> io::Result<&'a [u8]>, + mut visit: impl FnMut(u16) -> io::Result<()>, +) -> io::Result<()> { + let mut r = Reader { bytes }; + r.take(4)?; + let length = r.u32()? as usize; + r.take(length)?; + let catches = r.u16()? as usize; + r.take(catches * 8)?; + for _ in 0..r.u16()? { + let name = utf8(r.u16()?)?; + let length = r.u32()? as usize; + let data = r.take(length)?; + if name != b"StackMapTable" { + continue; + } + let mut frames = Reader { bytes: data }; + fn ty(r: &mut Reader<'_>, visit: &mut impl FnMut(u16) -> io::Result<()>) -> io::Result<()> { + match r.u8()? { + 0..=6 => {} + 7 => visit(r.u16()?)?, + 8 => { + r.u16()?; + } + _ => return Err(invalid()), + } + Ok(()) + } + for _ in 0..frames.u16()? { + match frames.u8()? { + 0..=63 => {} + 64..=127 => ty(&mut frames, &mut visit)?, + 247 => { + frames.u16()?; + ty(&mut frames, &mut visit)?; + } + 248..=251 => { + frames.u16()?; + } + frame @ 252..=254 => { + frames.u16()?; + for _ in 0..frame - 251 { + ty(&mut frames, &mut visit)?; + } + } + 255 => { + frames.u16()?; + for _ in 0..frames.u16()? { + ty(&mut frames, &mut visit)?; + } + for _ in 0..frames.u16()? { + ty(&mut frames, &mut visit)?; + } + } + _ => return Err(invalid()), + } + } + if !frames.bytes.is_empty() { + return Err(invalid()); + } + } + Ok(()) +} diff --git a/compiler-core/src/classfile/summary/forwarder.rs b/compiler-core/src/classfile/summary/forwarder.rs new file mode 100644 index 00000000..47a9355e --- /dev/null +++ b/compiler-core/src/classfile/summary/forwarder.rs @@ -0,0 +1,216 @@ +//! Prove exact static forwarding without a bytecode graph. +//! Permit local spills. Reject casts, branches, handlers, effects and reordered +//! arguments. Require identical caller and target descriptors. +use super::*; +use crate::classfile::{ByteReader, attributes::Instruction}; + +#[derive(Clone, Copy, PartialEq, Eq)] +struct Value { + source: u16, + kind: u8, +} + +fn signature(bytes: &[u8]) -> Option<(Vec<(usize, Value)>, u8)> { + let mut cursor = bytes.strip_prefix(b"(")?; + let mut params = Vec::new(); + let mut slot = 0; + while !cursor.starts_with(b")") { + let (&first, rest) = cursor.split_first()?; + cursor = rest; + let kind = match first { + b'Z' | b'B' | b'C' | b'S' | b'I' => 0, + b'J' => 1, + b'F' => 2, + b'D' => 3, + b'L' | b'[' => { + let mut element = first; + while element == b'[' { + let (&next, rest) = cursor.split_first()?; + element = next; + cursor = rest; + } + if element == b'L' { + let end = cursor.iter().position(|&b| b == b';')?; + cursor = &cursor[end + 1..]; + } else if !b"ZBCSIJFD".contains(&element) { + return None; + } + 4 + } + _ => return None, + }; + params.push(( + slot, + Value { + source: params.len() as u16, + kind, + }, + )); + slot += if kind == 1 || kind == 3 { 2 } else { 1 }; + if slot > 64 { + return None; + } + } + let result = match cursor.get(1)? { + b'Z' | b'B' | b'C' | b'S' | b'I' => 0, + b'J' => 1, + b'F' => 2, + b'D' => 3, + b'L' | b'[' => 4, + b'V' => 5, + _ => return None, + }; + Some((params, result)) +} + +pub(super) fn target<'a>( + bytes: &[u8], + descriptor: &[u8], + pool: &Pool<'a>, +) -> Option> { + let (params, result) = signature(descriptor)?; + let mut raw = Reader { bytes }; + raw.take(4).ok()?; + let length = raw.u32().ok()? as usize; + if length > 512 { + return None; + } + let bytes = raw.take(length).ok()?; + if raw.u16().ok()? != 0 { + return None; + } + let mut code = ByteReader::new(bytes); + let mut locals = [None; 64]; + for &(slot, value) in ¶ms { + locals[slot] = Some(value); + } + let mut stack = Vec::with_capacity(params.len() + 1); + let mut target = None; + while code.remaining() != 0 { + use Instruction::*; + let instruction = Instruction::from_bytes(&mut code).ok()?; + let load = match instruction { + Iload(i) => Some((i as usize, 0)), + Lload(i) => Some((i as usize, 1)), + Fload(i) => Some((i as usize, 2)), + Dload(i) => Some((i as usize, 3)), + Aload(i) => Some((i as usize, 4)), + Iload_0 => Some((0, 0)), + Iload_1 => Some((1, 0)), + Iload_2 => Some((2, 0)), + Iload_3 => Some((3, 0)), + Lload_0 => Some((0, 1)), + Lload_1 => Some((1, 1)), + Lload_2 => Some((2, 1)), + Lload_3 => Some((3, 1)), + Fload_0 => Some((0, 2)), + Fload_1 => Some((1, 2)), + Fload_2 => Some((2, 2)), + Fload_3 => Some((3, 2)), + Dload_0 => Some((0, 3)), + Dload_1 => Some((1, 3)), + Dload_2 => Some((2, 3)), + Dload_3 => Some((3, 3)), + Aload_0 => Some((0, 4)), + Aload_1 => Some((1, 4)), + Aload_2 => Some((2, 4)), + Aload_3 => Some((3, 4)), + _ => None, + }; + if let Some((slot, kind)) = load { + let value = (*locals.get(slot)?)?; + if value.kind != kind { + return None; + } + stack.push(value); + continue; + } + let store = match instruction { + Istore(i) => Some((i as usize, 0)), + Lstore(i) => Some((i as usize, 1)), + Fstore(i) => Some((i as usize, 2)), + Dstore(i) => Some((i as usize, 3)), + Astore(i) => Some((i as usize, 4)), + Istore_0 => Some((0, 0)), + Istore_1 => Some((1, 0)), + Istore_2 => Some((2, 0)), + Istore_3 => Some((3, 0)), + Lstore_0 => Some((0, 1)), + Lstore_1 => Some((1, 1)), + Lstore_2 => Some((2, 1)), + Lstore_3 => Some((3, 1)), + Fstore_0 => Some((0, 2)), + Fstore_1 => Some((1, 2)), + Fstore_2 => Some((2, 2)), + Fstore_3 => Some((3, 2)), + Dstore_0 => Some((0, 3)), + Dstore_1 => Some((1, 3)), + Dstore_2 => Some((2, 3)), + Dstore_3 => Some((3, 3)), + Astore_0 => Some((0, 4)), + Astore_1 => Some((1, 4)), + Astore_2 => Some((2, 4)), + Astore_3 => Some((3, 4)), + _ => None, + }; + if let Some((slot, kind)) = store { + let value = stack.pop()?; + if value.kind != kind { + return None; + } + *locals.get_mut(slot)? = Some(value); + continue; + } + match instruction { + Nop => {} + Goto(target) if usize::from(target) == code.position() => {} + Goto_w(target) if usize::try_from(target).ok() == Some(code.position()) => {} + Invokestatic(index) if target.is_none() => { + if !stack.iter().eq(params.iter().map(|(_, value)| value)) { + return None; + } + let Constant::Member(owner, member, true) = *pool.constants.get(index as usize)? + else { + return None; + }; + let callee = pool.member(owner, member).ok()?; + if callee.descriptor != descriptor { + return None; + } + target = Some(callee); + stack.clear(); + if result != 5 { + stack.push(Value { + source: u16::MAX, + kind: result, + }); + } + } + Ireturn | Lreturn | Freturn | Dreturn | Areturn | Return => { + let kind = match instruction { + Ireturn => 0, + Lreturn => 1, + Freturn => 2, + Dreturn => 3, + Areturn => 4, + _ => 5, + }; + if kind != result || code.remaining() != 0 { + return None; + } + if result != 5 + && stack.pop()? + != (Value { + source: u16::MAX, + kind: result, + }) + { + return None; + } + return if stack.is_empty() { target } else { None }; + } + _ => return None, + } + } + None +} diff --git a/java-linker/src/tests.rs b/java-linker/src/tests.rs index 2d64cbc3..174bc6f6 100644 --- a/java-linker/src/tests.rs +++ b/java-linker/src/tests.rs @@ -649,8 +649,12 @@ fn borrowed_metadata_matches_full_reader_and_rejects_truncation() { assert_eq!( summary::read(&bytes).unwrap(), summary::Summary { + private: false, + opaque_reflection: false, + method_demands: false, name: name.into(), - has_main: true + has_main: true, + carrier: None, } ); for length in 0..bytes.len() { From 1afb380c8be34163b6b9f391f167d1e61866abfb Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Wed, 30 Sep 2026 22:55:54 +1000 Subject: [PATCH 14/61] lower component storage operations --- compiler-core/src/analysis/arrays.rs | 175 +++++ compiler-core/src/analysis/mod.rs | 7 + compiler-core/src/analysis/origins.rs | 105 +++ compiler-core/src/analysis/users.rs | 73 ++ compiler-core/src/ir/body.rs | 303 ++++++++- compiler-core/src/ir/fold.rs | 74 ++- compiler-core/src/ir/ids.rs | 1 + compiler-core/src/ir/mod.rs | 4 + compiler-core/src/ir/remap.rs | 187 +++++- compiler-core/src/ir/shape.rs | 114 ++++ compiler-core/src/ir/types.rs | 114 +++- compiler-core/src/ir/verify/mod.rs | 46 +- compiler-core/src/ir/verify/types.rs | 623 +++++++++++++++++- compiler-core/src/jvm/abi.rs | 116 +++- compiler-core/src/jvm/frames/tests.rs | 44 ++ compiler-core/src/jvm/frames/transfer.rs | 19 +- compiler-core/src/jvm/select/addresses.rs | 449 +++++++++++++ compiler-core/src/jvm/select/allocate.rs | 11 +- compiler-core/src/jvm/select/arrays.rs | 111 +++- compiler-core/src/jvm/select/debug_tests.rs | 9 + compiler-core/src/jvm/select/fields.rs | 290 +++++++- compiler-core/src/jvm/select/fields/bytes.rs | 70 ++ compiler-core/src/jvm/select/forward.rs | 31 +- compiler-core/src/jvm/select/general.rs | 15 + compiler-core/src/jvm/select/memory.rs | 179 +++-- compiler-core/src/jvm/select/mod.rs | 81 ++- compiler-core/src/jvm/select/object_tests.rs | 2 - compiler-core/src/jvm/select/objects.rs | 37 +- .../src/jvm/select/representation.rs | 8 +- compiler-core/src/jvm/select/scalar.rs | 20 +- compiler-core/src/jvm/select/tests.rs | 152 ++++- compiler-core/src/jvm/select/unwind_tests.rs | 24 + compiler-core/src/jvm/select/views.rs | 109 +++ compiler-core/src/lib.rs | 1 + compiler-core/src/opt/fields_tests.rs | 1 - 35 files changed, 3319 insertions(+), 286 deletions(-) create mode 100644 compiler-core/src/analysis/arrays.rs create mode 100644 compiler-core/src/analysis/mod.rs create mode 100644 compiler-core/src/analysis/origins.rs create mode 100644 compiler-core/src/analysis/users.rs create mode 100644 compiler-core/src/ir/shape.rs create mode 100644 compiler-core/src/jvm/select/addresses.rs create mode 100644 compiler-core/src/jvm/select/fields/bytes.rs diff --git a/compiler-core/src/analysis/arrays.rs b/compiler-core/src/analysis/arrays.rs new file mode 100644 index 00000000..415a0679 --- /dev/null +++ b/compiler-core/src/analysis/arrays.rs @@ -0,0 +1,175 @@ +//! Prove that fresh primitive arrays have no decoded or encoded aliases. +//! Escape checks cover calls, storage, addresses and mixed joins. Local checks +//! also permit native writes before the first escape in a block. +use super::{NO_ORIGIN, origins}; +use crate::ir::*; +use crate::opt::Live; + +pub(crate) fn native_array_accesses( + body: &Body, + types: &Types, + live: &Live, +) -> Vec> { + let candidates = body + .instructions + .iter() + .enumerate() + .filter_map(|(id, inst)| { + if !live.instructions[id] || !matches!(inst.op, Op::NewArray(_)) { + return None; + } + let value = inst.result?; + let Type::Array(element) = types.get(body.value_type(value))? else { + return None; + }; + StorageSlot::scalar(element, types).map(|_| (value, element)) + }) + .collect::>(); + if candidates.is_empty() { + return Vec::new(); + } + let mut roots = vec![NO_ORIGIN; body.values.len()]; + for (index, &(value, _)) in candidates.iter().enumerate() { + roots[value.index()] = index as u32; + } + let origins = origins(body, &roots); + let mut escaped = vec![false; candidates.len()]; + for (id, inst) in body.instructions.iter().enumerate() { + if !live.instructions[id] { + continue; + } + inst.op.visit_uses(&body.args, |value| { + let origin = origins[value.index()]; + if origin == NO_ORIGIN { + return; + } + let element = candidates[origin as usize].1; + let safe = safe_use(body, *inst, value, element); + escaped[origin as usize] |= !safe; + }); + } + let mut escape = |value: ValueId| { + let origin = origins[value.index()]; + if origin != NO_ORIGIN { + escaped[origin as usize] = true; + } + }; + for edge in &body.edges { + for (&value, ¶m) in edge + .args + .iter() + .zip(&body.blocks[edge.target.index()].params) + { + if live.values[body.resolve(param).index()] + && origins[value.index()] != origins[param.index()] + { + escape(value); + } + } + } + for (index, block) in body.blocks.iter().enumerate() { + if live.blocks[index] { + block.terminator.unwrap().visit_uses(&mut escape); + } + } + // Permit native initialization before the first escape, including a later return. + // Restrict this proof to one block because other paths can expose the array. + let mut fresh = vec![None; body.blocks.len()]; + for block in &body.blocks { + if let Some(Terminator::Invoke { inst, normal, .. }) = block.terminator { + if let Some(value) = body.instructions[inst.index()].result { + let root = roots[value.index()]; + if root != NO_ORIGIN { + fresh[body.edges[normal.index()].target.index()] = Some(root); + } + } + } + } + let mut available = vec![None; candidates.len()]; + let mut native = vec![None; body.instructions.len()]; + for (index, block) in body.blocks.iter().enumerate() { + if !live.blocks[index] { + continue; + } + let current = Some(BlockId::new(index)); + if let Some(root) = fresh[index] { + available[root as usize] = current; + } + let invoke = match block.terminator { + Some(Terminator::Invoke { inst, .. }) => Some(inst), + _ => None, + }; + for &id in block.instructions.iter().chain(invoke.iter()) { + if !live.instructions[id.index()] { + continue; + } + let inst = body.instructions[id.index()]; + if let Some((array, element)) = access(body, inst) { + let origin = origins[array.index()]; + if origin != NO_ORIGIN + && candidates[origin as usize].1 == element + && (!escaped[origin as usize] || available[origin as usize] == current) + { + native[id.index()] = Some(element); + } + } + inst.op.visit_uses(&body.args, |value| { + let origin = origins[value.index()]; + if origin != NO_ORIGIN + && !safe_use(body, inst, value, candidates[origin as usize].1) + { + available[origin as usize] = None; + } + }); + if let Some(value) = inst.result { + let root = roots[value.index()]; + if root != NO_ORIGIN { + available[root as usize] = current; + } + } + } + } + native +} + +fn access(body: &Body, inst: Inst) -> Option<(ValueId, TypeId)> { + Some(match inst.op { + Op::ArrayGet { array, .. } => (array, body.value_type(inst.result?)), + Op::ArraySet { array, value, .. } | Op::ArrayFill { array, value } => { + (array, body.value_type(value)) + } + Op::ViewGet(parts) => ( + body.args[parts.start as usize], + body.value_type(inst.result?), + ), + Op::ViewSet { parts, value } => (body.args[parts.start as usize], body.value_type(value)), + _ => return None, + }) +} + +fn safe_use(body: &Body, inst: Inst, value: ValueId, element: TypeId) -> bool { + match inst.op { + Op::Reinterpret(_) | Op::Refine(_) | Op::ArrayLength(_) => true, + Op::ArrayGet { array, .. } => { + array == value && inst.result.is_some_and(|r| body.value_type(r) == element) + } + Op::ArraySet { + array, + value: stored, + .. + } + | Op::ArrayFill { + array, + value: stored, + } => array == value && body.value_type(stored) == element, + Op::ViewGet(parts) => { + body.args[parts.start as usize] == value + && inst.result.is_some_and(|r| body.value_type(r) == element) + } + Op::ViewSet { + parts, + value: stored, + } => body.args[parts.start as usize] == value && body.value_type(stored) == element, + _ => false, + } +} diff --git a/compiler-core/src/analysis/mod.rs b/compiler-core/src/analysis/mod.rs new file mode 100644 index 00000000..2728e883 --- /dev/null +++ b/compiler-core/src/analysis/mod.rs @@ -0,0 +1,7 @@ +//! Bounded facts shared by representation and storage optimizations. +mod origins; +pub(crate) use origins::{NO_ORIGIN, origins}; +mod users; +pub(crate) use users::ValueUsers; +mod arrays; +pub(crate) use arrays::native_array_accesses; diff --git a/compiler-core/src/analysis/origins.rs b/compiler-core/src/analysis/origins.rs new file mode 100644 index 00000000..d68784d2 --- /dev/null +++ b/compiler-core/src/analysis/origins.rs @@ -0,0 +1,105 @@ +//! Allocation origins through identity annotations and equal-origin CFG joins. +use crate::ir::*; + +pub(crate) const NO_ORIGIN: u32 = u32::MAX; +const UNKNOWN: u32 = u32::MAX - 1; + +/// Track allocation identity through aliases and joins. +/// Only roots in an entry block without incoming edges can cross joins. +/// Conflicting inputs make joins unknown. Each value changes at most twice. +pub(crate) fn origins(body: &Body, roots: &[u32]) -> Vec { + let count = body.values.len(); + let mut stable = vec![ + false; + roots + .iter() + .filter(|&&r| r != NO_ORIGIN) + .max() + .map_or(0, |r| *r as usize + 1) + ]; + if !body.edges.iter().any(|edge| edge.target == body.entry) { + let entry = &body.blocks[body.entry.index()]; + let invoke = match entry.terminator { + Some(Terminator::Invoke { inst, .. }) => Some(inst), + _ => None, + }; + for &id in entry.instructions.iter().chain(invoke.iter()) { + if let Some(value) = body.instructions[id.index()].result { + let root = roots[value.index()]; + if root != NO_ORIGIN { + stable[root as usize] = true; + } + } + } + } + let mut state = roots.to_vec(); + let mut joins = vec![false; count]; + let mut users = super::ValueUsers::new(count); + for (index, value) in body.values.iter().enumerate() { + if roots[index] != NO_ORIGIN { + continue; + } + let source = match value.def { + ValueDef::Alias(source) => Some(source), + ValueDef::Inst(inst) => match body.instructions[inst.index()].op { + Op::Reinterpret(source) | Op::Refine(source) => Some(source), + _ => None, + }, + ValueDef::Param(block) if block != body.entry => { + joins[index] = true; + state[index] = UNKNOWN; + None + } + _ => None, + }; + if let Some(source) = source { + state[index] = UNKNOWN; + users.connect(source, index); + } + } + for edge in &body.edges { + for (&source, &target) in edge + .args + .iter() + .zip(&body.blocks[edge.target.index()].params) + { + if joins[target.index()] { + users.connect(source, target.index()); + } + } + } + let mut pending = (0..count) + .filter(|&i| state[i] != UNKNOWN) + .collect::>(); + for phase in 0..2 { + while let Some(source) = pending.pop() { + for target in users.users(source) { + let mut incoming = state[source]; + if joins[target] && incoming != NO_ORIGIN && !stable[incoming as usize] { + incoming = NO_ORIGIN; + } + let previous = state[target]; + let merged = if previous == UNKNOWN || previous == incoming { + incoming + } else { + NO_ORIGIN + }; + if merged != previous { + state[target] = merged; + pending.push(target); + } + } + } + if phase == 0 { + // A cycle without a defining root cannot prove allocation identity, + // even if another incoming edge has a known root. + for (index, value) in state.iter_mut().enumerate() { + if *value == UNKNOWN { + *value = NO_ORIGIN; + pending.push(index); + } + } + } + } + state +} diff --git a/compiler-core/src/analysis/users.rs b/compiler-core/src/analysis/users.rs new file mode 100644 index 00000000..ea9c7bed --- /dev/null +++ b/compiler-core/src/analysis/users.rs @@ -0,0 +1,73 @@ +//! Store reverse value dependencies in two flat allocations. +//! Representation passes share this graph instead of per-value vectors. +use crate::ir::ValueId; + +const END: u32 = u32::MAX; + +pub(crate) struct ValueUsers { + heads: Vec, + edges: Vec<(u32, u32)>, +} + +impl ValueUsers { + pub fn new(count: usize) -> Self { + Self { + heads: vec![END; count], + edges: Vec::new(), + } + } + + pub fn connect(&mut self, source: ValueId, user: usize) { + let next = self.heads[source.index()]; + self.heads[source.index()] = u32::try_from(self.edges.len()).expect("too many value uses"); + self.edges + .push((next, u32::try_from(user).expect("too many values"))); + } + + pub fn users(&self, source: usize) -> impl Iterator + '_ { + let mut next = self.heads[source]; + std::iter::from_fn(move || { + if next == END { + return None; + } + let (following, user) = self.edges[next as usize]; + next = following; + Some(user as usize) + }) + } + + /// Reject a candidate if any input is ineligible. + /// Queue each rejected value once, including values in cycles. + pub fn close(&self, eligible: &mut [bool]) { + let mut pending = eligible + .iter() + .enumerate() + .filter_map(|(index, &yes)| (!yes).then_some(index as u32)) + .collect::>(); + while let Some(index) = pending.pop() { + for user in self.users(index as usize) { + if std::mem::replace(&mut eligible[user], false) { + pending.push(user as u32); + } + } + } + } + + /// Propagate proven identity aliases. Joins retain their own components. + /// They must not reuse the components of one input. + pub fn propagate(&self, eligible: &[bool], known: &mut [Option]) { + let mut pending = known + .iter() + .enumerate() + .filter_map(|(index, value)| value.is_some().then_some(index as u32)) + .collect::>(); + while let Some(index) = pending.pop() { + for user in self.users(index as usize) { + if eligible[user] && known[user].is_none() { + known[user] = known[index as usize]; + pending.push(user as u32); + } + } + } + } +} diff --git a/compiler-core/src/ir/body.rs b/compiler-core/src/ir/body.rs index 22396348..67fd6011 100644 --- a/compiler-core/src/ir/body.rs +++ b/compiler-core/src/ir/body.rs @@ -25,6 +25,15 @@ pub enum CallKind { Indirect, } +/// Storage operations for the JVM platform allocator. +/// Custom Rust allocators retain their calls and effects. +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +pub enum HeapOp { + Allocate, + Reallocate, + Deallocate, +} + /// A body-local call target. Signatures describe value-bearing JVM arguments; /// source-language unit arguments never enter the operand pool. #[derive(Clone, Debug, PartialEq, Eq, Hash)] @@ -42,8 +51,6 @@ pub struct FieldRef { pub name: String, pub ty: TypeId, pub is_static: bool, - /// Generated Rust pointer fields store a base plus two displacement fields. - pub relative_pointer: bool, } /// A field view preserves Rust byte layout and the runtime allocation identity. @@ -73,7 +80,7 @@ impl StorageSlot { }; let size = match scalar { Bool | I8 | U8 => 1, - I16 | U16 => 2, + I16 | U16 | F16 => 2, I32 | U32 | F32 => 4, I64 | U64 | F64 => 8, _ => return None, @@ -116,11 +123,40 @@ pub enum Op { Opaque(ValueId), Neg(ValueId), Cast(ValueId), + /// Refine a reference type after representation analysis proves the cast. + /// Emit the JVM cast only if a consumer needs the result. + Refine(ValueId), /// Representation adaptation at a JVM ABI boundary (boxing, views, casts). Adapt(ValueId), /// Same physical JVM carrier with a different semantic type annotation. Reinterpret(ValueId), NewArray(ValueId), + /// Primitive repeat initialization, retaining the scalar value unboxed. + ArrayFill { + array: ValueId, + value: ValueId, + }, + /// Fresh, naturally aligned scalar storage with exact initial contents. + /// Lower to a primitive array only if its address survives SSA promotion. + ScalarCell(ValueId), + /// Allocate [size, alignment], reallocate [root, offset, old size, + /// alignment, new size], or release [root, offset]. Return storage roots. + /// Keep address metadata separate. + Heap { + operation: HeapOp, + args: List, + }, + /// Select typed storage for a scalar field. + /// Use a boundary carrier if the root requires general memory access. + ProjectRoot { + address: List, + projection: ProjectionId, + }, + ProjectOffset { + root: ValueId, + base: ValueId, + offset: ValueId, + }, FunctionPointer { signature: MethodId, target: MethodId, @@ -132,6 +168,12 @@ pub enum Op { }, AddressOfSlot(SlotId), Load(ValueId), + /// An owned aggregate snapshot, distinct from a live decoded memory view. + LoadCopy(ValueId), + /// An owned Rust value snapshot, independent of subsequent source writes. + CopyValue(ValueId), + /// Complete mutations to a live decoded view using the original load owner. + Commit(ValueId), Store { pointer: ValueId, value: ValueId, @@ -145,11 +187,44 @@ pub enum Op { base: ValueId, projection: ProjectionId, }, + /// Copy an owned field without constructing an intermediate address. + /// Keep the read and snapshot together, including inside an Invoke. + LoadFieldCopy { + base: ValueId, + projection: ProjectionId, + }, StoreField { base: ValueId, projection: ProjectionId, value: ValueId, }, + LoadFieldPart { + base: ValueId, + projection: ProjectionId, + index: u8, + }, + /// Read a field through a storage root and displacement. + /// The index selects one physical component of a borrowed field. + LoadStorageField { + address: List, + projection: ProjectionId, + index: Option, + }, + LoadStorageFieldCopy { + address: List, + projection: ProjectionId, + }, + /// [root, displacement, value/components] for a projected field write. + StoreStorageField { + args: List, + projection: ProjectionId, + split: bool, + }, + StoreFieldParts { + base: ValueId, + projection: ProjectionId, + parts: List, + }, Offset { pointer: ValueId, offset: ValueId, @@ -178,24 +253,140 @@ pub enum Op { ArrayGet { array: ValueId, index: ValueId, + /// Compiler-owned ABI and outline scratch, inaccessible to Rust code. + native: bool, }, ArraySet { array: ValueId, index: ValueId, value: ValueId, + /// Same unaliased scratch-storage contract as ArrayGet::native. + native: bool, + }, + /// Slice access through [backing, start, index], without a view carrier. + ViewGet(List), + ViewSet { + parts: List, + value: ValueId, }, /// Logical Rust length, including zero-sized slices larger than JVM arrays. Length(ValueId), - /// Immutable fat-pointer carrier; data retains allocation identity and view. + /// Address with physical components [storage root, byte offset]. + AddressPack(List), + /// Scalar [payload, discriminant], boxed only at an opaque boundary. + TaggedPack(List), + TaggedPart { + value: ValueId, + index: u8, + }, + /// Exact Rust view layout belongs to the location, not its JVM carrier. + RetypeAddress { + pointer: ValueId, + size: u32, + codec: Option, + }, + /// A component location whose pointee layout is exact at this use. + TypedAddressPack { + parts: List, + size: u32, + codec: Option, + }, + /// An owned aggregate load with a statically known view layout. + LoadTypedCopy { + parts: List, + size: u32, + codec: Option, + }, + /// Parts are [root, displacement, optional JVM class name]. + /// Class and interface results request their type when no name is supplied. + /// Object and array results use getObject's null target. + /// No later consumer can commit the borrowed view's binding. + LoadTyped { + parts: List, + size: u32, + codec: Option, + }, + /// [root, displacement, value] with an exact source-language layout. + StoreTyped { + parts: List, + size: u32, + codec: Option, + }, + AddressEqual { + left: ValueId, + right: ValueId, + }, + LocationEqual(List), + /// Compare unsigned data addresses and return a signed comparison code. + AddressCompare { + left: ValueId, + right: ValueId, + }, + LocationCompare(List), + /// Discriminant of a valid null-niche pointer enum (None = 0, Some = 1). + AddressTag(ValueId), + LocationTag(List), + AddressPart { + address: ValueId, + index: u8, + }, + LoadAddress(List), + LoadAddressCopy(List), + /// [source root, source offset, destination root, destination offset, bytes]. + CopyStorage { + parts: List, + // Intern exact layouts instead of enlarging every instruction. + layouts: [TypeId; 2], + nonoverlapping: bool, + }, + StoreAddress { + parts: List, + value: ValueId, + }, + /// The authoritative primitive array backing an addressable scalar local. + SlotRoot(SlotId), View { data: ValueId, length: ValueId, }, + /// Boundary materialization from [backing, start, logical length]. + ViewPack(List), + /// Read one physical component without constructing an address. + ViewPart { + view: ValueId, + index: u8, + }, ViewData { view: ValueId, size: u32, codec: Option, }, + /// Extract an address from [backing, start, logical length]. + /// No view object is needed, including at memory boundaries. + ViewAddress { + parts: List, + size: u32, + codec: Option, + }, + /// Normalize a nonzero-sized slice backing to retain its exact element layout. + /// Keep the element displacement separate. + ViewRoot { + backing: ValueId, + size: u32, + codec: Option, + }, + /// Extract a view backing or element start from a scalar location. + /// Do not allocate a thin-pointer carrier. + AddressViewPart { + address: ValueId, + index: u8, + }, + TypedAddressViewPart { + parts: List, + size: u32, + codec: Option, + index: u8, + }, } impl Op { @@ -203,16 +394,32 @@ impl Op { /// scalar computations and literal loads do not. pub fn may_throw(self, body: &Body, types: &Types) -> bool { match self { - Self::Constant(id) => matches!(body.constants[id.index()], Constant::External { .. }), + Self::Constant(id) => matches!( + body.constants[id.index()], + Constant::External { pure: false, .. } + ), + // Native scratch is initialized and accessed within its declared bounds. + // Unused reads have no Rust-visible effect. + Self::ArrayGet { native: true, .. } => false, Self::Nop | Self::Exception + | Self::AddressTag(_) + | Self::LocationTag(_) | Self::Reinterpret(_) + | Self::TaggedPack(_) + | Self::TaggedPart { .. } + | Self::Refine(_) | Self::Not(_) | Self::Neg(_) | Self::Bit { .. } | Self::Overflow { .. } => false, // A typed field address has no source-language effects until used. - Self::Project { .. } => false, + Self::Project { .. } + | Self::ProjectRoot { .. } + | Self::ProjectOffset { .. } + | Self::RetypeAddress { size: 1.., .. } + | Self::ViewRoot { .. } + | Self::TypedAddressPack { .. } => false, Self::Binary { op: BinaryOp::Div | BinaryOp::Rem, left, @@ -229,32 +436,69 @@ impl Op { } pub fn visit_uses(self, args: &[ValueId], mut visit: impl FnMut(ValueId)) { match self { - Self::Binary { left, right, .. } => { + Self::ProjectRoot { address, .. } => { + for &value in &args[address.range()] { + visit(value); + } + } + Self::ProjectOffset { root, base, offset } => { + visit(root); + visit(base); + visit(offset); + } + Self::AddressViewPart { address, .. } => visit(address), + Self::ViewRoot { backing, .. } => visit(backing), + Self::RetypeAddress { pointer, .. } => visit(pointer), + Self::Binary { left, right, .. } + | Self::AddressEqual { left, right } + | Self::AddressCompare { left, right } => { visit(left); visit(right); } - Self::Not(v) + Self::TaggedPart { value: v, .. } + | Self::AddressTag(v) + | Self::Not(v) | Self::Opaque(v) | Self::Bit { value: v, .. } | Self::Neg(v) | Self::Cast(v) + | Self::Refine(v) | Self::Adapt(v) | Self::Reinterpret(v) | Self::NewArray(v) + | Self::ScalarCell(v) | Self::ArrayLength(v) | Self::Load(v) + | Self::CopyValue(v) + | Self::LoadCopy(v) + | Self::Commit(v) | Self::Length(v) => visit(v), Self::StoreSlot { value, .. } | Self::SetStatic { value, .. } => visit(value), - Self::Store { pointer, value } => { + Self::Store { pointer, value } + | Self::ArrayFill { + array: pointer, + value, + } => { visit(pointer); visit(value); } - Self::Project { base, .. } | Self::LoadField { base, .. } => visit(base), + Self::Project { base, .. } + | Self::LoadField { base, .. } + | Self::LoadFieldCopy { base, .. } + | Self::LoadFieldPart { base, .. } => visit(base), + Self::StoreFieldParts { base, parts, .. } => { + visit(base); + for &value in &args[parts.range()] { + visit(value); + } + } Self::StoreField { base, value, .. } => { visit(base); visit(value); } - Self::ViewData { view, .. } => visit(view), + Self::ViewData { view, .. } + | Self::ViewPart { view, .. } + | Self::AddressPart { address: view, .. } => visit(view), Self::View { data, length } => { visit(data); visit(length); @@ -265,17 +509,44 @@ impl Op { visit(pointer); visit(offset); } - Self::Call { args: list, .. } | Self::Overflow { args: list, .. } => { + Self::Call { args: list, .. } + | Self::Heap { args: list, .. } + | Self::StoreStorageField { args: list, .. } + | Self::LoadStorageField { address: list, .. } + | Self::LoadStorageFieldCopy { address: list, .. } + | Self::Overflow { args: list, .. } + | Self::TaggedPack(list) + | Self::ViewPack(list) + | Self::ViewGet(list) + | Self::AddressPack(list) + | Self::LoadAddress(list) + | Self::LoadAddressCopy(list) + | Self::CopyStorage { parts: list, .. } + | Self::LoadTypedCopy { parts: list, .. } + | Self::LoadTyped { parts: list, .. } + | Self::StoreTyped { parts: list, .. } + | Self::TypedAddressPack { parts: list, .. } + | Self::TypedAddressViewPart { parts: list, .. } + | Self::LocationTag(list) + | Self::LocationEqual(list) + | Self::LocationCompare(list) + | Self::ViewAddress { parts: list, .. } => { for &arg in &args[list.range()] { visit(arg); } } + Self::StoreAddress { parts, value } | Self::ViewSet { parts, value } => { + for &part in &args[parts.range()] { + visit(part); + } + visit(value); + } Self::GetField { object, .. } => visit(object), Self::SetField { object, value, .. } => { visit(object); visit(value); } - Self::ArrayGet { array, index } => { + Self::ArrayGet { array, index, .. } => { visit(array); visit(index); } @@ -283,6 +554,7 @@ impl Op { array, index, value, + .. } => { visit(array); visit(index); @@ -293,6 +565,7 @@ impl Op { | Self::Exception | Self::LoadSlot(_) | Self::AddressOfSlot(_) + | Self::SlotRoot(_) | Self::GetStatic(_) | Self::FunctionPointer { .. } => {} } @@ -318,6 +591,8 @@ pub enum Constant { External { index: u32, ty: TypeId, + /// A literal with no initialization or runtime helper effects. + pure: bool, }, } @@ -447,7 +722,7 @@ impl Body { } value } - pub(super) fn resolve_mut(&mut self, mut value: ValueId) -> ValueId { + pub(crate) fn resolve_mut(&mut self, mut value: ValueId) -> ValueId { let root = self.resolve(value); while let ValueDef::Alias(next) = self.values[value.index()].def { self.values[value.index()].def = ValueDef::Alias(root); diff --git a/compiler-core/src/ir/fold.rs b/compiler-core/src/ir/fold.rs index 9599a89c..3ecd7b58 100644 --- a/compiler-core/src/ir/fold.rs +++ b/compiler-core/src/ir/fold.rs @@ -1,41 +1,45 @@ use super::*; use crate::scalar::{BinaryFold, Scalar, fold_binary}; -impl Builder<'_> { - pub(super) fn scalar_value(&self, value: ValueId) -> Option { - let value = self.body.resolve(value); - let ValueDef::Inst(inst) = self.body.values[value.index()].def else { +pub(crate) enum Folded { + Value(ValueId), + Constant(Scalar), +} + +impl Body { + pub(crate) fn scalar_value(&self, value: ValueId) -> Option { + let value = self.resolve(value); + let ValueDef::Inst(inst) = self.values[value.index()].def else { return None; }; - let Op::Constant(constant) = self.body.instructions[inst.index()].op else { + let Op::Constant(constant) = self.instructions[inst.index()].op else { return None; }; - let Constant::Scalar(scalar) = self.body.constants[constant.index()] else { + let Constant::Scalar(scalar) = self.constants[constant.index()] else { return None; }; Some(scalar) } - pub(super) fn fold(&mut self, op: Op, result: TypeId) -> Option { + pub(crate) fn fold(&self, types: &Types, op: Op, result: TypeId) -> Option { if let Op::Length(view) = op { - let view = self.body.resolve(view); - if let ValueDef::Inst(inst) = self.body.values[view.index()].def - && let Op::View { length, .. } = self.body.instructions[inst.index()].op - && self.body.value_type(length) == result + let view = self.resolve(view); + if let ValueDef::Inst(inst) = self.values[view.index()].def + && let Op::View { length, .. } = self.instructions[inst.index()].op + && self.value_type(length) == result { - return Some(length); + return Some(Folded::Value(length)); } } - let Some(Type::Scalar(result_ty)) = self.types.get(result) else { + let Some(Type::Scalar(result_ty)) = types.get(result) else { return None; }; let constant = match op { Op::Binary { op, left, right } => { - let Some(Type::Scalar(left_ty)) = self.types.get(self.body.value_type(left)) else { + let Some(Type::Scalar(left_ty)) = types.get(self.value_type(left)) else { return None; }; - let Some(Type::Scalar(right_ty)) = self.types.get(self.body.value_type(right)) - else { + let Some(Type::Scalar(right_ty)) = types.get(self.value_type(right)) else { return None; }; let result = fold_binary( @@ -44,33 +48,57 @@ impl Builder<'_> { right_ty, self.scalar_value(left), self.scalar_value(right), - || self.body.resolve(left) == self.body.resolve(right), + || self.resolve(left) == self.resolve(right), )?; match result { - BinaryFold::Left if result_ty == left_ty => return Some(left), - BinaryFold::Right if result_ty == right_ty => return Some(right), + BinaryFold::Left if result_ty == left_ty => return Some(Folded::Value(left)), + BinaryFold::Right if result_ty == right_ty => { + return Some(Folded::Value(right)); + } BinaryFold::Constant(value) => value, _ => return None, } } + Op::Reinterpret(value) => { + // Full JVM integer-width signedness is only an annotation. + // Narrow int carriers can have different extension semantics. + let value = self.scalar_value(value)?; + let (width, _) = value.ty().integer()?; + if !matches!(width, 32 | 64) || result_ty.integer()?.0 != width { + return None; + } + value.cast(result_ty)? + } Op::Not(value) => self.scalar_value(value)?.not()?, Op::Bit { op, value } => self.scalar_value(value)?.bit(op)?, Op::Overflow { op, args } => { - let args = &self.body.args[args.range()]; + let args = &self.args[args.range()]; let a = self.scalar_value(args[0])?; let b = self.scalar_value(args[1])?; Scalar::boolean(a.overflows(op, b)?) } Op::Neg(value) => self.scalar_value(value)?.neg()?, Op::Cast(value) => { - if self.body.value_type(value) == result { - return Some(value); + if self.value_type(value) == result { + return Some(Folded::Value(value)); } self.scalar_value(value)?.cast(result_ty)? } _ => return None, }; - (constant.ty() == result_ty).then(|| self.constant(result, constant)) + (constant.ty() == result_ty).then_some(Folded::Constant(constant)) + } +} + +impl Builder<'_> { + pub(super) fn scalar_value(&self, value: ValueId) -> Option { + self.body.scalar_value(value) + } + pub(super) fn fold(&mut self, op: Op, result: TypeId) -> Option { + match self.body.fold(self.types, op, result)? { + Folded::Value(value) => Some(value), + Folded::Constant(value) => Some(self.constant(result, value)), + } } } diff --git a/compiler-core/src/ir/ids.rs b/compiler-core/src/ir/ids.rs index faeed842..5e88a514 100644 --- a/compiler-core/src/ir/ids.rs +++ b/compiler-core/src/ir/ids.rs @@ -23,6 +23,7 @@ ids!( EdgeId, VariableId, TypeId, + LayoutId, SymbolId, ConstId, SlotId, diff --git a/compiler-core/src/ir/mod.rs b/compiler-core/src/ir/mod.rs index 9194851a..725da86e 100644 --- a/compiler-core/src/ir/mod.rs +++ b/compiler-core/src/ir/mod.rs @@ -4,6 +4,7 @@ pub use debug::*; mod body; mod builder; mod fold; +pub(crate) use fold::Folded; mod ids; mod parameters; mod remap; @@ -22,3 +23,6 @@ mod tests; #[cfg(test)] mod storage_tests; + +mod shape; +pub use shape::ComponentShape; diff --git a/compiler-core/src/ir/remap.rs b/compiler-core/src/ir/remap.rs index 97afb8a3..414218e3 100644 --- a/compiler-core/src/ir/remap.rs +++ b/compiler-core/src/ir/remap.rs @@ -14,6 +14,39 @@ impl Op { pub fn remap(self, map: &mut impl Remap) -> Self { use Op::*; match self { + AddressViewPart { address, index } => AddressViewPart { + address: map.value(address), + index, + }, + RetypeAddress { + pointer, + size, + codec, + } => RetypeAddress { + pointer: map.value(pointer), + size, + codec, + }, + TypedAddressPack { parts, size, codec } => TypedAddressPack { + parts: map.args(parts), + size, + codec, + }, + LoadTypedCopy { parts, size, codec } => LoadTypedCopy { + parts: map.args(parts), + size, + codec, + }, + LoadTyped { parts, size, codec } => LoadTyped { + parts: map.args(parts), + size, + codec, + }, + StoreTyped { parts, size, codec } => StoreTyped { + parts: map.args(parts), + size, + codec, + }, Nop => Nop, Constant(c) => Constant(map.constant(c)), Exception => Exception, @@ -34,12 +67,33 @@ impl Op { }, Opaque(v) => Opaque(map.value(v)), Cast(v) => Cast(map.value(v)), + Refine(v) => Refine(map.value(v)), Adapt(v) => Adapt(map.value(v)), Reinterpret(v) => Reinterpret(map.value(v)), NewArray(v) => NewArray(map.value(v)), + ScalarCell(v) => ScalarCell(map.value(v)), + Heap { operation, args } => Heap { + operation, + args: map.args(args), + }, + ProjectRoot { + address, + projection, + } => ProjectRoot { + address: map.args(address), + projection: map.projection(projection), + }, + ProjectOffset { root, base, offset } => ProjectOffset { + root: map.value(root), + base: map.value(base), + offset: map.value(offset), + }, ArrayLength(v) => ArrayLength(map.value(v)), Length(v) => Length(map.value(v)), Load(v) => Load(map.value(v)), + LoadCopy(v) => LoadCopy(map.value(v)), + CopyValue(v) => CopyValue(map.value(v)), + Commit(v) => Commit(map.value(v)), FunctionPointer { signature, target } => FunctionPointer { signature: map.method(signature), target: map.method(target), @@ -62,6 +116,53 @@ impl Op { base: map.value(base), projection: map.projection(projection), }, + LoadFieldCopy { base, projection } => LoadFieldCopy { + base: map.value(base), + projection: map.projection(projection), + }, + LoadStorageFieldCopy { + address, + projection, + } => LoadStorageFieldCopy { + address: map.args(address), + projection: map.projection(projection), + }, + LoadStorageField { + address, + projection, + index, + } => LoadStorageField { + address: map.args(address), + projection: map.projection(projection), + index, + }, + StoreStorageField { + args, + projection, + split, + } => StoreStorageField { + args: map.args(args), + projection: map.projection(projection), + split, + }, + LoadFieldPart { + base, + projection, + index, + } => LoadFieldPart { + base: map.value(base), + projection: map.projection(projection), + index, + }, + StoreFieldParts { + base, + projection, + parts, + } => StoreFieldParts { + base: map.value(base), + projection: map.projection(projection), + parts: map.args(parts), + }, StoreField { base, projection, @@ -105,7 +206,12 @@ impl Op { field: map.field(field), value: map.value(value), }, - ArrayGet { array, index } => ArrayGet { + ArrayGet { + array, + index, + native, + } => ArrayGet { + native, array: map.value(array), index: map.value(index), }, @@ -113,7 +219,9 @@ impl Op { array, index, value, + native, } => ArraySet { + native, array: map.value(array), index: map.value(index), value: map.value(value), @@ -122,11 +230,88 @@ impl Op { data: map.value(data), length: map.value(length), }, + ViewPack(parts) => ViewPack(map.args(parts)), + TaggedPack(parts) => TaggedPack(map.args(parts)), + TaggedPart { value, index } => TaggedPart { + value: map.value(value), + index, + }, + ArrayFill { array, value } => ArrayFill { + array: map.value(array), + value: map.value(value), + }, + ViewGet(parts) => ViewGet(map.args(parts)), + ViewSet { parts, value } => ViewSet { + parts: map.args(parts), + value: map.value(value), + }, + AddressPack(parts) => AddressPack(map.args(parts)), + AddressEqual { left, right } => AddressEqual { + left: map.value(left), + right: map.value(right), + }, + AddressCompare { left, right } => AddressCompare { + left: map.value(left), + right: map.value(right), + }, + AddressTag(value) => AddressTag(map.value(value)), + LocationTag(parts) => LocationTag(map.args(parts)), + LocationEqual(parts) => LocationEqual(map.args(parts)), + LocationCompare(parts) => LocationCompare(map.args(parts)), + LoadAddress(parts) => LoadAddress(map.args(parts)), + LoadAddressCopy(parts) => LoadAddressCopy(map.args(parts)), + CopyStorage { + parts, + layouts, + nonoverlapping, + } => CopyStorage { + parts: map.args(parts), + layouts, + nonoverlapping, + }, + StoreAddress { parts, value } => StoreAddress { + parts: map.args(parts), + value: map.value(value), + }, + AddressPart { address, index } => AddressPart { + address: map.value(address), + index, + }, + SlotRoot(slot) => SlotRoot(map.slot(slot)), + ViewPart { view, index } => ViewPart { + view: map.value(view), + index, + }, ViewData { view, size, codec } => ViewData { view: map.value(view), size, codec, }, + TypedAddressViewPart { + parts, + size, + codec, + index, + } => TypedAddressViewPart { + parts: map.args(parts), + size, + codec, + index, + }, + ViewAddress { parts, size, codec } => ViewAddress { + parts: map.args(parts), + size, + codec, + }, + ViewRoot { + backing, + size, + codec, + } => ViewRoot { + backing: map.value(backing), + size, + codec, + }, } } } diff --git a/compiler-core/src/ir/shape.rs b/compiler-core/src/ir/shape.rs new file mode 100644 index 00000000..aac33418 --- /dev/null +++ b/compiler-core/src/ir/shape.rs @@ -0,0 +1,114 @@ +//! Describe physical components while retaining semantic pointee types. +//! Emit a boundary carrier only when a consumer needs one object. +use super::*; +use crate::scalar::ScalarType; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum ComponentShape { + View, + TaggedI64, + Address, + /// This root retains the allocation's runtime view layout. + /// A bare primitive array cannot replace it, unlike scalar address roots. + StorageAddress, +} + +impl ComponentShape { + pub fn view_carrier(types: &Types, ty: TypeId) -> bool { + match types.get(ty) { + Some(Type::Slice(_) | Type::Str) => true, + Some(Type::Class(name)) => matches!( + types.symbol_name(name), + Some("org/rustlang/runtime/SliceView" | "org/rustlang/runtime/Utf8View") + ), + _ => false, + } + } + /// Runtime helpers can erase the semantic pointee to a carrier class. + /// This annotation change does not require a carrier allocation. + pub fn accepts_annotation(self, types: &Types, ty: TypeId) -> bool { + if Self::of(types, ty) == Some(self) { + return true; + } + let Some(Type::Class(name)) = types.get(ty) else { + return false; + }; + match types.symbol_name(name) { + Some("java/lang/Object") => true, + Some("org/rustlang/runtime/Pointer") => self.is_address(), + Some("org/rustlang/runtime/SliceView" | "org/rustlang/runtime/Utf8View") => { + self == Self::View + } + _ => false, + } + } + pub fn len(self) -> usize { + match self { + Self::View => 3, + Self::TaggedI64 => 2, + Self::Address | Self::StorageAddress => 2, + } + } + pub fn of(types: &Types, ty: TypeId) -> Option { + match types.get(ty)? { + Type::Slice(_) | Type::Str => Some(Self::View), + Type::TaggedI64 => Some(Self::TaggedI64), + Type::Pointer(inner) if StorageSlot::scalar(inner, types).is_some() => { + Some(Self::Address) + } + Type::Pointer(_) => Some(Self::StorageAddress), + _ => None, + } + } + pub fn parts(self, types: &mut Types) -> impl ExactSizeIterator + use<> { + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let (parts, count) = match self { + Self::View => ( + [ + object, + types.scalar(ScalarType::I32), + types.scalar(ScalarType::U64), + ], + 3, + ), + Self::TaggedI64 => { + let long = types.scalar(ScalarType::I64); + ([long, long, object], 2) + } + Self::Address | Self::StorageAddress => { + ([object, types.scalar(ScalarType::I64), object], 2) + } + }; + parts.into_iter().take(count) + } + pub fn slots(self) -> usize { + match self { + Self::View | Self::TaggedI64 => 4, + Self::Address | Self::StorageAddress => 3, + } + } + pub fn pack(self, parts: List) -> Op { + match self { + Self::View => Op::ViewPack(parts), + Self::TaggedI64 => Op::TaggedPack(parts), + Self::Address | Self::StorageAddress => Op::AddressPack(parts), + } + } + pub fn part(self, value: ValueId, index: u8) -> Op { + match self { + Self::View => Op::ViewPart { view: value, index }, + Self::TaggedI64 => Op::TaggedPart { value, index }, + Self::Address | Self::StorageAddress => Op::AddressPart { + address: value, + index, + }, + } + } + pub fn is_borrowed(self) -> bool { + self != Self::TaggedI64 + } + pub fn is_address(self) -> bool { + matches!(self, Self::Address | Self::StorageAddress) + } +} diff --git a/compiler-core/src/ir/types.rs b/compiler-core/src/ir/types.rs index fcf5659a..b73b40e6 100644 --- a/compiler-core/src/ir/types.rs +++ b/compiler-core/src/ir/types.rs @@ -1,4 +1,4 @@ -use super::{SymbolId, TypeId}; +use super::{LayoutId, SymbolId, TypeId}; use crate::scalar::ScalarType; use rustc_hash::FxHashMap; use std::sync::Arc; @@ -8,6 +8,8 @@ pub enum Type { Unit, /// Body-local identity for an uninspected pointee. Never a JVM value. Opaque(u32), + /// Exact source-language layout of an address pointee. Never a JVM value. + Layout(LayoutId), Scalar(ScalarType), Class(SymbolId), Interface(SymbolId), @@ -15,6 +17,8 @@ pub enum Type { Array(TypeId), Slice(TypeId), Str, + /// A 64-bit payload and independent discriminant. + TaggedI64, } impl Type { @@ -22,7 +26,7 @@ impl Type { pub fn carrier(self) -> u8 { use ScalarType::*; match self { - Self::Unit | Self::Opaque(_) => 0, + Self::Unit | Self::Opaque(_) | Self::Layout(_) => 0, Self::Scalar(Bool | I8 | U8 | I16 | U16 | I32 | U32 | Char | F16) => 1, Self::Scalar(I64 | U64) => 2, Self::Scalar(F32) => 3, @@ -32,16 +36,74 @@ impl Type { } } +/// Shared exact pointee layout, kept out of the eight-byte ordinary type entry. +#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)] +pub struct AddressLayout { + pub value: TypeId, + pub size: u32, + pub codec: Option, +} + #[derive(Default, Debug, Clone)] pub struct Types { base: Option>, values: Vec, ids: FxHashMap, + layouts: Vec, + layout_ids: FxHashMap, symbols: Vec>, symbol_ids: FxHashMap, SymbolId>, } impl Types { + pub fn layout(&mut self, layout: AddressLayout) -> TypeId { + let id = if let Some(&id) = self + .layout_ids + .get(&layout) + .or_else(|| self.base.as_ref().and_then(|b| b.layout_ids.get(&layout))) + { + id + } else { + let id = LayoutId::new(self.base_layouts() + self.layouts.len()); + self.layouts.push(layout); + self.layout_ids.insert(layout, id); + id + }; + self.intern(Type::Layout(id)) + } + fn base_layouts(&self) -> usize { + self.base.as_ref().map_or(0, |b| b.layouts.len()) + } + pub fn get_layout(&self, id: LayoutId) -> AddressLayout { + let base = self.base_layouts(); + if id.index() < base { + self.base.as_ref().unwrap().layouts[id.index()] + } else { + self.layouts[id.index() - base] + } + } + pub fn pointee(&self, pointer: TypeId) -> Option { + let Type::Pointer(inner) = self.get(pointer)? else { + return None; + }; + Some(match self.get(inner)? { + Type::Layout(id) => self.get_layout(id).value, + _ => inner, + }) + } + pub fn address_layout(&self, pointer: TypeId) -> Option<(u32, Option)> { + let Type::Pointer(inner) = self.get(pointer)? else { + return None; + }; + match self.get(inner)? { + Type::Layout(id) => { + let layout = self.get_layout(id); + Some((layout.size, layout.codec)) + } + _ => None, + } + } + /// Extend an immutable common vocabulary without copying its tables. Only /// one shared layer is allowed, so lookup cost cannot grow across bodies. pub fn with_base(base: Arc) -> Self { @@ -110,6 +172,12 @@ impl Types { self.symbol_ids.insert(name, id); id } + pub fn find_symbol(&self, name: &str) -> Option { + self.symbol_ids + .get(name) + .or_else(|| self.base.as_ref().and_then(|b| b.symbol_ids.get(name))) + .copied() + } pub fn symbol_name(&self, symbol: SymbolId) -> Option<&str> { let base = self.base_symbols(); if symbol.index() < base { @@ -127,7 +195,14 @@ impl Types { impl PartialEq for Types { fn eq(&self, other: &Self) -> bool { self.len() == other.len() - && (0..self.len()).all(|i| self.get(TypeId::new(i)) == other.get(TypeId::new(i))) + && (0..self.len()).all(|i| { + match (self.get(TypeId::new(i)), other.get(TypeId::new(i))) { + (Some(Type::Layout(a)), Some(Type::Layout(b))) => { + self.get_layout(a) == other.get_layout(b) + } + (a, b) => a == b, + } + }) && self.base_symbols() + self.symbols.len() == other.base_symbols() + other.symbols.len() && (0..self.base_symbols() + self.symbols.len()) @@ -139,7 +214,13 @@ impl std::hash::Hash for Types { fn hash(&self, state: &mut H) { self.len().hash(state); for i in 0..self.len() { - self.get(TypeId::new(i)).unwrap().hash(state); + let ty = self.get(TypeId::new(i)).unwrap(); + if let Type::Layout(id) = ty { + std::mem::discriminant(&ty).hash(state); + self.get_layout(id).hash(state); + } else { + ty.hash(state); + } } (self.base_symbols() + self.symbols.len()).hash(state); for i in 0..self.base_symbols() + self.symbols.len() { @@ -147,3 +228,28 @@ impl std::hash::Hash for Types { } } } + +#[cfg(test)] +mod layout_tests { + use super::*; + #[test] + fn layouts_are_interned_across_shared_vocabularies_without_growing_types() { + assert_eq!(std::mem::size_of::(), 8); + let mut base = Types::default(); + let value = base.scalar(ScalarType::I64); + let first = AddressLayout { + value, + size: 8, + codec: None, + }; + let ty = base.layout(first); + let mut child = Types::with_base(Arc::new(base)); + assert_eq!(child.layout(first), ty); + assert!(!child.has_additions()); + let different = child.layout(AddressLayout { size: 16, ..first }); + assert_ne!(different, ty); + let pointer = child.intern(Type::Pointer(different)); + assert_eq!(child.pointee(pointer), Some(value)); + assert_eq!(child.address_layout(pointer), Some((16, None))); + } +} diff --git a/compiler-core/src/ir/verify/mod.rs b/compiler-core/src/ir/verify/mod.rs index 4ed6ef87..cec60be2 100644 --- a/compiler-core/src/ir/verify/mod.rs +++ b/compiler-core/src/ir/verify/mod.rs @@ -36,14 +36,14 @@ pub fn verify_with_debug( check!( types .get(body.return_type) - .is_some_and(|t| !matches!(t, Type::Opaque(_))), + .is_some_and(|t| !matches!(t, Type::Opaque(_) | Type::Layout(_))), "invalid return type" ); for value in &body.values { check!( types .get(value.ty) - .is_some_and(|t| !matches!(t, Type::Opaque(_))), + .is_some_and(|t| !matches!(t, Type::Opaque(_) | Type::Layout(_))), "invalid value type" ); } @@ -67,16 +67,11 @@ pub fn verify_with_debug( Some(Type::Class(symbol) | Type::Interface(symbol)) => { check!(types.symbol_name(symbol).is_some(), "invalid storage class") } - Some(Type::Pointer(_) | Type::Slice(_) | Type::Str) => {} + Some(Type::Pointer(_) | Type::Slice(_) | Type::Str | Type::TaggedI64) => {} _ => return Err(VerifyError("unsupported storage type".into())), } } for field in &body.fields { - check!( - !field.relative_pointer - || (!field.is_static && matches!(types.get(field.ty), Some(Type::Pointer(_)))), - "invalid relative pointer field" - ); check!( matches!(types.get(field.owner), Some(Type::Class(symbol) | Type::Interface(symbol)) if types.symbol_name(symbol).is_some()), "invalid field owner" @@ -84,7 +79,7 @@ pub fn verify_with_debug( check!( types .get(field.ty) - .is_some_and(|t| !matches!(t, Type::Unit | Type::Opaque(_))), + .is_some_and(|t| !matches!(t, Type::Unit | Type::Opaque(_) | Type::Layout(_))), "invalid field type" ); } @@ -191,11 +186,36 @@ pub fn verify_with_debug( check!(method.index() < body.methods.len(), "invalid call target"); check!(args.range().end <= body.args.len(), "invalid operand list"); } - Op::Overflow { args, .. } => { + Op::Overflow { args, .. } + | Op::Heap { args, .. } + | Op::LoadStorageField { address: args, .. } + | Op::LoadStorageFieldCopy { address: args, .. } + | Op::StoreStorageField { args, .. } + | Op::AddressPack(args) + | Op::LoadAddress(args) + | Op::LoadAddressCopy(args) + | Op::CopyStorage { parts: args, .. } + | Op::LoadTypedCopy { parts: args, .. } + | Op::LoadTyped { parts: args, .. } + | Op::StoreTyped { parts: args, .. } + | Op::TypedAddressPack { parts: args, .. } + | Op::LocationTag(args) + | Op::LocationEqual(args) + | Op::LocationCompare(args) + | Op::TaggedPack(args) + | Op::ViewPack(args) + | Op::ViewGet(args) + | Op::ViewSet { parts: args, .. } + | Op::ViewAddress { parts: args, .. } + | Op::StoreAddress { parts: args, .. } + | Op::StoreFieldParts { parts: args, .. } => { check!(args.range().end <= body.args.len(), "invalid operand list") } Op::Constant(id) => check!(id.index() < body.constants.len(), "invalid constant"), - Op::LoadSlot(id) | Op::AddressOfSlot(id) | Op::StoreSlot { slot: id, .. } => { + Op::LoadSlot(id) + | Op::AddressOfSlot(id) + | Op::StoreSlot { slot: id, .. } + | Op::SlotRoot(id) => { check!(id.index() < body.slots.len(), "invalid storage slot") } _ => {} @@ -259,7 +279,9 @@ pub fn verify_with_debug( for local in &debug.locals { match *local { DebugLocal::Value(ty) => check!( - types.get(ty).is_some_and(|t| !matches!(t, Type::Opaque(_))), + types + .get(ty) + .is_some_and(|t| !matches!(t, Type::Opaque(_) | Type::Layout(_))), "invalid debug local type" ), DebugLocal::Storage(slot) => { diff --git a/compiler-core/src/ir/verify/types.rs b/compiler-core/src/ir/verify/types.rs index b7a6776f..8eae380d 100644 --- a/compiler-core/src/ir/verify/types.rs +++ b/compiler-core/src/ir/verify/types.rs @@ -5,6 +5,176 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() let result = inst.result.map(|v| body.value_type(v)); let ty = |v| body.value_type(v); match inst.op { + Op::CopyStorage { parts, layouts, .. } => { + let parts = body + .args + .get(parts.range()) + .ok_or_else(|| VerifyError("invalid copy components".into()))?; + check!( + parts.len() == 5 && result.is_none(), + "invalid memory copy operands" + ); + for (i, layout) in layouts.into_iter().enumerate() { + let Some(Type::Layout(layout)) = types.get(layout) else { + return Err(VerifyError("copy needs an exact storage layout".into())); + }; + let AddressLayout { size, codec, .. } = types.get_layout(layout); + check!( + types + .get(ty(parts[i * 2])) + .is_some_and(|t| t.carrier() == 5) + && types.get(ty(parts[i * 2 + 1])) == Some(Type::Scalar(ScalarType::I64)), + "invalid copy location" + ); + check!( + size <= i32::MAX as u32 && codec.is_none_or(|s| types.symbol_name(s).is_some()), + "invalid copy layout" + ); + } + check!( + matches!( + types.get(ty(parts[4])), + Some(Type::Scalar( + ScalarType::I32 | ScalarType::U32 | ScalarType::I64 | ScalarType::U64 + )) + ), + "copy length needs an integer" + ); + } + Op::ViewRoot { + backing, + size, + codec, + } => check!( + types.get(ty(backing)).is_some_and(|t| t.carrier() == 5) + && result.is_some_and(|t| matches!(types.get(t), Some(Type::Class(s)) + if types.symbol_name(s) == Some("java/lang/Object"))) + && (1..=i32::MAX as u32).contains(&size) + && codec.is_none_or(|s| types.symbol_name(s).is_some()), + "invalid slice storage root" + ), + Op::Heap { operation, args } => { + let args = &body.args[args.range()]; + let (count, root, returns) = match operation { + HeapOp::Allocate => (2, false, true), + HeapOp::Reallocate => (5, true, true), + HeapOp::Deallocate => (2, true, false), + }; + check!(args.len() == count, "invalid heap operation arity"); + check!( + !root || types.get(ty(args[0])).is_some_and(|t| t.carrier() == 5), + "heap operation requires a storage root" + ); + check!( + args[usize::from(root)..] + .iter() + .all(|&arg| types.get(ty(arg)) == Some(Type::Scalar(ScalarType::I64))), + "heap sizes, alignments and offsets must be i64" + ); + check!( + if returns { + result.is_some_and(|ty| { + matches!(types.get(ty), Some(Type::Class(name)) + if types.symbol_name(name) == Some("java/lang/Object")) + }) + } else { + result.is_none() + }, + "invalid heap operation result" + ); + } + Op::CopyValue(value) => check!( + result == Some(ty(value)) && types.get(ty(value)).is_some_and(|t| t.carrier() == 5), + "copy requires the same reference value type" + ), + Op::TaggedPack(parts) => { + let parts = &body.args[parts.range()]; + check!( + parts.len() == 2 + && parts + .iter() + .all(|&v| types.get(ty(v)) == Some(Type::Scalar(ScalarType::I64))) + && result.is_some_and(|t| types.get(t) == Some(Type::TaggedI64)), + "invalid tagged scalar components" + ); + } + Op::TaggedPart { value, index } => { + check!( + index < 2 + && types.get(ty(value)) == Some(Type::TaggedI64) + && result.is_some_and(|t| types.get(t) == Some(Type::Scalar(ScalarType::I64))), + "invalid tagged scalar projection" + ); + } + Op::LoadStorageField { + address, + projection, + .. + } + | Op::LoadStorageFieldCopy { + address, + projection, + } + | Op::StoreStorageField { + args: address, + projection, + .. + } => { + let parts = body + .args + .get(address.range()) + .ok_or_else(|| VerifyError("invalid storage address".into()))?; + check!( + parts.len() >= 2 + && types.get(ty(parts[0])).is_some_and(|t| t.carrier() == 5) + && types.get(ty(parts[1])) == Some(Type::Scalar(ScalarType::I64)), + "invalid storage address components" + ); + let projection = body + .projections + .get(projection.index()) + .ok_or_else(|| VerifyError("invalid storage projection".into()))?; + let field = body + .fields + .get(projection.field.index()) + .ok_or_else(|| VerifyError("invalid storage field".into()))?; + check!( + !field.is_static && matches!(types.get(field.owner), Some(Type::Class(_))), + "storage projection needs a concrete owner" + ); + match inst.op { + Op::LoadStorageFieldCopy { .. } => check!( + parts.len() == 2 + && result == Some(field.ty) + && types.get(field.ty).is_some_and(|t| t.carrier() == 5), + "invalid owned storage field result" + ), + Op::LoadStorageField { index, .. } => check!( + parts.len() == 2 + && result.is_some_and(|result| match index { + Some(index) => + borrowed_part_matches(types, field.ty, index as usize, result), + None => result == field.ty, + }), + "invalid storage field result" + ), + Op::StoreStorageField { split, .. } => check!( + result.is_none() + && if split { + ComponentShape::of(types, field.ty) + .is_some_and(|s| s.len() == parts.len() - 2) + && parts[2..] + .iter() + .enumerate() + .all(|(i, &p)| borrowed_part_matches(types, field.ty, i, ty(p))) + } else { + parts.len() == 3 && ty(parts[2]) == field.ty + }, + "invalid storage field values" + ), + _ => unreachable!(), + } + } Op::Nop => check!(result.is_none(), "nop produces a value"), Op::ArrayLength(value) => check!( result.and_then(|t| types.get(t)) == Some(Type::Scalar(ScalarType::I32)) @@ -14,6 +184,13 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() ), "array length requires an array or view and int result" ), + Op::Refine(value) => { + check!( + result.and_then(|r| types.get(r)).map(Type::carrier) == Some(5) + && types.get(ty(value)).map(Type::carrier) == Some(5), + "reference refinement requires reference operands" + ); + } Op::Reinterpret(value) => { check!( result.and_then(|r| types.get(r)).map(Type::carrier) @@ -31,6 +208,38 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() "invalid ABI adaptation" ); } + Op::ScalarCell(value) => { + check!( + StorageSlot::scalar(ty(value), types).is_some() + && result.and_then(|r| types.get(r)) == Some(Type::Pointer(ty(value))), + "scalar storage requires an exact scalar pointee" + ); + } + Op::ProjectRoot { + address, + projection, + } => { + let values = &body.args[address.range()]; + check!( + values.len() == 2 + && types.get(ty(values[0])).is_some_and(|t| t.carrier() == 5) + && types.get(ty(values[1])) == Some(Type::Scalar(ScalarType::I64)) + && projection.index() < body.projections.len() + && result + .and_then(|r| types.get(r)) + .is_some_and(|t| t.carrier() == 5), + "invalid typed field location" + ); + } + Op::ProjectOffset { root, base, offset } => { + check!( + types.get(ty(root)).is_some_and(|t| t.carrier() == 5) + && types.get(ty(base)).is_some_and(|t| t.carrier() == 5) + && types.get(ty(offset)) == Some(Type::Scalar(ScalarType::I64)) + && result == Some(ty(offset)), + "invalid typed field displacement" + ); + } Op::NewArray(size) => { check!( types.get(ty(size)) == Some(Type::Scalar(ScalarType::I32)), @@ -41,6 +250,14 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() "array allocation requires array result" ); } + Op::ArrayFill { array, value } => { + check!( + result.is_none() + && types.get(ty(array)) == Some(Type::Array(ty(value))) + && StorageSlot::scalar(ty(value), types).is_some(), + "primitive array fill type mismatch" + ); + } Op::FunctionPointer { signature, target } => { let signature = body .methods @@ -251,8 +468,61 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() } } Op::Opaque(value) => check!(result == Some(ty(value)), "opaque value type mismatch"), + Op::AddressTag(value) => { + check!( + matches!(types.get(ty(value)), Some(Type::Pointer(_))) + && result.and_then(|ty| types.get(ty)) == Some(Type::Scalar(ScalarType::I64)), + "nullable tag needs pointer operand and i64 result" + ); + } + Op::AddressEqual { left, right } | Op::AddressCompare { left, right } => { + let output = if matches!(inst.op, Op::AddressCompare { .. }) { + ScalarType::I32 + } else { + ScalarType::Bool + }; + check!( + matches!(types.get(ty(left)), Some(Type::Pointer(_))) + && matches!(types.get(ty(right)), Some(Type::Pointer(_))) + && result.and_then(|ty| types.get(ty)) == Some(Type::Scalar(output)), + "invalid address comparison" + ); + } + Op::LocationEqual(parts) | Op::LocationCompare(parts) => { + let output = if matches!(inst.op, Op::LocationCompare(_)) { + ScalarType::I32 + } else { + ScalarType::Bool + }; + let values = body + .args + .get(parts.range()) + .ok_or_else(|| VerifyError("invalid location operands".into()))?; + check!( + values.len() == 4 + && result.and_then(|ty| types.get(ty)) == Some(Type::Scalar(output)), + "invalid location comparison" + ); + for (index, &value) in values.iter().enumerate() { + check!( + if index % 2 == 0 { + matches!(types.get(ty(value)), Some(Type::Class(s)) if types.symbol_name(s) == Some("java/lang/Object")) + } else { + types.get(ty(value)) == Some(Type::Scalar(ScalarType::I64)) + }, + "location component mismatch" + ); + } + } Op::Project { base, projection } | Op::LoadField { base, projection } + | Op::LoadFieldCopy { base, projection } + | Op::LoadFieldPart { + base, projection, .. + } + | Op::StoreFieldParts { + base, projection, .. + } | Op::StoreField { base, projection, .. } => { @@ -265,18 +535,26 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() .get(projection.field.index()) .ok_or_else(|| VerifyError("invalid projection field".into()))?; check!( - !field.is_static && types.get(ty(base)) == Some(Type::Pointer(field.owner)), + !field.is_static && types.pointee(ty(base)) == Some(field.owner), "projection owner mismatch" ); - if !matches!(inst.op, Op::Project { .. }) { + if !matches!(inst.op, Op::Project { .. } | Op::LoadFieldCopy { .. }) { check!( - matches!(types.get(field.ty), Some(Type::Scalar(_))), - "promoted field access requires a scalar field" + matches!( + types.get(field.ty), + Some(Type::Scalar(_) | Type::Pointer(_) | Type::Slice(_) | Type::Str) + ), + "promoted field access requires a scalar or pointer field" ); } match inst.op { + Op::LoadFieldCopy { .. } => check!( + result == Some(field.ty) + && types.get(field.ty).is_some_and(|t| t.carrier() == 5), + "owned field load type mismatch" + ), Op::Project { .. } => check!( - result.and_then(|id| types.get(id)) == Some(Type::Pointer(field.ty)), + result.and_then(|id| types.pointee(id)) == Some(field.ty), "projection result mismatch" ), Op::LoadField { .. } => { @@ -286,6 +564,41 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() result.is_none() && ty(value) == field.ty, "field store type mismatch" ), + Op::LoadFieldPart { index, .. } => { + check!( + result.is_some_and(|part| borrowed_part_matches( + types, + field.ty, + index as usize, + part + )), + "invalid split field load" + ); + } + Op::StoreFieldParts { parts, .. } => { + let values = body + .args + .get(parts.range()) + .ok_or_else(|| VerifyError("invalid split field operands".into()))?; + check!( + result.is_none() + && ComponentShape::of(types, field.ty) + .is_some_and(|shape| shape.len() == values.len()), + "invalid split field store" + ); + check!( + values + .iter() + .enumerate() + .all(|(index, &value)| borrowed_part_matches( + types, + field.ty, + index, + ty(value) + )), + "invalid split field components" + ); + } _ => unreachable!(), } check!( @@ -327,6 +640,193 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() "view length needs usize" ); } + Op::RetypeAddress { + pointer, + size, + codec, + } => { + check!( + matches!(types.get(ty(pointer)), Some(Type::Pointer(_))) + && matches!(result.and_then(|t| types.get(t)), Some(Type::Pointer(_))), + "retyped address needs pointer operands" + ); + check!( + size <= i32::MAX as u32 && codec.is_none_or(|id| types.symbol_name(id).is_some()), + "invalid address layout" + ); + } + Op::LocationTag(parts) + | Op::AddressPack(parts) + | Op::LoadAddress(parts) + | Op::LoadAddressCopy(parts) + | Op::LoadTypedCopy { parts, .. } + | Op::LoadTyped { parts, .. } + | Op::StoreTyped { parts, .. } + | Op::TypedAddressPack { parts, .. } + | Op::StoreAddress { parts, .. } => { + let parts = body + .args + .get(parts.range()) + .ok_or_else(|| VerifyError("invalid address components".into()))?; + check!( + match inst.op { + Op::LoadTyped { .. } => matches!(parts.len(), 2 | 3), + Op::StoreTyped { .. } => parts.len() == 3, + _ => parts.len() == 2, + }, + "address component count" + ); + check!( + types.get(ty(parts[0])).is_some_and(|t| t.carrier() == 5), + "address root needs a reference" + ); + check!( + types.get(ty(parts[1])) == Some(Type::Scalar(ScalarType::I64)), + "address offset needs signed long" + ); + if let Op::LoadTypedCopy { size, codec, .. } + | Op::LoadTyped { size, codec, .. } + | Op::TypedAddressPack { size, codec, .. } + | Op::StoreTyped { size, codec, .. } = inst.op + { + check!( + size > 0 + && size <= i32::MAX as u32 + && codec.is_none_or(|id| types.symbol_name(id).is_some()), + "invalid typed address layout" + ); + } + match inst.op { + Op::LocationTag(_) => check!( + result.and_then(|ty| types.get(ty)) == Some(Type::Scalar(ScalarType::I64)), + "nullable tag needs i64 result" + ), + Op::AddressPack(_) | Op::TypedAddressPack { .. } => check!( + result + .and_then(|t| ComponentShape::of(types, t)) + .is_some_and(ComponentShape::is_address), + "address pack needs pointer type" + ), + Op::LoadAddressCopy(_) | Op::LoadTypedCopy { .. } => check!( + result.is_some_and(|t| types.get(t).is_some_and(|t| t.carrier() == 5)), + "owned address load needs a reference carrier" + ), + Op::LoadTyped { .. } => { + check!( + result.is_some_and(|t| types.get(t).is_some_and(|t| t.carrier() == 5)), + "typed borrowed load needs a reference carrier" + ); + if let Some(&target) = parts.get(2) { + check!( + matches!(types.get(ty(target)), Some(Type::Class(name)) + if types.symbol_name(name) == Some("java/lang/String")), + "typed borrowed load needs a class name" + ); + } + } + Op::LoadAddress(_) => check!( + result.is_some_and(|t| StorageSlot::scalar(t, types).is_some() + || types.get(t).is_some_and(|t| t.carrier() == 5)), + "address load needs a stored value" + ), + Op::StoreTyped { .. } => check!( + result.is_none() && types.get(ty(parts[2])).is_some_and(|t| t.carrier() == 5), + "typed store needs an aggregate value" + ), + Op::StoreAddress { value, .. } => check!( + result.is_none() + && (StorageSlot::scalar(ty(value), types).is_some() + || types.get(ty(value)).is_some_and(|t| t.carrier() == 5)), + "address store needs a stored value" + ), + _ => unreachable!(), + } + } + Op::AddressPart { address, index } => { + check!( + ComponentShape::of(types, ty(address)).is_some_and(ComponentShape::is_address), + "address part needs pointer type" + ); + check!( + match index { + 0 => result + .and_then(|t| types.get(t)) + .is_some_and(|t| t.carrier() == 5), + 1 => result.and_then(|t| types.get(t)) == Some(Type::Scalar(ScalarType::I64)), + _ => false, + }, + "invalid address component" + ); + } + Op::SlotRoot(slot) => { + check!( + body.slots + .get(slot.index()) + .and_then(|s| StorageSlot::scalar(s.ty, types)) + .is_some(), + "slot root needs primitive storage" + ); + check!( + result + .and_then(|t| types.get(t)) + .is_some_and(|t| t.carrier() == 5), + "slot root needs reference result" + ); + } + Op::ViewPack(parts) | Op::ViewAddress { parts, .. } => { + let parts = body + .args + .get(parts.range()) + .ok_or_else(|| VerifyError("invalid view components".into()))?; + check!(parts.len() == 3, "view needs backing, start, and length"); + check!( + types.get(ty(parts[0])).is_some_and(|ty| ty.carrier() == 5), + "view backing needs a reference" + ); + check!( + types.get(ty(parts[1])) == Some(Type::Scalar(ScalarType::I32)), + "view start needs an int" + ); + check!( + types.get(ty(parts[2])) == Some(Type::Scalar(ScalarType::U64)), + "view length needs usize" + ); + if let Op::ViewPack(_) = inst.op { + check!( + result.is_some_and(|t| ComponentShape::view_carrier(types, t)), + "view pack result" + ); + return Ok(()); + } + let Op::ViewAddress { size, codec, .. } = inst.op else { + unreachable!() + }; + check!( + matches!(result.and_then(|ty| types.get(ty)), Some(Type::Pointer(_))), + "view address needs a pointer" + ); + check!( + size <= i32::MAX as u32 && codec.is_none_or(|id| types.symbol_name(id).is_some()), + "invalid view address layout" + ); + } + Op::ViewPart { view, index } => { + check!( + ComponentShape::view_carrier(types, ty(view)), + "view part source" + ); + check!( + match index { + 0 => result + .and_then(|t| types.get(t)) + .is_some_and(|t| t.carrier() == 5), + 1 => result.and_then(|t| types.get(t)) == Some(Type::Scalar(ScalarType::I32)), + 2 => result.and_then(|t| types.get(t)) == Some(Type::Scalar(ScalarType::U64)), + _ => false, + }, + "view part result" + ); + } Op::ViewData { view, size, codec } => { let element = view_element(types.get(ty(view)), types); check!( @@ -339,6 +839,44 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() "invalid view element layout" ); } + Op::AddressViewPart { address, index } => { + check!( + ComponentShape::of(types, ty(address)).is_some_and(ComponentShape::is_address), + "view part needs an address" + ); + check!( + match (index, result.and_then(|t| types.get(t))) { + (0, Some(Type::Class(s))) => types.symbol_name(s) == Some("java/lang/Object"), + (1, Some(Type::Scalar(ScalarType::I32))) => true, + _ => false, + }, + "invalid address view component" + ); + } + Op::TypedAddressViewPart { + parts, + size, + codec, + index, + } => { + let parts = &body.args[parts.range()]; + check!(parts.len() == 2, "invalid typed slice location arity"); + check!( + types.get(ty(parts[0])).is_some_and(|t| t.carrier() == 5) + && types.get(ty(parts[1])) == Some(Type::Scalar(ScalarType::I64)) + && (1..=i32::MAX as u32).contains(&size) + && codec.is_none_or(|s| types.symbol_name(s).is_some()), + "invalid typed slice location" + ); + check!( + match (index, result.and_then(|t| types.get(t))) { + (0, Some(Type::Class(s))) => types.symbol_name(s) == Some("java/lang/Object"), + (1, Some(Type::Scalar(ScalarType::I32))) => true, + _ => false, + }, + "invalid typed slice component" + ); + } Op::Cast(value) => { check!(result.is_some(), "cast has no result"); if matches!(types.get(ty(value)), Some(Type::Scalar(_))) { @@ -356,12 +894,25 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() result.and_then(|ty| types.get(ty)) == Some(Type::Pointer(body.slots[slot.index()].ty)), "slot address type mismatch" ), + Op::Commit(pointer) => check!( + result.is_none() && matches!(types.get(ty(pointer)), Some(Type::Pointer(_))), + "view commit needs its address owner" + ), + Op::LoadCopy(pointer) => check!( + result.is_some_and(|t| types.get(t).is_some_and(|t| t.carrier() == 5)) + && types.pointee(ty(pointer)) == result, + "owned pointer load type mismatch" + ), Op::Load(pointer) => check!( - result.is_some() && types.get(ty(pointer)) == result.map(Type::Pointer), - "pointer load type mismatch" + result.is_some() && types.pointee(ty(pointer)) == result, + "pointer load type mismatch: address {:?} ({:?}), result {:?} ({:?})", + pointer, + types.get(ty(pointer)), + result, + result.and_then(|ty| types.get(ty)) ), Op::Store { pointer, value } => check!( - result.is_none() && types.get(ty(pointer)) == Some(Type::Pointer(ty(value))), + result.is_none() && types.pointee(ty(pointer)) == Some(ty(value)), "pointer store type mismatch" ), Op::StoreSlot { slot, value } => { @@ -395,8 +946,45 @@ pub(super) fn verify_types(inst: &Inst, body: &Body, types: &Types) -> Result<() check!(result == Some(member.ty), "field load type mismatch"); } } - Op::ArraySet { .. } => check!(result.is_none(), "store produces value"), - _ => {} + Op::ViewGet(parts) | Op::ViewSet { parts, .. } => { + let parts = &body.args[parts.range()]; + check!( + parts.len() == 3 + && types.get(ty(parts[0])).is_some_and(|t| t.carrier() == 5) + && parts[1..] + .iter() + .all(|&p| types.get(ty(p)) == Some(Type::Scalar(ScalarType::I32))), + "slice access requires backing, start and index" + ); + check!( + matches!(inst.op, Op::ViewGet(_)) == result.is_some(), + "slice access result mismatch" + ); + } + Op::ArrayGet { + array, + index, + native, + } + | Op::ArraySet { + array, + index, + native, + .. + } => { + check!( + !native || matches!(types.get(ty(array)), Some(Type::Array(_))), + "native scratch access requires a JVM array" + ); + check!( + types.get(ty(index)) == Some(Type::Scalar(ScalarType::I32)), + "array index requires JVM int" + ); + check!( + matches!(inst.op, Op::ArrayGet { .. }) == result.is_some(), + "array access result mismatch" + ); + } } Ok(()) } @@ -408,3 +996,18 @@ fn view_element(ty: Option, types: &Types) -> Option { _ => None, } } + +fn borrowed_part_matches(types: &Types, logical: TypeId, index: usize, part: TypeId) -> bool { + match (ComponentShape::of(types, logical), index, types.get(part)) { + (Some(ComponentShape::TaggedI64), 0..=1, Some(Type::Scalar(ScalarType::I64))) => true, + (Some(_), 0, Some(Type::Class(s))) => types.symbol_name(s) == Some("java/lang/Object"), + ( + Some(ComponentShape::Address | ComponentShape::StorageAddress), + 1, + Some(Type::Scalar(ScalarType::I64)), + ) + | (Some(ComponentShape::View), 1, Some(Type::Scalar(ScalarType::I32))) + | (Some(ComponentShape::View), 2, Some(Type::Scalar(ScalarType::U64))) => true, + _ => false, + } +} diff --git a/compiler-core/src/jvm/abi.rs b/compiler-core/src/jvm/abi.rs index 0acb87e3..883ba21d 100644 --- a/compiler-core/src/jvm/abi.rs +++ b/compiler-core/src/jvm/abi.rs @@ -1,15 +1,117 @@ //! JVM names shared by body selection and generated representation schemas. pub const SLICE_VIEW_CLASS: &str = "org/rustlang/runtime/SliceView"; pub const UTF8_VIEW_CLASS: &str = "org/rustlang/runtime/Utf8View"; +pub const TAGGED_LONG_CLASS: &str = "org/rustlang/runtime/TaggedLong"; pub const POINTER_CLASS: &str = "org/rustlang/runtime/Pointer"; -pub const RELATIVE_POINTER_METHOD_SUFFIX: &str = "$relative"; -pub const RELATIVE_POINTER_ELEMENT_OFFSET_SUFFIX: &str = "$rcj$elementOffset"; -pub const RELATIVE_POINTER_BYTE_OFFSET_SUFFIX: &str = "$rcj$byteOffset"; -pub fn relative_pointer_element_offset_field(field: &str) -> String { - format!("{field}{RELATIVE_POINTER_ELEMENT_OFFSET_SUFFIX}") +/// Reconstruction plans for addresses of stored borrows. These are ABI tags. +/// Ordinary scalar plans retain their 1/2/4/8-byte strides. +pub const STORED_VIEW: u32 = 64; +pub const STORED_ADDRESS: u32 = 128; + +pub fn address_plan(types: &crate::ir::Types, pointee: crate::ir::TypeId) -> u32 { + use crate::ir::{StorageSlot, Type}; + match types.get(pointee) { + Some(Type::Slice(_) | Type::Str) => STORED_VIEW, + Some(Type::Pointer(_)) => STORED_ADDRESS, + _ => StorageSlot::scalar(pointee, types).map_or(0, |slot| slot.size), + } +} + +/// Consume an Object root and long displacement. +/// Construct a boundary carrier only when a consumer needs one reference. +pub fn materialize_address( + cp: &mut crate::classfile::constant_pool::InternedConstantPool, + code: &mut Vec, + plan: u32, +) -> crate::classfile::Result<()> { + use crate::classfile::attributes::Instruction; + code.push(crate::jvm::constants::get_int_const_instr(cp, plan as i32)); + let owner = cp.add_class(POINTER_CLASS)?; + code.push(Instruction::Invokestatic(cp.add_method_ref( + owner, + "addressFromParts", + "(Ljava/lang/Object;JI)Lorg/rustlang/runtime/Pointer;", + )?)); + Ok(()) } -pub fn relative_pointer_byte_offset_field(field: &str) -> String { - format!("{field}{RELATIVE_POINTER_BYTE_OFFSET_SUFFIX}") +pub fn materialize_typed_address( + cp: &mut crate::classfile::constant_pool::InternedConstantPool, + code: &mut Vec, + size: u32, + codec: Option<&str>, +) -> crate::classfile::Result<()> { + use crate::classfile::attributes::Instruction; + code.push(crate::jvm::constants::get_int_const_instr(cp, size as i32)); + code.push(if let Some(codec) = codec { + Instruction::Ldc_w(cp.add_name_string(codec)?) + } else { + Instruction::Aconst_null + }); + let owner = cp.add_class(POINTER_CLASS)?; + code.push(Instruction::Invokestatic(cp.add_method_ref( + owner, + "fromTypedStorageLocation", + "(Ljava/lang/Object;JILjava/lang/String;)Lorg/rustlang/runtime/Pointer;", + )?)); + Ok(()) +} + +pub const VIEW_PARTS: [(&str, &str); 3] = [ + ("array", "Ljava/lang/Object;"), + ("offset", "I"), + ("rustLength", "J"), +]; + +/// Use null for an absent optional borrow at boundaries. +/// All adapters use the same default component slots for this niche. +pub fn view_part_access( + cp: &mut crate::classfile::constant_pool::InternedConstantPool, + index: usize, +) -> crate::classfile::Result { + let (name, descriptor) = VIEW_PARTS[index]; + let owner = cp.add_class(SLICE_VIEW_CLASS)?; + Ok(crate::classfile::attributes::Instruction::Invokestatic( + cp.add_method_ref( + owner, + format!("$part${name}"), + format!("(L{SLICE_VIEW_CLASS};){descriptor}"), + )?, + )) +} +/// Synthetic displacement field paired with a scalar-address root field. +pub fn tagged_field_names(name: &str) -> [String; 2] { + [name.into(), format!("$rust$t${name}")] +} +pub fn address_field_name(name: &str, size: u32) -> String { + format!("$rust${size}${name}") +} +/// Exact layouts are metadata on the carrier schema, not on every borrow. +pub fn typed_address_field_name(name: &str, size: u32) -> String { + format!("$rust$a{size}${name}") +} +pub fn address_codec_field_name(name: &str) -> String { + format!("$rust$c${name}") +} +pub fn address_displacement_name( + types: &crate::ir::Types, + pointer: crate::ir::TypeId, + name: &str, +) -> String { + if let Some((size, _)) = types.address_layout(pointer) { + typed_address_field_name(name, size) + } else { + let Some(crate::ir::Type::Pointer(inner)) = types.get(pointer) else { + unreachable!() + }; + address_field_name(name, address_plan(types, inner)) + } +} +pub fn view_field_names(name: &str, utf8: bool) -> [String; 3] { + [ + name.into(), + format!("$rust${}${name}", if utf8 { "u" } else { "s" }), + format!("$rust$l${name}"), + ] } diff --git a/compiler-core/src/jvm/frames/tests.rs b/compiler-core/src/jvm/frames/tests.rs index c08136c7..b4fddc10 100644 --- a/compiler-core/src/jvm/frames/tests.rs +++ b/compiler-core/src/jvm/frames/tests.rs @@ -89,3 +89,47 @@ fn definite_assignment_accepts_a_store_on_every_path() { let loads = locals_loaded_before_definite_store(&instructions, &[], 2, "test", &[]).unwrap(); assert!(loads.is_empty()); } + +#[test] +fn duplicate_two_words_under_one_preserves_categories() { + use FrameValue::*; + for (input, expected) in [ + (vec![Null, Long], vec![Long, Null, Long]), + ( + vec![Null, Integer, Float], + vec![Integer, Float, Null, Integer, Float], + ), + ] { + let mut state = FrameState::new(Vec::new(), 0); + for value in input { + state.push(value); + } + super::transfer::transfer_instruction( + 0, + &Instruction::Dup2_x1, + &mut state, + &[], + &ConstantPool::default(), + "dup2_x1", + &mut SignatureCache::default(), + ) + .unwrap(); + assert_eq!(state.stack, expected); + assert_eq!(state.stack_words, 5); + } + let mut state = FrameState::new(Vec::new(), 0); + state.push(Long); + state.push(Long); + assert!( + super::transfer::transfer_instruction( + 0, + &Instruction::Dup2_x1, + &mut state, + &[], + &ConstantPool::default(), + "dup2_x1", + &mut SignatureCache::default() + ) + .is_err() + ); +} diff --git a/compiler-core/src/jvm/frames/transfer.rs b/compiler-core/src/jvm/frames/transfer.rs index 1592aefe..33cd114c 100644 --- a/compiler-core/src/jvm/frames/transfer.rs +++ b/compiler-core/src/jvm/frames/transfer.rs @@ -233,7 +233,24 @@ pub(super) fn transfer_instruction( state.push(value1); } } - I::Dup_x2 | I::Dup2_x1 | I::Dup2_x2 => { + I::Dup2_x1 => { + let first = state.pop(context, instruction_index)?; + if first.is_category2() { + let below = state.pop_category1(context, instruction_index)?; + state.push(first.clone()); + state.push(below); + state.push(first); + } else { + let second = state.pop_category1(context, instruction_index)?; + let below = state.pop_category1(context, instruction_index)?; + state.push(second.clone()); + state.push(first.clone()); + state.push(below); + state.push(second); + state.push(first); + } + } + I::Dup_x2 | I::Dup2_x2 => { return Err(jvm::Error::VerificationError { context: context.to_string(), message: format!( diff --git a/compiler-core/src/jvm/select/addresses.rs b/compiler-core/src/jvm/select/addresses.rs new file mode 100644 index 00000000..91dd67ec --- /dev/null +++ b/compiler-core/src/jvm/select/addresses.rs @@ -0,0 +1,449 @@ +//! Physical address components and scalar memory operations. +use super::*; + +impl Selector<'_> { + pub(super) fn address(&mut self, inst: Inst) -> jvm::Result { + match inst.op { + Op::CopyStorage { + parts, + layouts, + nonoverlapping, + } => { + let parts = &self.body.args[parts.range()]; + for (i, layout) in layouts.into_iter().enumerate() { + let Some(Type::Layout(layout)) = self.types.get(layout) else { + return Err(error("copy needs an exact storage layout")); + }; + let AddressLayout { size, codec, .. } = self.types.get_layout(layout); + self.load(parts[i * 2])?; + self.load(parts[i * 2 + 1])?; + self.address_layout(size, codec)?; + } + self.load(parts[4])?; + if self + .types + .get(self.body.value_type(parts[4])) + .unwrap() + .carrier() + == 0 + { + self.assembly.code.push(Instruction::I2l); + } + self.assembly.code.push(if nonoverlapping { + Instruction::Iconst_1 + } else { + Instruction::Iconst_0 + }); + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref(owner, "copyStorage", + "(Ljava/lang/Object;JILjava/lang/String;Ljava/lang/Object;JILjava/lang/String;JZ)V")?; + self.assembly.code.push(Instruction::Invokestatic(method)); + } + Op::TypedAddressViewPart { + parts, + size, + codec, + index, + } => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + self.address_layout(size, codec)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let (name, result) = if index == 0 { + ("typedLocationSliceBacking", "Ljava/lang/Object;") + } else { + ("typedLocationSliceOffset", "I") + }; + let method = self.cp.add_method_ref( + owner, + name, + format!("(Ljava/lang/Object;JILjava/lang/String;){result}"), + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + } + Op::RetypeAddress { + pointer, + size, + codec, + } => { + self.load(pointer)?; + self.address_layout(size, codec)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + "retype", + "(ILjava/lang/String;)Lorg/rustlang/runtime/Pointer;", + )?; + self.assembly.code.push(Instruction::Invokevirtual(method)); + } + Op::TypedAddressPack { parts, size, codec } => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + self.address_layout(size, codec)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + "fromTypedStorageLocation", + "(Ljava/lang/Object;JILjava/lang/String;)Lorg/rustlang/runtime/Pointer;", + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + } + Op::LoadTypedCopy { parts, size, codec } => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + self.address_layout(size, codec)?; + let ty = self.body.value_type(inst.result.unwrap()); + let name = self.address_target(ty)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + "loadTypedStorageCopy", + "(Ljava/lang/Object;JILjava/lang/String;Ljava/lang/String;)Ljava/lang/Object;", + )?; + let class = self.cp.add_class(&name)?; + self.assembly.code.extend([ + Instruction::Invokestatic(method), + Instruction::Checkcast(class), + ]); + } + Op::LoadTyped { parts, size, codec } => { + let parts = &self.body.args[parts.range()]; + self.load(parts[0])?; + self.load(parts[1])?; + self.address_layout(size, codec)?; + if let Some(&target) = parts.get(2) { + self.load(target)?; + } else if let Some(Type::Class(name) | Type::Interface(name)) = + self.types.get(self.body.value_type(inst.result.unwrap())) + && self.types.symbol_name(name) != Some("java/lang/Object") + { + self.assembly.code.push(Instruction::Ldc_w( + self.cp + .add_name_string(self.types.symbol_name(name).unwrap())?, + )); + } else { + self.assembly.code.push(Instruction::Aconst_null); + } + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + "loadTypedStorage", + "(Ljava/lang/Object;JILjava/lang/String;Ljava/lang/String;)Ljava/lang/Object;", + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + let mut descriptor = String::new(); + representation::descriptor( + self.types, + self.body.value_type(inst.result.unwrap()), + &mut descriptor, + )?; + if descriptor != "Ljava/lang/Object;" { + let target = descriptor + .strip_prefix('L') + .and_then(|s| s.strip_suffix(';')) + .unwrap_or(&descriptor); + self.assembly + .code + .push(Instruction::Checkcast(self.cp.add_class(target)?)); + } + } + Op::StoreTyped { parts, size, codec } => { + let parts = &self.body.args[parts.range()]; + self.load(parts[0])?; + self.load(parts[1])?; + self.address_layout(size, codec)?; + self.load(parts[2])?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + "storeTypedStorage", + "(Ljava/lang/Object;JILjava/lang/String;Ljava/lang/Object;)V", + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + } + Op::ProjectRoot { + address, + projection, + } => { + for &part in &self.body.args[address.range()] { + self.load(part)?; + } + self.project_field_arguments(projection)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let field = + &self.body.fields[self.body.projections[projection.index()].field.index()]; + let name = if ComponentShape::of(self.types, field.ty) + .is_some_and(ComponentShape::is_borrowed) + { + "storageBorrowedFieldRoot" + } else { + "storageFieldRoot" + }; + let method = self.cp.add_method_ref(owner, name, + "(Ljava/lang/Object;JLjava/lang/String;Ljava/lang/String;JJLjava/lang/String;)Ljava/lang/Object;")?; + self.assembly.code.push(Instruction::Invokestatic(method)); + } + Op::ProjectOffset { root, base, offset } => { + self.load(root)?; + self.load(base)?; + self.load(offset)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + "storageFieldOffset", + "(Ljava/lang/Object;Ljava/lang/Object;J)J", + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + } + Op::AddressTag(value) => { + self.load(value)?; + self.assembly.code.push(Instruction::Lconst_0); + self.location_tag()?; + } + Op::LocationTag(parts) => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + self.location_tag()?; + } + Op::AddressViewPart { address, index } => { + self.load(address)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let (name, descriptor) = if index == 0 { + ("sliceBackingArray", "()Ljava/lang/Object;") + } else { + ("sliceElementOffset", "()I") + }; + let method = self.cp.add_method_ref(owner, name, descriptor)?; + self.assembly.code.push(Instruction::Invokevirtual(method)); + } + Op::AddressEqual { left, right } | Op::AddressCompare { left, right } => { + self.load(left)?; + self.assembly.code.push(Instruction::Lconst_0); + self.load(right)?; + self.assembly.code.push(Instruction::Lconst_0); + self.location_comparison(matches!(inst.op, Op::AddressCompare { .. }))?; + } + Op::LocationEqual(parts) | Op::LocationCompare(parts) => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + self.location_comparison(matches!(inst.op, Op::LocationCompare(_)))?; + } + Op::AddressPart { address, index: 0 } => self.load(address)?, + Op::AddressPart { index: 1, .. } => self.assembly.code.push(Instruction::Lconst_0), + Op::SlotRoot(slot) => self + .assembly + .code + .push(Kind::Reference.load(self.storage[slot.index()])), + Op::AddressPack(parts) => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + let Some(Type::Pointer(inner)) = + self.types.get(self.body.value_type(inst.result.unwrap())) + else { + return Err(error("address pack needs pointer type")); + }; + self.materialize_address(inner)?; + } + Op::LoadAddressCopy(parts) => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + self.read_object_address(self.body.value_type(inst.result.unwrap()), true)?; + } + Op::LoadAddress(parts) => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + self.read_address(self.body.value_type(inst.result.unwrap()))?; + } + Op::StoreAddress { parts, value } => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + self.load(value)?; + self.write_address(self.body.value_type(value))?; + } + _ => return Ok(false), + } + Ok(true) + } + pub(super) fn materialize_address(&mut self, ty: TypeId) -> jvm::Result<()> { + if let Some(Type::Layout(id)) = self.types.get(ty) { + let AddressLayout { size, codec, .. } = self.types.get_layout(id); + return super::super::abi::materialize_typed_address( + self.cp, + &mut self.assembly.code, + size, + codec.map(|s| self.types.symbol_name(s).unwrap()), + ); + } + super::super::abi::materialize_address( + self.cp, + &mut self.assembly.code, + super::super::abi::address_plan(self.types, ty), + ) + } + fn location_tag(&mut self) -> jvm::Result<()> { + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = + self.cp + .add_method_ref(owner, "nullableLocationTag", "(Ljava/lang/Object;J)J")?; + self.assembly.code.push(Instruction::Invokestatic(method)); + Ok(()) + } + fn location_comparison(&mut self, ordering: bool) -> jvm::Result<()> { + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + if ordering { + "compareLocations" + } else { + "sameLocation" + }, + if ordering { + "(Ljava/lang/Object;JLjava/lang/Object;J)I" + } else { + "(Ljava/lang/Object;JLjava/lang/Object;J)Z" + }, + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + Ok(()) + } + pub(super) fn read_address(&mut self, ty: TypeId) -> jvm::Result<()> { + if self.types.get(ty).is_some_and(|t| t.carrier() == 5) { + self.read_object_address(ty, false)?; + return Ok(()); + } + let size = StorageSlot::scalar(ty, self.types) + .ok_or_else(|| error("scalar load layout"))? + .size; + self.assembly + .code + .push(get_int_const_instr(self.cp, size as i32)); + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = + self.cp + .add_method_ref(owner, "loadLocationBits", "(Ljava/lang/Object;JI)J")?; + self.assembly.code.push(Instruction::Invokestatic(method)); + self.bits_to_scalar(ty) + } + pub(super) fn bits_to_scalar(&mut self, ty: TypeId) -> jvm::Result<()> { + match self.types.get(ty) { + Some(Type::Scalar(ScalarType::F64)) => { + self.address_float("java/lang/Double", "longBitsToDouble", "(J)D")? + } + Some(Type::Scalar(ScalarType::I64 | ScalarType::U64)) => {} + Some(Type::Scalar(scalar)) => { + self.assembly.code.push(Instruction::L2i); + if scalar == ScalarType::F32 { + self.address_float("java/lang/Float", "intBitsToFloat", "(I)F")?; + } + } + _ => return Err(error("non-scalar address load")), + } + Ok(()) + } + pub(super) fn read_object_address(&mut self, ty: TypeId, owned: bool) -> jvm::Result<()> { + let name = self.address_target(ty)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + if owned { + "loadStorageCopy" + } else if matches!(self.types.get(ty), Some(Type::Array(_))) { + "loadStorageArray" + } else { + "loadStorageLocation" + }, + "(Ljava/lang/Object;JLjava/lang/String;)Ljava/lang/Object;", + )?; + let class = self.cp.add_class(&name)?; + self.assembly.code.extend([ + Instruction::Invokestatic(method), + Instruction::Checkcast(class), + ]); + Ok(()) + } + pub(super) fn address_layout(&mut self, size: u32, codec: Option) -> jvm::Result<()> { + self.assembly + .code + .push(get_int_const_instr(self.cp, size as i32)); + self.assembly.code.push(match codec { + Some(id) => Instruction::Ldc_w( + self.cp + .add_name_string(self.types.symbol_name(id).unwrap())?, + ), + None => Instruction::Aconst_null, + }); + Ok(()) + } + pub(super) fn address_target(&mut self, ty: TypeId) -> jvm::Result { + let mut descriptor = String::new(); + representation::descriptor(self.types, ty, &mut descriptor)?; + let name = descriptor + .strip_prefix('L') + .and_then(|s| s.strip_suffix(';')) + .unwrap_or(&descriptor); + self.assembly + .code + .push(if matches!(self.types.get(ty), Some(Type::Pointer(_))) { + Instruction::Aconst_null + } else { + Instruction::Ldc_w(self.cp.add_name_string(name)?) + }); + Ok(name.into()) + } + pub(super) fn write_address(&mut self, ty: TypeId) -> jvm::Result<()> { + if self.types.get(ty).is_some_and(|t| t.carrier() == 5) { + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + "storeStorageLocation", + "(Ljava/lang/Object;JLjava/lang/Object;)V", + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + return Ok(()); + } + let size = StorageSlot::scalar(ty, self.types) + .ok_or_else(|| error("scalar store layout"))? + .size; + self.scalar_to_bits(ty)?; + self.assembly + .code + .push(get_int_const_instr(self.cp, size as i32)); + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = + self.cp + .add_method_ref(owner, "storeLocationBits", "(Ljava/lang/Object;JJI)V")?; + self.assembly.code.push(Instruction::Invokestatic(method)); + Ok(()) + } + pub(super) fn scalar_to_bits(&mut self, ty: TypeId) -> jvm::Result<()> { + match self.types.get(ty) { + Some(Type::Scalar(ScalarType::F64)) => { + self.address_float("java/lang/Double", "doubleToRawLongBits", "(D)J")? + } + Some(Type::Scalar(ScalarType::I64 | ScalarType::U64)) => {} + Some(Type::Scalar(scalar)) => { + if scalar == ScalarType::F32 { + self.address_float("java/lang/Float", "floatToRawIntBits", "(F)I")?; + } + self.assembly.code.push(Instruction::I2l); + } + _ => return Err(error("non-scalar address store")), + } + Ok(()) + } + fn address_float(&mut self, class: &str, name: &str, descriptor: &str) -> jvm::Result<()> { + let owner = self.cp.add_class(class)?; + let method = self.cp.add_method_ref(owner, name, descriptor)?; + self.assembly.code.push(Instruction::Invokestatic(method)); + Ok(()) + } +} diff --git a/compiler-core/src/jvm/select/allocate.rs b/compiler-core/src/jvm/select/allocate.rs index 8d288c97..79275443 100644 --- a/compiler-core/src/jvm/select/allocate.rs +++ b/compiler-core/src/jvm/select/allocate.rs @@ -18,7 +18,6 @@ pub(super) fn allocate( types: &Types, live: &crate::opt::Live, forwarded: &[bool], - relative_pointer_abi: bool, debug: Option<&DebugInfo>, order: &[BlockId], ) -> jvm::Result { @@ -35,14 +34,7 @@ pub(super) fn allocate( .count .checked_add(value_kind(body, types, param)?.width()) .ok_or_else(|| error("JVM local limit"))?; - if relative_pointer_abi - && matches!(types.get(body.value_type(param)), Some(Type::Pointer(_))) - { - result.count = result - .count - .checked_add(4) - .ok_or_else(|| error("JVM parameter limit"))?; - } + active.push(Reverse((intervals[param.index()].1, param))); } let mut order = (0..body.values.len()) @@ -270,7 +262,6 @@ mod tests { &types, &crate::opt::live(&body, &types), &vec![false; body.values.len()], - false, None, &body.layout(), ) diff --git a/compiler-core/src/jvm/select/arrays.rs b/compiler-core/src/jvm/select/arrays.rs index 409523d4..c9e05c5c 100644 --- a/compiler-core/src/jvm/select/arrays.rs +++ b/compiler-core/src/jvm/select/arrays.rs @@ -1,7 +1,36 @@ //! Array operations preserve lazy repeat copies and pointer-backed slice views. use super::*; +#[derive(Clone, Copy, PartialEq, Eq)] +enum Access { + Native, + Array, + View, +} + impl Selector<'_> { - pub(super) fn array(&mut self, inst: Inst) -> jvm::Result { + pub(super) fn array(&mut self, id: InstId, inst: Inst) -> jvm::Result { + if let Op::ArrayFill { array, value } = inst.op { + let element = self.body.value_type(value); + let native = self.native_array_accesses.get(id.index()) == Some(&Some(element)); + self.load(array)?; + self.argument(value)?; + let mut descriptor = String::from("("); + representation::descriptor(self.types, self.body.value_type(array), &mut descriptor)?; + representation::descriptor(self.types, element, &mut descriptor)?; + descriptor.push_str(")V"); + let owner = self.cp.add_class(if native { + "java/util/Arrays" + } else { + POINTER_CLASS + })?; + let target = self.cp.add_method_ref( + owner, + if native { "fill" } else { "fillArray" }, + descriptor, + )?; + self.assembly.code.push(Instruction::Invokestatic(target)); + return Ok(true); + } if let Op::ArrayLength(array) = inst.op { self.load(array)?; let op = if matches!( @@ -16,13 +45,49 @@ impl Selector<'_> { self.assembly.code.push(op); return Ok(true); } - let (array, index, value) = match inst.op { - Op::ArrayGet { array, index } => (array, index, None), + if let Op::ViewGet(parts) | Op::ViewSet { parts, .. } = inst.op { + let [backing, start, index] = self.body.args[parts.range()] else { + return Err(error("slice access components")); + }; + let value = match inst.op { + Op::ViewSet { value, .. } => Some(value), + _ => None, + }; + let element = self.body.value_type(value.or(inst.result).unwrap()); + let native = self.native_array_accesses.get(id.index()) == Some(&Some(element)); + self.load(backing)?; + if native { + let mut descriptor = String::from("["); + representation::descriptor(self.types, element, &mut descriptor)?; + self.assembly + .code + .push(Instruction::Checkcast(self.cp.add_class(&descriptor)?)); + } + self.load(start)?; + self.load(index)?; + self.assembly.code.push(Instruction::Iadd); + if let Some(value) = value { + self.argument(value)?; + } + self.array_access( + element, + value.is_some(), + if native { Access::Native } else { Access::View }, + )?; + return Ok(true); + } + let (array, index, value, native) = match inst.op { + Op::ArrayGet { + array, + index, + native, + } => (array, index, None, native), Op::ArraySet { array, index, value, - } => (array, index, Some(value)), + native, + } => (array, index, Some(value), native), _ => return Ok(false), }; let representation = self.types.get(self.body.value_type(array)); @@ -72,6 +137,25 @@ impl Selector<'_> { if let Some(value) = value { self.argument(value)?; } + let native = native || self.native_array_accesses.get(id.index()) == Some(&Some(element)); + let primitive = matches!(self.types.get(element), Some(Type::Scalar(_))); + // A raw pointer can expose this array to a decoded aggregate alias later. + // Only a whole-lifetime proof can remove coherence checks. + self.array_access( + element, + value.is_some(), + if native { + Access::Native + } else if view || primitive { + Access::View + } else { + Access::Array + }, + )?; + Ok(true) + } + + fn array_access(&mut self, element: TypeId, store: bool, access: Access) -> jvm::Result<()> { use ScalarType::*; let (suffix, descriptor, read, write) = match self.types.get(element) { Some(Type::Scalar(Bool)) => ("Boolean", "Z", Instruction::Baload, Instruction::Bastore), @@ -97,24 +181,21 @@ impl Selector<'_> { Instruction::Aastore, ), }; - if view || (suffix == "Object" && value.is_none()) { + if access == Access::View || (access == Access::Array && suffix == "Object" && !store) { let owner = self.cp.add_class(POINTER_CLASS)?; - let name = if view { - format!( - "slice{}{suffix}", - if value.is_some() { "Set" } else { "Get" } - ) + let name = if access == Access::View { + format!("slice{}{suffix}", if store { "Set" } else { "Get" }) } else { "arrayGetObject".into() }; - let signature = if value.is_some() { + let signature = if store { format!("(Ljava/lang/Object;I{descriptor})V") } else { format!("(Ljava/lang/Object;I){descriptor}") }; let method = self.cp.add_method_ref(owner, name, signature)?; self.assembly.code.push(Instruction::Invokestatic(method)); - if suffix == "Object" && value.is_none() { + if suffix == "Object" && !store { let mut name = String::new(); representation::descriptor(self.types, element, &mut name)?; let name = name @@ -126,10 +207,8 @@ impl Selector<'_> { .push(Instruction::Checkcast(self.cp.add_class(name)?)); } } else { - self.assembly - .code - .push(if value.is_some() { write } else { read }); + self.assembly.code.push(if store { write } else { read }); } - Ok(true) + Ok(()) } } diff --git a/compiler-core/src/jvm/select/debug_tests.rs b/compiler-core/src/jvm/select/debug_tests.rs index 981d9bb8..72499201 100644 --- a/compiler-core/src/jvm/select/debug_tests.rs +++ b/compiler-core/src/jvm/select/debug_tests.rs @@ -46,6 +46,15 @@ fn debug_mirrors_keep_dead_values_and_hide_uninitialized_join_bindings() { value: seven, }, ); + // Keep executable work after initializing the conditional binding. + // An empty edge has no bytecode range for its debug information. + debug.push( + &b, + DebugChange::Set { + local: 1, + value: parameter, + }, + ); b.jump(join, vec![]); b.switch_to(join); debug.push(&b, DebugChange::Scope(0)); diff --git a/compiler-core/src/jvm/select/fields.rs b/compiler-core/src/jvm/select/fields.rs index bb840018..4b31a8d0 100644 --- a/compiler-core/src/jvm/select/fields.rs +++ b/compiler-core/src/jvm/select/fields.rs @@ -1,15 +1,63 @@ -//! Select promoted field pointers without reflective runtime field cells. +//! Select projected fields through authoritative typed storage when available. use super::*; +mod bytes; + impl Selector<'_> { pub(super) fn field_memory(&mut self, inst: Inst) -> jvm::Result { - let (base, projection, value) = match inst.op { - Op::LoadField { base, projection } => (base, projection, None), + let address = match inst.op { + Op::LoadStorageField { address, .. } | Op::LoadStorageFieldCopy { address, .. } => { + Some(&self.body.args[address.range()]) + } + Op::StoreStorageField { args, .. } => Some(&self.body.args[args.range()][..2]), + _ => None, + }; + let owned = matches!( + inst.op, + Op::LoadFieldCopy { .. } | Op::LoadStorageFieldCopy { .. } + ); + let (base, projection, value, part, components) = match inst.op { + Op::LoadStorageFieldCopy { projection, .. } => { + (address.unwrap()[0], projection, None, None, None) + } + Op::LoadStorageField { + projection, index, .. + } => (address.unwrap()[0], projection, None, index, None), + Op::StoreStorageField { + args, + projection, + split, + } => { + let values = List { + start: args.start + 2, + len: args.len - 2, + }; + ( + address.unwrap()[0], + projection, + (!split).then(|| self.body.args[values.range().start]), + None, + split.then_some(values), + ) + } + Op::LoadField { base, projection } | Op::LoadFieldCopy { base, projection } => { + (base, projection, None, None, None) + } Op::StoreField { base, projection, value, - } => (base, projection, Some(value)), + } => (base, projection, Some(value), None, None), + Op::LoadFieldPart { + base, + projection, + index, + } => (base, projection, None, Some(index), None), + Op::StoreFieldParts { + base, + projection, + parts, + } => (base, projection, None, None, Some(parts)), _ => return Ok(false), }; let field = &self.body.fields[self.body.projections[projection.index()].field.index()]; @@ -17,50 +65,252 @@ impl Selector<'_> { return Err(error("field memory requires a concrete owner")); }; let owner = self.cp.add_class(self.types.symbol_name(owner).unwrap())?; - let mut descriptor = String::new(); - representation::descriptor(self.types, field.ty, &mut descriptor)?; - let member = self.cp.add_field_ref(owner, &field.name, &descriptor)?; let pointer = self.cp.add_class(POINTER_CLASS)?; let direct = self.cp.add_method_ref( pointer, "directAggregate", "(Ljava/lang/Class;)Ljava/lang/Object;", )?; + let split = part.is_some() || components.is_some(); + let physical = if split { + let names = match self.types.get(field.ty) { + Some(Type::TaggedI64) => super::super::abi::tagged_field_names(&field.name) + .into_iter() + .zip(["J", "J"]) + .collect(), + Some(Type::Pointer(_)) => { + vec![ + (field.name.clone(), "Ljava/lang/Object;"), + ( + super::super::abi::address_displacement_name( + self.types, + field.ty, + &field.name, + ), + "J", + ), + ] + } + Some(Type::Slice(_) | Type::Str) => super::super::abi::view_field_names( + &field.name, + matches!(self.types.get(field.ty), Some(Type::Str)), + ) + .into_iter() + .zip(["Ljava/lang/Object;", "I", "J"]) + .collect(), + _ => return Err(error("invalid split field")), + }; + Some( + names + .into_iter() + .map(|(name, ty)| self.cp.add_field_ref(owner, &name, ty)) + .collect::>>()?, + ) + } else { + None + }; + let member = if let Some(physical) = &physical { + physical[part.unwrap_or(0) as usize] + } else { + let mut descriptor = String::new(); + representation::descriptor(self.types, field.ty, &mut descriptor)?; + self.cp.add_field_ref(owner, &field.name, &descriptor)? + }; let slow = self.assembly.label(); let done = self.assembly.label(); - // Keep the pointer for committing the live carrier, or projecting just - // this field when storage is raw or only partially initialized. + let store = value.is_some() || components.is_some(); self.load(base)?; - self.assembly.code.extend([ - Instruction::Dup, - Instruction::Ldc_w(owner), - Instruction::Invokevirtual(direct), - Instruction::Dup, - ]); + // Adjacent field components share one resolved aggregate object. + // Mutations, opaque calls and control-flow boundaries invalidate it. + let cached = part.is_some(); + let key = ( + self.body.resolve(base), + address.map(|a| self.body.resolve(a[1])), + field.owner, + ); + if cached && self.aggregate_cache == Some(key) { + self.assembly + .code + .push(Kind::Reference.load(self.aggregate_slot.unwrap())); + } else { + self.assembly.code.push(Instruction::Dup); + if let Some(parts) = address { + self.load(parts[1])?; + self.assembly.code.push(Instruction::Ldc_w(owner)); + let resolve = self.cp.add_method_ref( + pointer, + "directStorageAggregate", + "(Ljava/lang/Object;JLjava/lang/Class;)Ljava/lang/Object;", + )?; + self.assembly.code.push(Instruction::Invokestatic(resolve)); + } else { + self.assembly.code.extend([ + Instruction::Ldc_w(owner), + Instruction::Invokevirtual(direct), + ]); + } + if cached { + let slot = if let Some(slot) = self.aggregate_slot { + slot + } else { + let slot = self.next_slot; + self.next_slot = self + .next_slot + .checked_add(1) + .ok_or_else(|| error("JVM local limit"))?; + self.aggregate_slot = Some(slot); + slot + }; + self.assembly + .code + .extend([Instruction::Dup, Kind::Reference.store(slot)]); + self.aggregate_cache = Some(key); + } + } + self.assembly.code.push(Instruction::Dup); self.assembly.branch(Instruction::Ifnull(0), slow); - if value.is_none() { + if !store { self.assembly .code .extend([Instruction::Swap, Instruction::Pop]); } self.assembly.code.push(Instruction::Checkcast(owner)); - if let Some(value) = value { + if let Some(parts) = components { + let values = &self.body.args[parts.range()]; + for (index, &value) in values.iter().enumerate() { + if index + 1 < values.len() { + self.assembly.code.push(Instruction::Dup); + } + self.load(value)?; + self.assembly + .code + .push(Instruction::Putfield(physical.as_ref().unwrap()[index])); + } + } else if let Some(value) = value { self.argument(value)?; self.assembly.code.push(Instruction::Putfield(member)); - let commit = self.cp.add_method_ref(pointer, "commitMemoryView", "()V")?; - self.assembly.code.push(Instruction::Invokevirtual(commit)); } else { self.assembly.code.push(Instruction::Getfield(member)); + if owned { + self.copy_value(field.ty)?; + } + } + if store { + if let Some(parts) = address { + self.load(parts[1])?; + let commit = self.cp.add_method_ref( + pointer, + "commitStorageLocation", + "(Ljava/lang/Object;J)V", + )?; + self.assembly.code.push(Instruction::Invokestatic(commit)); + } else { + let commit = self.cp.add_method_ref(pointer, "commitMemoryView", "()V")?; + self.assembly.code.push(Instruction::Invokevirtual(commit)); + } } self.assembly.branch(Instruction::Goto_w(0), done); self.assembly.bind(slow); self.assembly.code.push(Instruction::Pop); + if owned { + if let Some(parts) = address { + self.load(parts[1])?; + } else { + self.assembly.code.push(Instruction::Lconst_0); + } + self.project_field_arguments(projection)?; + let target = self.address_target(field.ty)?; + let method = self.cp.add_method_ref(pointer, "loadStorageFieldCopy", + "(Ljava/lang/Object;JLjava/lang/String;Ljava/lang/String;JJLjava/lang/String;Ljava/lang/String;)Ljava/lang/Object;")?; + let target = self.cp.add_class(target)?; + self.assembly.code.extend([ + Instruction::Invokestatic(method), + Instruction::Checkcast(target), + ]); + self.assembly.bind(done); + return Ok(true); + } + // Keep managed access direct. Share byte and fallback dispatch in the + // runtime because larger generated methods can prevent JVM inlining. + if let Some(value) = value + && matches!( + self.types.get(field.ty), + Some(Type::Class(_) | Type::Array(_)) + ) + { + if let Some(parts) = address { + self.load(parts[1])?; + } else { + self.assembly.code.push(Instruction::Lconst_0); + } + self.project_field_arguments(projection)?; + self.argument(value)?; + let method = self.cp.add_method_ref(pointer, "storeStorageField", + "(Ljava/lang/Object;JLjava/lang/String;Ljava/lang/String;JJLjava/lang/String;Ljava/lang/Object;)V")?; + self.assembly.code.push(Instruction::Invokestatic(method)); + self.assembly.bind(done); + return Ok(true); + } + if !split && self.scalar_field(address.map(|parts| parts[1]), projection, value)? { + self.assembly.bind(done); + return Ok(true); + } + if let Some(parts) = address { + self.load(parts[1])?; + let materialize = self.cp.add_method_ref( + pointer, + "fromStorageLocation", + "(Ljava/lang/Object;J)Lorg/rustlang/runtime/Pointer;", + )?; + self.assembly + .code + .push(Instruction::Invokestatic(materialize)); + } self.project_field(projection)?; - if let Some(value) = value { + if let Some(parts) = components { + if let Some(Type::Pointer(inner)) = self.types.get(field.ty) { + for &value in &self.body.args[parts.range()] { + self.load(value)?; + } + self.materialize_address(inner)?; + } else if matches!(self.types.get(field.ty), Some(Type::TaggedI64)) { + self.materialize_tagged(parts)?; + } else { + self.materialize_view(field.ty, parts)?; + } + self.write_memory(field.ty)?; + } else if let Some(value) = value { self.argument(value)?; self.write_memory(field.ty)?; } else { self.read_memory(field.ty)?; + if let Some(index) = part { + if matches!(self.types.get(field.ty), Some(Type::Pointer(_))) { + if index == 1 { + self.assembly + .code + .extend([Instruction::Pop, Instruction::Lconst_0]); + } + } else if matches!(self.types.get(field.ty), Some(Type::TaggedI64)) { + let owner = self.cp.add_class(super::super::abi::TAGGED_LONG_CLASS)?; + self.assembly + .code + .push(Instruction::Invokestatic(self.cp.add_method_ref( + owner, + if index == 0 { "value" } else { "tag" }, + "(Lorg/rustlang/runtime/TaggedLong;)J", + )?)); + } else { + let slice = self.cp.add_class(representation::SLICE_VIEW_CLASS)?; + let (name, ty) = [ + ("array", "Ljava/lang/Object;"), + ("offset", "I"), + ("rustLength", "J"), + ][index as usize]; + let member = self.cp.add_field_ref(slice, name, ty)?; + self.assembly.code.push(Instruction::Getfield(member)); + } + } } self.assembly.bind(done); Ok(true) diff --git a/compiler-core/src/jvm/select/fields/bytes.rs b/compiler-core/src/jvm/select/fields/bytes.rs new file mode 100644 index 00000000..0ded393d --- /dev/null +++ b/compiler-core/src/jvm/select/fields/bytes.rs @@ -0,0 +1,70 @@ +//! Scalar field fallback with byte-storage dispatch shared in the runtime. +use super::*; + +impl Selector<'_> { + /// Consume the root already on the stack, preserving stack forwarding. + pub(super) fn scalar_field( + &mut self, + offset: Option, + projection: ProjectionId, + value: Option, + ) -> jvm::Result { + let layout = &self.body.projections[projection.index()]; + let field = &self.body.fields[layout.field.index()]; + let ty = field.ty; + if layout.codec.is_some() + || !StorageSlot::scalar(ty, self.types) + .is_some_and(|s| (1..=8).contains(&s.size) && u64::from(s.size) == layout.size) + { + return Ok(false); + } + let Some(Type::Class(owner)) = self.types.get(field.owner) else { + return Err(error("scalar field requires a concrete owner")); + }; + if let Some(offset) = offset { + self.load(offset)?; + } else { + self.assembly.code.push(Instruction::Lconst_0); + } + self.assembly.code.extend([ + Instruction::Ldc_w( + self.cp + .add_name_string(self.types.symbol_name(owner).unwrap())?, + ), + Instruction::Ldc_w(self.cp.add_string(&field.name)?), + get_long_const_instr(self.cp, layout.offset as i64), + ]); + if let Some(value) = value { + self.argument(value)?; + self.scalar_to_bits(ty)?; + } + self.assembly + .code + .push(get_int_const_instr(self.cp, layout.size as i32)); + let pointer = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + pointer, + if value.is_some() { + "storeScalarField" + } else { + "loadScalarField" + }, + if value.is_some() { + "(Ljava/lang/Object;JLjava/lang/String;Ljava/lang/String;JJI)V" + } else { + "(Ljava/lang/Object;JLjava/lang/String;Ljava/lang/String;JI)J" + }, + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + if value.is_none() { + self.bits_to_scalar(ty)?; + // Integer result normalization is shared with all other IR loads. + match self.types.get(ty) { + Some(Type::Scalar(ScalarType::F16)) => self.assembly.code.push(Instruction::I2s), + Some(Type::Scalar(ScalarType::Char)) => self.assembly.code.push(Instruction::I2c), + _ => {} + } + } + Ok(true) + } +} diff --git a/compiler-core/src/jvm/select/forward.rs b/compiler-core/src/jvm/select/forward.rs index 09852e38..8b1f1329 100644 --- a/compiler-core/src/jvm/select/forward.rs +++ b/compiler-core/src/jvm/select/forward.rs @@ -57,24 +57,25 @@ fn first_operand(body: &Body, op: Op) -> Option { Op::Neg(value) | Op::Not(value) | Op::Cast(value) + | Op::Refine(value) | Op::Reinterpret(value) | Op::Adapt(value) | Op::NewArray(value) | Op::ArrayLength(value) | Op::Opaque(value) | Op::Load(value) + | Op::LoadCopy(value) | Op::Length(value) | Op::SetStatic { value, .. } => Some(value), - Op::Project { base, .. } | Op::LoadField { base, .. } | Op::StoreField { base, .. } => { - Some(base) - } + Op::Project { base, .. } + | Op::LoadField { base, .. } + | Op::LoadFieldCopy { base, .. } + | Op::StoreField { base, .. } + | Op::LoadFieldPart { base, .. } + | Op::StoreFieldParts { base, .. } => Some(base), Op::Offset { pointer, .. } | Op::Store { pointer, .. } => Some(pointer), Op::ViewData { view, .. } => Some(view), - Op::GetField { object, field } | Op::SetField { object, field, .. } - if !body.fields[field.index()].relative_pointer => - { - Some(object) - } + Op::GetField { object, .. } | Op::SetField { object, .. } => Some(object), Op::Call { kind, args, .. } if kind != CallKind::Constructor => { body.args[args.range()].first().copied() } @@ -98,17 +99,9 @@ mod tests { let result = b.constant(int, Scalar::integer(ScalarType::I32, 7).unwrap()); b.terminate(Terminator::Return(Some(result))); let body = b.finish().unwrap(); - let code = compile_with_options( - &body, - &types, - &mut Default::default(), - Options { - relative_pointer_abi: true, - ..Default::default() - }, - ) - .unwrap(); - assert_eq!(code.max_locals, 6); // Pointer + two offsets + byte ABI slots. + let code = compile_with_options(&body, &types, &mut Default::default(), Options::default()) + .unwrap(); + assert_eq!(code.max_locals, 2); // Pointer and byte ABI slots. assert_eq!( code.instructions, vec![Instruction::Bipush(7), Instruction::Ireturn] diff --git a/compiler-core/src/jvm/select/general.rs b/compiler-core/src/jvm/select/general.rs index c063d52f..de3b2371 100644 --- a/compiler-core/src/jvm/select/general.rs +++ b/compiler-core/src/jvm/select/general.rs @@ -5,6 +5,21 @@ use jvm::attributes::{ArrayType, BootstrapMethod}; impl Selector<'_> { pub(super) fn general(&mut self, inst: Inst) -> jvm::Result { match inst.op { + Op::Heap { operation, args } => { + let (name, descriptor) = match operation { + HeapOp::Allocate => ("allocate", "(JJ)Ljava/lang/Object;"), + HeapOp::Reallocate => { + ("reallocate", "(Ljava/lang/Object;JJJJ)Ljava/lang/Object;") + } + HeapOp::Deallocate => ("deallocate", "(Ljava/lang/Object;J)V"), + }; + for &value in &self.body.args[args.range()] { + self.load(value)?; + } + let class = self.cp.add_class("org/rustlang/runtime/Heap")?; + let method = self.cp.add_method_ref(class, name, descriptor)?; + self.assembly.code.push(Instruction::Invokestatic(method)); + } Op::Nop => {} Op::Exception => self.assembly.code.push( Kind::Reference.load( diff --git a/compiler-core/src/jvm/select/memory.rs b/compiler-core/src/jvm/select/memory.rs index 90d27001..2dd1511c 100644 --- a/compiler-core/src/jvm/select/memory.rs +++ b/compiler-core/src/jvm/select/memory.rs @@ -9,35 +9,6 @@ fn scalar(types: &Types, ty: TypeId) -> jvm::Result { } impl Selector<'_> { - pub(super) fn materialize_parameters(&mut self) -> jvm::Result<()> { - for ¶m in &self.body.blocks[self.body.entry.index()].params { - // The ABI root reserves the slot; only additional uses need an - // actual pointer value. Unused arguments require no materialization. - if self.live.uses[param.index()] == 1 - || !matches!( - self.types.get(self.body.value_type(param)), - Some(Type::Pointer(_)) - ) - { - continue; - } - let slot = self.slot(param); - let owner = self.cp.add_class(POINTER_CLASS)?; - let method = self.cp.add_method_ref( - owner, - "materializeRelative", - "(Lorg/rustlang/runtime/Pointer;JJ)Lorg/rustlang/runtime/Pointer;", - )?; - self.assembly.code.extend([ - Kind::Reference.load(slot), - Kind::Long.load(slot + 1), - Kind::Long.load(slot + 3), - Instruction::Invokestatic(method), - Kind::Reference.store(slot), - ]); - } - Ok(()) - } pub(super) fn initialize_storage(&mut self) -> jvm::Result<()> { for storage in &self.body.slots { let owner = self.cp.add_class(POINTER_CLASS)?; @@ -46,7 +17,7 @@ impl Selector<'_> { let code = match scalar { Bool => 4, I8 | U8 => 8, - I16 => 9, + I16 | F16 => 9, U16 => 5, I32 | U32 => 10, I64 | U64 => 11, @@ -56,17 +27,10 @@ impl Selector<'_> { }; let array = jvm::attributes::ArrayType::from_bytes(&mut jvm::ByteReader::new(&[code]))?; - self.assembly.code.extend([ - Instruction::Iconst_1, - Instruction::Newarray(array), - Instruction::Iconst_0, - get_int_const_instr(self.cp, storage.size as i32), - ]); - self.cp.add_method_ref( - owner, - "array", - "(Ljava/lang/Object;II)Lorg/rustlang/runtime/Pointer;", - )? + self.assembly + .code + .extend([Instruction::Iconst_1, Instruction::Newarray(array)]); + None } else { self.assembly.code.push(Instruction::Aconst_null); self.assembly @@ -83,11 +47,11 @@ impl Selector<'_> { self.assembly .code .push(get_int_const_instr(self.cp, storage.alignment as i32)); - self.cp.add_method_ref( + Some(self.cp.add_method_ref( owner, "cellAligned", "(Ljava/lang/Object;ILjava/lang/String;I)Lorg/rustlang/runtime/Pointer;", - )? + )?) }; let slot = self.next_slot; self.next_slot = self @@ -95,10 +59,10 @@ impl Selector<'_> { .checked_add(1) .ok_or_else(|| error("JVM storage slot limit"))?; self.storage.push(slot); - self.assembly.code.extend([ - Instruction::Invokestatic(target), - Kind::Reference.store(slot), - ]); + if let Some(target) = target { + self.assembly.code.push(Instruction::Invokestatic(target)); + } + self.assembly.code.push(Kind::Reference.store(slot)); } Ok(()) } @@ -108,6 +72,10 @@ impl Selector<'_> { return Ok(true); } match inst.op { + Op::CopyValue(value) => { + self.load(value)?; + self.copy_value(self.body.value_type(value))?; + } Op::Opaque(value) => self.load(value)?, Op::Project { base, projection } => { self.load(base)?; @@ -140,20 +108,37 @@ impl Selector<'_> { )?; self.assembly.code.push(Instruction::Invokestatic(method)); } - Op::AddressOfSlot(slot) => self - .assembly - .code - .push(Kind::Reference.load(self.storage[slot.index()])), + Op::AddressOfSlot(slot) => { + self.assembly + .code + .push(Kind::Reference.load(self.storage[slot.index()])); + let ty = self.body.slots[slot.index()].ty; + if StorageSlot::scalar(ty, self.types).is_some() { + self.assembly.code.push(Instruction::Lconst_0); + self.materialize_address(ty)?; + } + } Op::LoadSlot(slot) => { self.assembly .code .push(Kind::Reference.load(self.storage[slot.index()])); - self.read_memory(self.body.slots[slot.index()].ty)?; + let ty = self.body.slots[slot.index()].ty; + if StorageSlot::scalar(ty, self.types).is_some() { + self.assembly.code.push(Instruction::Lconst_0); + self.read_address(ty)?; + } else { + self.read_memory(ty)?; + } } Op::StoreSlot { slot, value } => { self.assembly .code .push(Kind::Reference.load(self.storage[slot.index()])); + let storage = &self.body.slots[slot.index()]; + let scalar = StorageSlot::scalar(storage.ty, self.types).is_some(); + if scalar { + self.assembly.code.push(Instruction::Lconst_0); + } self.argument(value)?; let storage = &self.body.slots[slot.index()]; if storage.size == 0 { @@ -164,10 +149,23 @@ impl Selector<'_> { "(Ljava/lang/Object;)V", )?; self.assembly.code.push(Instruction::Invokevirtual(init)); + } else if scalar { + self.write_address(storage.ty)?; } else { self.write_memory(storage.ty)?; } } + Op::Commit(pointer) => { + self.load(pointer)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref(owner, "commitMemoryView", "()V")?; + self.assembly.code.push(Instruction::Invokevirtual(method)); + } + Op::LoadCopy(pointer) => { + self.load(pointer)?; + self.assembly.code.push(Instruction::Lconst_0); + self.read_object_address(self.body.value_type(inst.result.unwrap()), true)?; + } Op::Load(pointer) => { self.load(pointer)?; self.read_memory(self.body.value_type(inst.result.unwrap()))?; @@ -193,13 +191,9 @@ impl Selector<'_> { else { return Err(error("unsupported reference cast")); }; - let to = scalar(self.types, pointee)?; - let bytes = match to { - ScalarType::Bool => 1, - ScalarType::F32 => 4, - ScalarType::F64 => 8, - _ => to.integer().ok_or_else(|| error("pointer view type"))?.0 / 8, - }; + let bytes = StorageSlot::scalar(pointee, self.types) + .ok_or_else(|| error("pointer view type"))? + .size; self.load(value)?; self.assembly.code.extend([ get_int_const_instr(self.cp, bytes as i32), @@ -218,8 +212,37 @@ impl Selector<'_> { Ok(true) } + pub(super) fn copy_value(&mut self, ty: TypeId) -> jvm::Result<()> { + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + "copyManagedValue", + "(Ljava/lang/Object;)Ljava/lang/Object;", + )?; + let mut name = String::new(); + descriptor(self.types, ty, &mut name)?; + let name = name + .strip_prefix('L') + .and_then(|s| s.strip_suffix(';')) + .unwrap_or(&name); + let class = self.cp.add_class(name)?; + self.assembly.code.extend([ + Instruction::Invokestatic(method), + Instruction::Checkcast(class), + ]); + Ok(()) + } + /// Project the pointer already on the operand stack. pub(super) fn project_field(&mut self, projection: ProjectionId) -> jvm::Result<()> { + self.project_field_arguments(projection)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref(owner, "projectStructField", "(Ljava/lang/String;Ljava/lang/String;JJLjava/lang/String;)Lorg/rustlang/runtime/Pointer;")?; + self.assembly.code.push(Instruction::Invokevirtual(method)); + Ok(()) + } + + pub(super) fn project_field_arguments(&mut self, projection: ProjectionId) -> jvm::Result<()> { let projection = &self.body.projections[projection.index()]; let field = &self.body.fields[projection.field.index()]; let Some(Type::Class(symbol)) = self.types.get(field.owner) else { @@ -230,7 +253,7 @@ impl Selector<'_> { self.cp .add_name_string(self.types.symbol_name(symbol).unwrap())?, ), - Instruction::Ldc_w(self.cp.add_name_string(&field.name)?), + Instruction::Ldc_w(self.cp.add_string(&field.name)?), get_long_const_instr(self.cp, projection.offset as i64), get_long_const_instr(self.cp, projection.size as i64), match &projection.codec { @@ -238,12 +261,23 @@ impl Selector<'_> { None => Instruction::Aconst_null, }, ]); - let owner = self.cp.add_class(POINTER_CLASS)?; - let method = self.cp.add_method_ref(owner, "projectStructField", "(Ljava/lang/String;Ljava/lang/String;JJLjava/lang/String;)Lorg/rustlang/runtime/Pointer;")?; - self.assembly.code.push(Instruction::Invokevirtual(method)); Ok(()) } pub(super) fn read_memory(&mut self, ty: TypeId) -> jvm::Result<()> { + if matches!(self.types.get(ty), Some(Type::Array(_))) { + let mut descriptor = String::new(); + representation::descriptor(self.types, ty, &mut descriptor)?; + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self + .cp + .add_method_ref(owner, "getObject", "()Ljava/lang/Object;")?; + let class = self.cp.add_class(descriptor)?; + self.assembly.code.extend([ + Instruction::Invokevirtual(method), + Instruction::Checkcast(class), + ]); + return Ok(()); + } let class = match self.types.get(ty) { Some(Type::Class(symbol) | Type::Interface(symbol)) => { Some(self.types.symbol_name(symbol).unwrap()) @@ -251,13 +285,20 @@ impl Selector<'_> { Some(Type::Pointer(_)) => Some(POINTER_CLASS), Some(Type::Slice(_)) => Some(representation::SLICE_VIEW_CLASS), Some(Type::Str) => Some(representation::UTF8_VIEW_CLASS), + Some(Type::TaggedI64) => Some(super::super::abi::TAGGED_LONG_CLASS), _ => None, }; if let Some(class) = class { let owner = self.cp.add_class(POINTER_CLASS)?; let target = if matches!( self.types.get(ty), - Some(Type::Class(_) | Type::Interface(_) | Type::Slice(_) | Type::Str) + Some( + Type::Class(_) + | Type::Interface(_) + | Type::Slice(_) + | Type::Str + | Type::TaggedI64 + ) ) { self.assembly .code @@ -282,7 +323,7 @@ impl Selector<'_> { let (name, result) = match scalar(self.types, ty)? { Bool => ("getBoolean", "Z"), I8 | U8 => ("getI8", "B"), - I16 | U16 => ("getI16", "S"), + I16 | U16 | F16 => ("getI16", "S"), I32 | U32 => ("getI32", "I"), I64 | U64 => ("getI64", "J"), F32 => ("getF32", "F"), @@ -301,7 +342,13 @@ impl Selector<'_> { if matches!( self.types.get(ty), Some( - Type::Class(_) | Type::Interface(_) | Type::Pointer(_) | Type::Slice(_) | Type::Str + Type::Class(_) + | Type::Interface(_) + | Type::Pointer(_) + | Type::Slice(_) + | Type::Str + | Type::TaggedI64 + | Type::Array(_) ) ) { signature.push_str("Ljava/lang/Object;"); diff --git a/compiler-core/src/jvm/select/mod.rs b/compiler-core/src/jvm/select/mod.rs index fa9ad141..a1855919 100644 --- a/compiler-core/src/jvm/select/mod.rs +++ b/compiler-core/src/jvm/select/mod.rs @@ -1,4 +1,5 @@ //! Direct selection from compact SSA with interval-based JVM local reuse. +mod addresses; mod allocate; mod arrays; mod assemble; @@ -90,7 +91,6 @@ pub struct Options<'a> { /// verify; production selection otherwise consumes trusted internal IR. pub verify: bool, pub lines: Option<&'a SourceLines>, - pub relative_pointer_abi: bool, pub constants: Option<&'a dyn Constants>, pub debug: Option<&'a DebugInfo>, pub bootstrap: Option<&'a mut Vec>, @@ -112,7 +112,10 @@ struct Selector<'a> { handlers: Vec<(Label, BlockId)>, copies: Vec<(u16, u16, Kind)>, live: crate::opt::Live, + native_array_accesses: Vec>, storage: Vec, + aggregate_cache: Option<(ValueId, Option, TypeId)>, + aggregate_slot: Option, constants: Option<&'a dyn Constants>, bootstrap: Option<&'a mut Vec>, } @@ -135,7 +138,6 @@ pub fn compile_with_options( let Options { verify, lines, - relative_pointer_abi, constants, debug, bootstrap, @@ -154,15 +156,8 @@ pub fn compile_with_options( crate::opt::live_with_roots(body, types, debug.into_iter().flat_map(|d| d.roots(body))); let order = body.layout(); let forwarded = forward::values(body, &live, debug); - let allocation = allocate::allocate( - body, - types, - &live, - &forwarded, - relative_pointer_abi, - debug, - &order, - )?; + let allocation = allocate::allocate(body, types, &live, &forwarded, debug, &order)?; + let native_array_accesses = crate::analysis::native_array_accesses(body, types, &live); let mut debug = debug.map(|d| debug::Debug::new(d, body)); let mut s = Selector { body, @@ -180,7 +175,10 @@ pub fn compile_with_options( handlers: Vec::new(), copies: Vec::new(), live, + native_array_accesses, storage: Vec::new(), + aggregate_cache: None, + aggregate_slot: None, constants, bootstrap, }; @@ -207,12 +205,6 @@ pub fn compile_with_options( for ¶m in &body.blocks[body.entry.index()].params { let value = s.initial_value(param)?; frames::push_local_value(&mut initial, value); - if relative_pointer_abi - && matches!(types.get(body.value_type(param)), Some(Type::Pointer(_))) - { - frames::push_local_value(&mut initial, frames::FrameValue::Long); - frames::push_local_value(&mut initial, frames::FrameValue::Long); - } } if body.blocks.iter().any(|b| { matches!( @@ -233,16 +225,21 @@ pub fn compile_with_options( .checked_add(1) .ok_or_else(|| error("JVM local limit"))?; } - for block in order { + for (position, &block) in order.iter().enumerate() { + s.aggregate_cache = None; + // Keep loop headers after method entry. Frame-offset conversion reserves + // offset zero for the implicit entry frame. Removing an entry jump must + // not create a branch target at zero. + if s.assembly.code.is_empty() && body.edges.iter().any(|edge| edge.target == block) { + s.assembly.code.push(Instruction::Nop); + } s.assembly.bind(s.blocks[block.index()]); if block == body.entry { s.initialize_storage()?; if let Some(debug) = &mut debug { debug.allocate(&mut s)?; } - if relative_pointer_abi { - s.materialize_parameters()?; - } + // JVM byte/short parameters are sign-extended by Java callers. for ¶m in &body.blocks[block.index()].params { if s.live.uses[param.index()] > 1 @@ -272,6 +269,7 @@ pub fn compile_with_options( .is_some_and(|event| event.position as usize == position) { let event = events.next().unwrap(); + s.aggregate_cache = None; let start = s.assembly.code.len(); debug.as_mut().unwrap().apply(&mut s, event.change)?; if lines.is_some() { @@ -295,7 +293,10 @@ pub fn compile_with_options( } } let start = s.assembly.code.len(); - s.terminator(body.blocks[block.index()].terminator.unwrap())?; + s.terminator( + body.blocks[block.index()].terminator.unwrap(), + order.get(position + 1).copied(), + )?; debug_assert!(s.stack_value.is_none()); if let Some(debug) = &mut debug { debug.mark(start, s.assembly.code.len()); @@ -434,23 +435,35 @@ impl Selector<'_> { Ok(()) } fn jump(&mut self, edge: EdgeId) -> jvm::Result<()> { + self.jump_to(edge, None) + } + fn jump_to(&mut self, edge: EdgeId, fallthrough: Option) -> jvm::Result<()> { self.copies(edge)?; - self.assembly.branch( - Instruction::Goto_w(0), - self.blocks[self.body.edges[edge.index()].target.index()], - ); + let target = self.body.edges[edge.index()].target; + if Some(target) != fallthrough { + self.assembly + .branch(Instruction::Goto_w(0), self.blocks[target.index()]); + } Ok(()) } - fn terminator(&mut self, term: Terminator) -> jvm::Result<()> { + fn terminator(&mut self, term: Terminator, fallthrough: Option) -> jvm::Result<()> { match term { - Terminator::Jump(edge) => self.jump(edge)?, + Terminator::Jump(edge) => self.jump_to(edge, fallthrough)?, Terminator::Branch { condition, yes, no } => { - let yes_label = self.assembly.label(); + // Put the fallthrough edge last, after its parameter copies. + // The first edge must skip those copies even when both targets match. + let (branch, first, last) = + if Some(self.body.edges[no.index()].target) == fallthrough { + (Instruction::Ifeq(0), yes, no) + } else { + (Instruction::Ifne(0), no, yes) + }; + let last_label = self.assembly.label(); self.load(condition)?; - self.assembly.branch(Instruction::Ifne(0), yes_label); - self.jump(no)?; - self.assembly.bind(yes_label); - self.jump(yes)?; + self.assembly.branch(branch, last_label); + self.jump(first)?; + self.assembly.bind(last_label); + self.jump_to(last, fallthrough)?; } Terminator::Switch { value, @@ -474,7 +487,7 @@ impl Selector<'_> { self.instruction(inst)?; let end = u16::try_from(self.assembly.code.len())?; self.protect(start, end, unwind); - self.jump(normal)?; + self.jump_to(normal, fallthrough)?; } Terminator::Rethrow => { self.assembly.code.extend([ diff --git a/compiler-core/src/jvm/select/object_tests.rs b/compiler-core/src/jvm/select/object_tests.rs index 25e85acc..740edc5c 100644 --- a/compiler-core/src/jvm/select/object_tests.rs +++ b/compiler-core/src/jvm/select/object_tests.rs @@ -61,7 +61,6 @@ fn jvm_executes_constructors_fields_and_virtual_interface_calls() { name: "narrow".into(), ty: int, is_static: false, - relative_pointer: false, }); let old = b .emit( @@ -237,7 +236,6 @@ fn field_projection_verifies_both_ends_of_the_pointer_view() { name: "value".into(), ty: int, is_static: false, - relative_pointer: false, }); let projection = b.projection(PointerProjection { field, diff --git a/compiler-core/src/jvm/select/objects.rs b/compiler-core/src/jvm/select/objects.rs index a0114a7e..993dd137 100644 --- a/compiler-core/src/jvm/select/objects.rs +++ b/compiler-core/src/jvm/select/objects.rs @@ -1,11 +1,8 @@ -use super::super::abi::{ - relative_pointer_byte_offset_field, relative_pointer_element_offset_field, -}; use super::*; impl Selector<'_> { pub(super) fn object(&mut self, inst: Inst) -> jvm::Result { - if let Op::Cast(value) = inst.op { + if let Op::Cast(value) | Op::Refine(value) = inst.op { if self.value_kind(value)? == Kind::Reference && self.value_kind(inst.result.unwrap())? == Kind::Reference { @@ -63,38 +60,6 @@ impl Selector<'_> { _ => unreachable!(), }; self.assembly.code.push(op); - if field.relative_pointer { - let (object, store) = match inst.op { - Op::GetField { object, .. } => (object, false), - Op::SetField { object, .. } => (object, true), - _ => return Err(error("relative pointer fields require an instance")), - }; - for name in [ - relative_pointer_element_offset_field(&field.name), - relative_pointer_byte_offset_field(&field.name), - ] { - let offset = self.cp.add_field_ref(class, name, "J")?; - self.load(object)?; - if store { - self.assembly - .code - .extend([Instruction::Lconst_0, Instruction::Putfield(offset)]); - } else { - self.assembly.code.push(Instruction::Getfield(offset)); - } - } - if !store { - let owner = self.cp.add_class(POINTER_CLASS)?; - let materialize = self.cp.add_method_ref( - owner, - "materializeRelative", - "(Lorg/rustlang/runtime/Pointer;JJ)Lorg/rustlang/runtime/Pointer;", - )?; - self.assembly - .code - .push(Instruction::Invokestatic(materialize)); - } - } Ok(true) } } diff --git a/compiler-core/src/jvm/select/representation.rs b/compiler-core/src/jvm/select/representation.rs index 98f6106d..8e0e2945 100644 --- a/compiler-core/src/jvm/select/representation.rs +++ b/compiler-core/src/jvm/select/representation.rs @@ -1,6 +1,6 @@ use super::*; -pub use super::super::abi::{POINTER_CLASS, SLICE_VIEW_CLASS, UTF8_VIEW_CLASS}; +pub use super::super::abi::{POINTER_CLASS, SLICE_VIEW_CLASS, TAGGED_LONG_CLASS, UTF8_VIEW_CLASS}; pub(super) fn value_kind(types: &Types, ty: TypeId) -> jvm::Result { match types.get(ty) { @@ -11,7 +11,8 @@ pub(super) fn value_kind(types: &Types, ty: TypeId) -> jvm::Result { | Type::Interface(_) | Type::Array(_) | Type::Slice(_) - | Type::Str, + | Type::Str + | Type::TaggedI64, ) => Ok(Kind::Reference), _ => Err(error("unsupported JVM value representation")), } @@ -20,11 +21,12 @@ pub(super) fn value_kind(types: &Types, ty: TypeId) -> jvm::Result { pub(super) fn descriptor(types: &Types, ty: TypeId, output: &mut String) -> jvm::Result<()> { use ScalarType::*; match types.get(ty) { - Some(t @ (Type::Pointer(_) | Type::Slice(_) | Type::Str)) => { + Some(t @ (Type::Pointer(_) | Type::Slice(_) | Type::Str | Type::TaggedI64)) => { output.push('L'); output.push_str(match t { Type::Slice(_) => SLICE_VIEW_CLASS, Type::Str => UTF8_VIEW_CLASS, + Type::TaggedI64 => TAGGED_LONG_CLASS, _ => POINTER_CLASS, }); output.push(';'); diff --git a/compiler-core/src/jvm/select/scalar.rs b/compiler-core/src/jvm/select/scalar.rs index 22cf75bc..21591b45 100644 --- a/compiler-core/src/jvm/select/scalar.rs +++ b/compiler-core/src/jvm/select/scalar.rs @@ -3,6 +3,13 @@ use super::*; impl Selector<'_> { pub(super) fn instruction(&mut self, id: InstId) -> jvm::Result<()> { let inst = self.body.instructions[id.index()]; + if !matches!( + inst.op, + Op::LoadFieldPart { .. } | Op::LoadStorageField { index: Some(_), .. } + ) && inst.op.may_throw(self.body, self.types) + { + self.aggregate_cache = None; + } if inst.result.is_some_and(|v| literal(self.body, v).is_some()) { return Ok(()); } @@ -25,7 +32,8 @@ impl Selector<'_> { return Ok(()); } if self.general(inst)? - || self.array(inst)? + || self.array(id, inst)? + || self.address(inst)? || self.memory(inst)? || self.object(inst)? || self.view(inst)? @@ -149,9 +157,13 @@ impl Selector<'_> { self.assembly.code.push(I::L2i); } let (width, signed) = ty.integer().ok_or_else(|| error("non-integer shift"))?; - self.assembly - .code - .extend([get_int_const_instr(self.cp, (width - 1) as i32), I::Iand]); + // JVM int/long shifts already mask to 5/6 bits. Only narrower + // Rust integer widths require a stricter mask. + if width < 32 { + self.assembly + .code + .extend([get_int_const_instr(self.cp, (width - 1) as i32), I::Iand]); + } self.assembly.code.push(match (op, kind, signed) { (Shl, Kind::Int, _) => I::Ishl, (Shl, Kind::Long, _) => I::Lshl, diff --git a/compiler-core/src/jvm/select/tests.rs b/compiler-core/src/jvm/select/tests.rs index 5ed1d3d0..aec6ceca 100644 --- a/compiler-core/src/jvm/select/tests.rs +++ b/compiler-core/src/jvm/select/tests.rs @@ -25,6 +25,78 @@ fn binary_body(types: &Types, ty: TypeId, ret: TypeId, op: BinaryOp) -> Body { b.finish().unwrap() } +fn branch_arguments(types: &Types, int: TypeId, boolean: TypeId) -> Body { + let mut b = Builder::new(types, int); + let condition = b.parameter(b.current(), boolean); + let left = b.parameter(b.current(), int); + let right = b.parameter(b.current(), int); + let join = b.create_block(); + let value = b.parameter(join, int); + let yes = b.edge(join, vec![left]); + let no = b.edge(join, vec![right]); + b.terminate(Terminator::Branch { condition, yes, no }); + b.switch_to(join); + b.terminate(Terminator::Return(Some(value))); + b.finish().unwrap() +} + +fn entry_loop(types: &Types, int: TypeId, boolean: TypeId, array: TypeId) -> Body { + let mut b = Builder::new(types, int); + let count = b.parameter(b.current(), array); + let header = b.create_block(); + let update = b.create_block(); + let exit = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + let zero = integer(&mut b, int, ScalarType::I32, 0); + let n = b + .emit( + Op::ArrayGet { + array: count, + index: zero, + native: true, + }, + Some(int), + ) + .unwrap(); + let condition = b + .emit( + Op::Binary { + op: BinaryOp::Gt, + left: n, + right: zero, + }, + Some(boolean), + ) + .unwrap(); + b.branch(condition, update, exit); + b.switch_to(update); + let one = integer(&mut b, int, ScalarType::I32, 1); + let next = b + .emit( + Op::Binary { + op: BinaryOp::Sub, + left: n, + right: one, + }, + Some(int), + ) + .unwrap(); + b.emit( + Op::ArraySet { + array: count, + index: zero, + value: next, + native: true, + }, + None, + ); + b.jump(header, vec![]); + b.switch_to(exit); + b.terminate(Terminator::Return(Some(n))); + b.finish().unwrap() +} + fn swap_loop( types: &Types, ty: TypeId, @@ -267,6 +339,7 @@ fn jvm_verifies_and_executes_ssa_loops_parallel_copies_scalars_and_throw_points( let byte = types.scalar(ScalarType::U8); let short = types.scalar(ScalarType::I16); let signed_byte = types.scalar(ScalarType::I8); + let int_array = types.intern(Type::Array(int)); use crate::scalar::BitOp; let methods = [ ( @@ -356,6 +429,16 @@ fn jvm_verifies_and_executes_ssa_loops_parallel_copies_scalars_and_throw_points( overflow_body(&types, ulong, boolean, BinaryOp::Mul), ), ("callHandler", "(II)I", call_with_handler(&types, int)), + ( + "entryLoop", + "([I)I", + entry_loop(&types, int, boolean, int_array), + ), + ( + "branchArguments", + "(ZII)I", + branch_arguments(&types, int, boolean), + ), ( "tableSwitch", "(I)I", @@ -402,6 +485,21 @@ fn jvm_verifies_and_executes_ssa_loops_parallel_copies_scalars_and_throw_points( "(SS)S", binary_body(&types, short, short, BinaryOp::Shr), ), + ( + "shiftInt", + "(II)I", + binary_body(&types, int, int, BinaryOp::Shl), + ), + ( + "shiftLong", + "(JJ)J", + binary_body(&types, long, long, BinaryOp::Shr), + ), + ( + "shiftUnsignedLong", + "(JJ)J", + binary_body(&types, ulong, ulong, BinaryOp::Shr), + ), ( "addByte", "(BB)B", @@ -455,6 +553,21 @@ fn jvm_verifies_and_executes_ssa_loops_parallel_copies_scalars_and_throw_points( .into_iter() .map(|(name, descriptor, body)| { let code = compile(&body, &types, &mut cp).unwrap(); + if name == "entryLoop" { + assert!( + code.instructions + .iter() + .any(|inst| matches!(inst, Instruction::Goto_w(1))) + ); + } + if matches!(name, "swap" | "swapWide" | "branchArguments") { + assert!(!code.instructions.iter().enumerate().any(|(i, inst)| { + matches!(inst, Instruction::Goto_w(target) if *target as usize == i + 1) + })); + } + if matches!(name, "shiftInt" | "shiftLong" | "shiftUnsignedLong") { + assert!(!code.instructions.contains(&Instruction::Iand)); + } Method { access_flags: MethodAccessFlags::PUBLIC | MethodAccessFlags::STATIC, name_index: cp.add_utf8(name).unwrap(), @@ -520,6 +633,9 @@ public class Run { } long[] boundaries = {Long.MIN_VALUE, Long.MIN_VALUE+1, -3037000500L, -1, 0, 1, 3037000500L, Long.MAX_VALUE-1, Long.MAX_VALUE}; for (long a : boundaries) for (long b : boundaries) { + check(SsaFixture.shiftInt((int)a, (int)b) == ((int)a << (int)b)); + check(SsaFixture.shiftLong(a, b) == (a >> b)); + check(SsaFixture.shiftUnsignedLong(a, b) == (a >>> b)); boolean add=false, sub=false, mul=false; try { Math.addExact(a,b); } catch (ArithmeticException ex) { add=true; } try { Math.subtractExact(a,b); } catch (ArithmeticException ex) { sub=true; } @@ -537,6 +653,11 @@ public class Run { } check(SsaFixture.callHandler(-7, 2)==-4); check(SsaFixture.callHandler(123, 0)==123); + check(SsaFixture.branchArguments(true, 11, 29)==11); + check(SsaFixture.branchArguments(false, 11, 29)==29); + check(SsaFixture.entryLoop(new int[]{100})==0); + check(SsaFixture.entryLoop(new int[]{0})==0); + check(SsaFixture.entryLoop(new int[]{-1})==-1); for (int n=-100; n<100; n++) check(SsaFixture.tableSwitch(n)==(n>=-2 && n<=2 ? n+3 : -1)); check(SsaFixture.lookupSwitch(Integer.MIN_VALUE)==1); check(SsaFixture.lookupSwitch(0)==2); check(SsaFixture.lookupSwitch(Integer.MAX_VALUE)==3); @@ -582,6 +703,29 @@ public class Run { fs::remove_dir_all(directory).unwrap(); } +#[test] +fn empty_blocks_fall_through_without_emitting_branches() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let mut b = Builder::new(&types, int); + let value = integer(&mut b, int, ScalarType::I32, 7); + for _ in 0..4 { + let block = b.create_block(); + b.jump(block, vec![]); + b.switch_to(block); + } + b.terminate(Terminator::Return(Some(value))); + let code = compile(&b.finish().unwrap(), &types, &mut Default::default()).unwrap(); + assert_eq!( + code.instructions, + vec![ + Instruction::Nop, + Instruction::Bipush(7), + Instruction::Ireturn + ] + ); +} + #[test] fn source_line_tables_follow_selected_instructions_and_terminators() { let mut types = Types::default(); @@ -697,9 +841,11 @@ fn representation_constants_and_opaque_values_retain_ordered_effects() { let int = types.scalar(ScalarType::I32); let unit = types.intern(Type::Unit); let mut b = Builder::new(&types, unit); - b.body - .constants - .push(Constant::External { index: 7, ty: int }); + b.body.constants.push(Constant::External { + index: 7, + ty: int, + pure: false, + }); let value = b.emit(Op::Constant(ConstId::new(0)), Some(int)).unwrap(); b.emit(Op::Opaque(value), Some(int)); b.terminate(Terminator::Return(None)); diff --git a/compiler-core/src/jvm/select/unwind_tests.rs b/compiler-core/src/jvm/select/unwind_tests.rs index 9dc2a470..856c9b9b 100644 --- a/compiler-core/src/jvm/select/unwind_tests.rs +++ b/compiler-core/src/jvm/select/unwind_tests.rs @@ -13,6 +13,28 @@ fn jvm_preserves_throwable_identity_through_catch_and_rethrow() { let this_class = cp.add_class("UnwindSsa").unwrap(); let super_class = cp.add_class("java/lang/Object").unwrap(); let mut methods = Vec::new(); + // The JVM requires a terminal instruction on Rust-UB paths, + // including methods that return a value. + let mut b = Builder::new(&types, int); + b.terminate(Terminator::Unreachable); + let code = compile(&b.finish().unwrap(), &types, &mut cp).unwrap(); + assert_eq!( + code.instructions, + vec![Instruction::Aconst_null, Instruction::Athrow] + ); + methods.push(Method { + access_flags: MethodAccessFlags::PUBLIC | MethodAccessFlags::STATIC, + name_index: cp.add_utf8("unreachable").unwrap(), + descriptor_index: cp.add_utf8("()I").unwrap(), + attributes: vec![Attribute::Code { + name_index: cp.add_utf8("Code").unwrap(), + max_stack: code.max_stack, + max_locals: code.max_locals, + code: code.instructions, + exception_table: code.exceptions, + attributes: code.attributes, + }], + }); for rethrow in [false, true] { let mut b = Builder::new(&types, throwable); let original = b.parameter(b.current(), throwable); @@ -81,6 +103,8 @@ fn jvm_preserves_throwable_identity_through_catch_and_rethrow() { r#" public class UnwindRun { public static void main(String[] args) throws Throwable { + try { UnwindSsa.unreachable(); throw new AssertionError("returned"); } + catch (NullPointerException expected) { } Throwable original = new IllegalArgumentException("identity"); if (UnwindSsa.capture(original) != original) throw new AssertionError("catch"); try { UnwindSsa.rethrow(original); throw new AssertionError("returned"); } diff --git a/compiler-core/src/jvm/select/views.rs b/compiler-core/src/jvm/select/views.rs index 4c29ed3e..1116a2aa 100644 --- a/compiler-core/src/jvm/select/views.rs +++ b/compiler-core/src/jvm/select/views.rs @@ -5,6 +5,20 @@ use representation::{SLICE_VIEW_CLASS, UTF8_VIEW_CLASS}; impl Selector<'_> { pub(super) fn view(&mut self, inst: Inst) -> jvm::Result { match inst.op { + Op::TaggedPack(parts) => { + self.materialize_tagged(parts)?; + } + Op::TaggedPart { value, index } => { + self.load(value)?; + let owner = self.cp.add_class(super::super::abi::TAGGED_LONG_CLASS)?; + self.assembly + .code + .push(Instruction::Invokestatic(self.cp.add_method_ref( + owner, + if index == 0 { "value" } else { "tag" }, + "(Lorg/rustlang/runtime/TaggedLong;)J", + )?)); + } Op::Length(view) => { self.load(view)?; let owner = self.cp.add_class(SLICE_VIEW_CLASS)?; @@ -30,6 +44,16 @@ impl Selector<'_> { self.load(length)?; self.assembly.code.push(Instruction::Invokespecial(init)); } + Op::ViewPart { view, index } => { + self.load(view)?; + self.assembly.code.push(super::super::abi::view_part_access( + self.cp, + index as usize, + )?); + } + Op::ViewPack(parts) => { + self.materialize_view(self.body.value_type(inst.result.unwrap()), parts)?; + } Op::ViewData { view, size, codec } => { self.load(view)?; self.assembly @@ -50,8 +74,93 @@ impl Selector<'_> { )?; self.assembly.code.push(Instruction::Invokestatic(method)); } + Op::ViewAddress { parts, size, codec } => { + for &part in &self.body.args[parts.range()] { + self.load(part)?; + } + self.assembly + .code + .push(get_int_const_instr(self.cp, size as i32)); + self.assembly.code.push(match codec { + Some(id) => Instruction::Ldc_w( + self.cp + .add_name_string(self.types.symbol_name(id).unwrap())?, + ), + None => Instruction::Aconst_null, + }); + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + "fromSliceParts", + "(Ljava/lang/Object;IJILjava/lang/String;)Lorg/rustlang/runtime/Pointer;", + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + } + Op::ViewRoot { + backing, + size, + codec, + } => { + self.load(backing)?; + if codec.is_some() { + self.address_layout(size, codec)?; + } else { + self.assembly + .code + .push(get_int_const_instr(self.cp, size as i32)); + } + let owner = self.cp.add_class(POINTER_CLASS)?; + let method = self.cp.add_method_ref( + owner, + if codec.is_some() { + "sliceAddressRoot" + } else { + "scalarSliceRoot" + }, + if codec.is_some() { + "(Ljava/lang/Object;ILjava/lang/String;)Ljava/lang/Object;" + } else { + "(Ljava/lang/Object;I)Ljava/lang/Object;" + }, + )?; + self.assembly.code.push(Instruction::Invokestatic(method)); + } _ => return Ok(false), } Ok(true) } + pub(super) fn materialize_tagged(&mut self, parts: List) -> jvm::Result<()> { + let owner = self.cp.add_class(super::super::abi::TAGGED_LONG_CLASS)?; + self.assembly + .code + .extend([Instruction::New(owner), Instruction::Dup]); + for &value in &self.body.args[parts.range()] { + self.load(value)?; + } + self.assembly.code.push(Instruction::Invokespecial( + self.cp.add_method_ref(owner, "", "(JJ)V")?, + )); + Ok(()) + } + pub(super) fn materialize_view(&mut self, ty: TypeId, parts: List) -> jvm::Result<()> { + let owner = self + .cp + .add_class(if matches!(self.types.get(ty), Some(Type::Str)) + || matches!(self.types.get(ty), Some(Type::Class(s)) if self.types.symbol_name(s) == Some(UTF8_VIEW_CLASS)) { + UTF8_VIEW_CLASS + } else { + SLICE_VIEW_CLASS + })?; + let init = self + .cp + .add_method_ref(owner, "", "(Ljava/lang/Object;IJ)V")?; + self.assembly + .code + .extend([Instruction::New(owner), Instruction::Dup]); + for &value in &self.body.args[parts.range()] { + self.load(value)?; + } + self.assembly.code.push(Instruction::Invokespecial(init)); + Ok(()) + } } diff --git a/compiler-core/src/lib.rs b/compiler-core/src/lib.rs index f3b70949..d6905789 100644 --- a/compiler-core/src/lib.rs +++ b/compiler-core/src/lib.rs @@ -1,4 +1,5 @@ //! Rustc-independent IR and JVM compiler machinery. +pub mod analysis; pub mod ir; pub mod scalar; diff --git a/compiler-core/src/opt/fields_tests.rs b/compiler-core/src/opt/fields_tests.rs index f7e85ba4..6c832b3f 100644 --- a/compiler-core/src/opt/fields_tests.rs +++ b/compiler-core/src/opt/fields_tests.rs @@ -23,7 +23,6 @@ fn project( name: "value".into(), ty: scalar, is_static: false, - relative_pointer: false, }); let projection = b.projection(PointerProjection { field, From 125beaca445bfc2ada7f790c6dfe78f67316a58a Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 00:12:52 +1000 Subject: [PATCH 15/61] simplify components and unreachable edges --- compiler-core/src/opt/live.rs | 5 +- compiler-core/src/opt/mod.rs | 6 + compiler-core/src/opt/simplify.rs | 239 +++++++++++++++++++++++++ compiler-core/src/opt/tagged.rs | 205 +++++++++++++++++++++ compiler-core/src/opt/unreachable.rs | 257 +++++++++++++++++++++++++++ 5 files changed, 711 insertions(+), 1 deletion(-) create mode 100644 compiler-core/src/opt/simplify.rs create mode 100644 compiler-core/src/opt/tagged.rs create mode 100644 compiler-core/src/opt/unreachable.rs diff --git a/compiler-core/src/opt/live.rs b/compiler-core/src/opt/live.rs index 7fa30186..4ca8c30b 100644 --- a/compiler-core/src/opt/live.rs +++ b/compiler-core/src/opt/live.rs @@ -85,7 +85,10 @@ pub fn live_with_roots( fn has_effects(op: Op, body: &Body, types: &Types) -> bool { // An unused fat-pointer carrier has no observable identity. - !matches!(op, Op::View { .. }) && op.may_throw(body, types) + !matches!( + op, + Op::View { .. } | Op::ViewPack(_) | Op::AddressPack(_) | Op::SlotRoot(_) + ) && op.may_throw(body, types) } #[cfg(test)] diff --git a/compiler-core/src/opt/mod.rs b/compiler-core/src/opt/mod.rs index 6621eb4a..1e67a044 100644 --- a/compiler-core/src/opt/mod.rs +++ b/compiler-core/src/opt/mod.rs @@ -8,6 +8,12 @@ pub use live::{Live, live, live_with_roots}; mod cells; pub use cells::promote_cells; +mod tagged; +pub use tagged::decompose_tagged; #[cfg(test)] mod cells_tests; + +mod simplify; +mod unreachable; +pub use simplify::simplify_components; diff --git a/compiler-core/src/opt/simplify.rs b/compiler-core/src/opt/simplify.rs new file mode 100644 index 00000000..562236c6 --- /dev/null +++ b/compiler-core/src/opt/simplify.rs @@ -0,0 +1,239 @@ +//! Fold arithmetic and aliases after representation lowering. +//! Each definition becomes an identity or constant at most twice. +//! Revisit only its consumers. Preserve source positions and instruction IDs. +use crate::ir::*; + +pub fn simplify_components(body: &mut Body, types: &Types) { + let mut users = crate::analysis::ValueUsers::new(body.values.len()); + let mut pending = Vec::new(); + let mut queued = vec![false; body.instructions.len()]; + for (index, inst) in body.instructions.iter().enumerate() { + if inst.result.is_some() + && matches!( + inst.op, + Op::Binary { .. } + | Op::Cast(_) + | Op::Reinterpret(_) + | Op::Neg(_) + | Op::Not(_) + | Op::Bit { .. } + | Op::Overflow { .. } + | Op::Length(_) + ) + { + inst.op + .visit_uses(&body.args, |v| users.connect(body.resolve(v), index)); + pending.push(index); + queued[index] = true; + } + } + // Visit definitions in order to avoid repeated work for straight-line code. + // The worklist also handles earlier users of appended components. + pending.reverse(); + while let Some(index) = pending.pop() { + queued[index] = false; + let inst = body.instructions[index]; + let Some(result) = inst.result else { + continue; + }; + let folded = match inst.op { + Op::Reinterpret(value) if body.value_type(value) == body.value_type(result) => Some( + body.scalar_value(value) + .map_or(Folded::Value(value), Folded::Constant), + ), + _ => body.fold(types, inst.op, body.value_type(result)), + }; + let Some(folded) = folded else { + continue; + }; + match folded { + Folded::Value(source) => { + let source = body.resolve(source); + if source == result { + continue; + } + let replacement = if let Some(value) = body.scalar_value(source) { + let id = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar(value)); + Op::Constant(id) + } else { + Op::Reinterpret(source) + }; + if inst.op == replacement { + continue; + } + // Retain the result until constant propagation ends. + // An appended source can become constant later and must reach these consumers. + body.instructions[index].op = replacement; + } + Folded::Constant(value) => { + let id = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar(value)); + body.instructions[index].op = Op::Constant(id); + } + } + for user in users.users(result.index()) { + if !std::mem::replace(&mut queued[user], true) { + pending.push(user); + } + } + } + // Remove exact identities only after the constant worklist has settled. + for index in 0..body.instructions.len() { + let inst = body.instructions[index]; + if let (Some(result), Op::Reinterpret(source)) = (inst.result, inst.op) + && body.value_type(source) == body.value_type(result) + { + let source = body.resolve_mut(source); + body.values[result.index()].def = ValueDef::Alias(source); + body.instructions[index] = Inst { + op: Op::Nop, + result: None, + }; + } + } + for index in 0..body.values.len() { + body.resolve_mut(ValueId::new(index)); + } + for index in 0..body.blocks.len() { + let terminator = body.blocks[index].terminator.unwrap(); + let replacement = match terminator { + Terminator::Invoke { inst, normal, .. } + if matches!(body.instructions[inst.index()].op, Op::Nop) + || matches!(body.instructions[inst.index()].op, Op::Constant(id) if matches!(body.constants[id.index()], Constant::Scalar(_))) => + { + body.blocks[index].instructions.push(inst); + Some(Terminator::Jump(normal)) + } + Terminator::Branch { condition, yes, no } => body + .scalar_value(condition) + .map(|value| Terminator::Jump(if value.bits() != 0 { yes } else { no })), + Terminator::Switch { + value, + cases, + otherwise, + } => body.scalar_value(value).map(|value| { + let edge = body.cases[cases.range()] + .iter() + .find_map(|&(key, edge)| (key == value).then_some(edge)) + .unwrap_or(otherwise); + Terminator::Jump(edge) + }), + _ => None, + }; + let replacement = replacement.or_else(|| super::unreachable::simplify(body, terminator)); + if let Some(replacement) = replacement { + body.blocks[index].terminator = Some(replacement); + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::scalar::{BinaryOp, Scalar, ScalarType}; + + #[test] + fn constants_reach_earlier_users_through_late_signedness_annotations() { + let mut types = Types::default(); + let long = types.scalar(ScalarType::I64); + let unsigned = types.scalar(ScalarType::U64); + let mut b = Builder::new(&types, long); + let x = b.parameter(b.current(), long); + let first = b.emit(Op::Reinterpret(x), Some(long)).unwrap(); + let result = b.emit(Op::Neg(first), Some(long)).unwrap(); + let late = b.emit(Op::Reinterpret(x), Some(unsigned)).unwrap(); + let a = b.constant(unsigned, Scalar::integer(ScalarType::U64, 8).unwrap()); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + let ValueDef::Inst(first_inst) = body.values[first.index()].def else { + panic!() + }; + let ValueDef::Inst(late_inst) = body.values[late.index()].def else { + panic!() + }; + body.instructions[first_inst.index()].op = Op::Reinterpret(late); + body.instructions[late_inst.index()].op = Op::Binary { + op: BinaryOp::Mul, + left: a, + right: a, + }; + // Definitions have newer IDs but occur before their users in the block. + body.blocks[body.entry.index()] + .instructions + .sort_by_key(|id| { + if *id == late_inst { + 0 + } else if *id == first_inst { + 1 + } else { + 2 + } + }); + simplify_components(&mut body, &types); + assert_eq!(body.scalar_value(result).unwrap().signed(), Some(-64)); + verify(&body, &types).unwrap(); + } + + // Representation passes append instructions after the builder's fold point. + // Propagate facts from these newer IDs to older users. + #[test] + fn folds_appended_components_without_dropping_traps_or_float_work() { + let mut types = Types::default(); + let long = types.scalar(ScalarType::I64); + let float = types.scalar(ScalarType::F64); + let mut b = Builder::new(&types, long); + let x = b.parameter(b.current(), long); + let f = b.parameter(b.current(), float); + let first = b.emit(Op::Reinterpret(x), Some(long)).unwrap(); + let second = b.emit(Op::Reinterpret(first), Some(long)).unwrap(); + let zero = b.constant(long, Scalar::integer(ScalarType::I64, 0).unwrap()); + let eight = b.constant(long, Scalar::integer(ScalarType::I64, 8).unwrap()); + let trap = b + .emit( + Op::Binary { + op: BinaryOp::Div, + left: zero, + right: x, + }, + Some(long), + ) + .unwrap(); + let float_zero = b.constant(float, Scalar::f64(0.0)); + let float_sum = b + .emit( + Op::Binary { + op: BinaryOp::Add, + left: f, + right: float_zero, + }, + Some(float), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(second))); + let mut body = b.finish().unwrap(); + let ValueDef::Inst(id) = body.values[first.index()].def else { + panic!() + }; + body.instructions[id.index()].op = Op::Binary { + op: BinaryOp::Mul, + left: eight, + right: eight, + }; + simplify_components(&mut body, &types); + assert_eq!(body.scalar_value(second).unwrap().bits(), 64); + assert!(body.scalar_value(trap).is_none()); + assert!(body.scalar_value(float_sum).is_none()); + verify(&body, &types).unwrap(); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!( + code.instructions + .contains(&crate::classfile::attributes::Instruction::Ldiv) + ); + assert!( + !code + .instructions + .contains(&crate::classfile::attributes::Instruction::Lmul) + ); + } +} diff --git a/compiler-core/src/opt/tagged.rs b/compiler-core/src/opt/tagged.rs new file mode 100644 index 00000000..8568bc1f --- /dev/null +++ b/compiler-core/src/opt/tagged.rs @@ -0,0 +1,205 @@ +//! Keep tagged scalar products decomposed through joins and loop backedges. +use crate::ir::*; + +pub fn decompose_tagged(body: &mut Body, types: &Types, debug: Option<&mut DebugInfo>) { + if !body + .instructions + .iter() + .any(|i| matches!(i.op, Op::TaggedPack(_) | Op::TaggedPart { .. })) + { + return; + } + let count = body.values.len(); + let mut eligible = vec![false; count]; + let mut users = crate::analysis::ValueUsers::new(count); + let predecessors = body.predecessors(); + for (index, value) in body.values.iter().enumerate() { + if types.get(value.ty) != Some(Type::TaggedI64) { + continue; + } + match value.def { + ValueDef::Inst(id) => match body.instructions[id.index()].op { + Op::TaggedPack(_) => eligible[index] = true, + Op::Reinterpret(source) | Op::Adapt(source) | Op::Refine(source) => { + eligible[index] = true; + users.connect(source, index); + } + _ => {} + }, + ValueDef::Alias(source) => { + eligible[index] = true; + users.connect(source, index); + } + ValueDef::Param(block) if block != body.entry => { + eligible[index] = !predecessors[block.index()].is_empty(); + let position = body.blocks[block.index()] + .params + .iter() + .position(|v| v.index() == index) + .unwrap(); + for &(_, edge) in &predecessors[block.index()] { + users.connect(body.edges[edge.index()].args[position], index); + } + } + _ => {} + } + } + users.close(&mut eligible); + let long = types + .find(Type::Scalar(crate::scalar::ScalarType::I64)) + .unwrap(); + let mut components = vec![None::<[ValueId; 2]>; count]; + let mut joins = vec![Vec::new(); body.blocks.len()]; + for index in 0..count { + if !eligible[index] { + continue; + } + match body.values[index].def { + ValueDef::Inst(id) => { + if let Op::TaggedPack(parts) = body.instructions[id.index()].op { + components[index] = Some(body.args[parts.range()].try_into().unwrap()); + } + } + ValueDef::Param(block) => { + let parts = std::array::from_fn(|_| { + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty: long, + def: ValueDef::Param(block), + }); + value + }); + let position = body.blocks[block.index()] + .params + .iter() + .position(|v| v.index() == index) + .unwrap(); + joins[block.index()].push((index, position)); + components[index] = Some(parts); + } + _ => {} + } + } + users.propagate(&eligible, &mut components); + for (block, joins) in joins.iter_mut().enumerate() { + joins.sort_unstable_by_key(|&(_, position)| position); + let mut prologue = Vec::new(); + for &(index, position) in joins.iter().rev() { + let parts = components[index].unwrap(); + body.blocks[block] + .params + .splice(position..position + 1, parts); + for &(_, edge) in &predecessors[block] { + let input = body.edges[edge.index()].args[position]; + body.edges[edge.index()] + .args + .splice(position..position + 1, components[input.index()].unwrap()); + } + let inst = InstId::new(body.instructions.len()); + body.values[index].def = ValueDef::Inst(inst); + body.instructions.push(Inst { + op: Op::TaggedPack(List::append(&mut body.args, parts)), + result: Some(ValueId::new(index)), + }); + prologue.push(inst); + } + prologue.append(&mut body.blocks[block].instructions); + body.blocks[block].instructions = prologue; + } + for inst in &mut body.instructions { + if let Op::TaggedPart { value, index } = inst.op { + if let Some(Some(parts)) = components.get(value.index()) { + inst.op = Op::Reinterpret(parts[index as usize]); + } + } + } + if let Some(debug) = debug { + for event in &mut debug.events { + event.position += joins[event.block.index()].len() as u32; + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::scalar::{BinaryOp, Scalar, ScalarType}; + + #[test] + fn loop_joins_retain_scalars_without_a_boundary_carrier() { + let mut types = Types::default(); + let long = types.scalar(ScalarType::I64); + let boolean = types.scalar(ScalarType::Bool); + let tagged = types.intern(Type::TaggedI64); + let mut b = Builder::new(&types, long); + let limit = b.parameter(b.current(), long); + let zero = b.constant(long, Scalar::integer(ScalarType::I64, 0).unwrap()); + let one = b.constant(long, Scalar::integer(ScalarType::I64, 1).unwrap()); + let parts = b.args([zero, one]); + let first = b.emit(Op::TaggedPack(parts), Some(tagged)).unwrap(); + let value = b.variable(tagged); + b.define(value, first); + let header = b.create_block(); + let step = b.create_block(); + let done = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + let pair = b.read(value); + let payload = b + .emit( + Op::TaggedPart { + value: pair, + index: 0, + }, + Some(long), + ) + .unwrap(); + let more = b + .emit( + Op::Binary { + op: BinaryOp::Lt, + left: payload, + right: limit, + }, + Some(boolean), + ) + .unwrap(); + b.branch(more, step, done); + b.switch_to(step); + let next = b + .emit( + Op::Binary { + op: BinaryOp::Add, + left: payload, + right: one, + }, + Some(long), + ) + .unwrap(); + let parts = b.args([next, one]); + let pair = b.emit(Op::TaggedPack(parts), Some(tagged)).unwrap(); + b.define(value, pair); + b.jump(header, vec![]); + b.switch_to(done); + b.terminate(Terminator::Return(Some(payload))); + let mut body = b.finish().unwrap(); + decompose_tagged(&mut body, &types, None); + verify(&body, &types).unwrap(); + let live = crate::opt::live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(index, inst)| live.instructions[index] + && matches!(inst.op, Op::TaggedPack(_) | Op::TaggedPart { .. })) + ); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!( + !code + .instructions + .iter() + .any(|i| matches!(i, crate::classfile::attributes::Instruction::New(_))) + ); + } +} diff --git a/compiler-core/src/opt/unreachable.rs b/compiler-core/src/opt/unreachable.rs new file mode 100644 index 00000000..af8ae65d --- /dev/null +++ b/compiler-core/src/opt/unreachable.rs @@ -0,0 +1,257 @@ +//! Bypass empty traps on undefined paths. +//! Keep preceding instructions because they can throw or unwind. +use crate::ir::*; + +fn impossible(body: &Body, edge: EdgeId) -> bool { + let block = &body.blocks[body.edges[edge.index()].target.index()]; + block.terminator == Some(Terminator::Unreachable) + && block + .instructions + .iter() + .all(|i| body.instructions[i.index()].op == Op::Nop) +} + +pub(super) fn simplify(body: &mut Body, term: Terminator) -> Option { + match term { + Terminator::Branch { yes, no, .. } => match (impossible(body, yes), impossible(body, no)) { + (true, true) => Some(Terminator::Unreachable), + (true, false) => Some(Terminator::Jump(no)), + (false, true) => Some(Terminator::Jump(yes)), + _ => None, + }, + Terminator::Switch { + value, + cases, + otherwise, + } => { + let targets = &body.cases[cases.range()]; + let default = if impossible(body, otherwise) { + let Some((_, edge)) = targets.iter().find(|(_, edge)| !impossible(body, *edge)) + else { + return Some(Terminator::Unreachable); + }; + *edge + } else { + otherwise + }; + if default == otherwise + && !targets + .iter() + .any(|&(_, edge)| edge == default || impossible(body, edge)) + { + return None; + } + let retained = targets + .iter() + .copied() + .filter(|&(_, edge)| edge != default && !impossible(body, edge)) + .collect::>(); + if retained.is_empty() { + return Some(Terminator::Jump(default)); + } + let cases = List { + start: body.cases.len().try_into().expect("switch pool capacity"), + len: retained.len().try_into().expect("switch capacity"), + }; + body.cases.extend(retained); + Some(Terminator::Switch { + value, + cases, + otherwise: default, + }) + } + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::scalar::{BinaryOp, Scalar, ScalarType}; + + #[test] + fn impossible_edges_disappear_without_losing_the_returned_value() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let bool_ = types.scalar(ScalarType::Bool); + for switch in [false, true] { + let mut b = Builder::new(&types, int); + let condition = b.parameter(b.current(), if switch { int } else { bool_ }); + let value = b.parameter(b.current(), int); + let valid = b.create_block(); + let result = b.parameter(valid, int); + let invalid = b.create_block(); + let yes = b.edge(valid, vec![value]); + let no = b.edge(invalid, vec![]); + let cases = List { start: 0, len: 1 }; + b.body + .cases + .push((Scalar::integer(ScalarType::I32, 17).unwrap(), yes)); + b.terminate(if switch { + Terminator::Switch { + value: condition, + cases, + otherwise: no, + } + } else { + Terminator::Branch { condition, yes, no } + }); + b.switch_to(valid); + b.terminate(Terminator::Return(Some(result))); + b.switch_to(invalid); + b.terminate(Terminator::Unreachable); + let mut body = b.finish().unwrap(); + super::super::simplify_components(&mut body, &types); + assert_eq!( + body.blocks[body.entry.index()].terminator, + Some(Terminator::Jump(yes)) + ); + assert_eq!(body.resolve(result), body.resolve(value)); + assert!(!body.reachable()[invalid.index()]); + verify(&body, &types).unwrap(); + } + } + + #[test] + fn work_before_unreachable_may_diverge_so_the_edge_must_remain() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let bool_ = types.scalar(ScalarType::Bool); + let mut b = Builder::new(&types, int); + let condition = b.parameter(b.current(), bool_); + let value = b.parameter(b.current(), int); + let divisor = b.parameter(b.current(), int); + let valid = b.create_block(); + let invalid = b.create_block(); + b.branch(condition, valid, invalid); + b.switch_to(valid); + b.terminate(Terminator::Return(Some(value))); + b.switch_to(invalid); + b.emit( + Op::Binary { + op: BinaryOp::Div, + left: value, + right: divisor, + }, + Some(int), + ); + b.terminate(Terminator::Unreachable); + let mut body = b.finish().unwrap(); + super::super::simplify_components(&mut body, &types); + assert!(matches!( + body.blocks[body.entry.index()].terminator, + Some(Terminator::Branch { .. }) + )); + assert!(body.reachable()[invalid.index()]); + verify(&body, &types).unwrap(); + } + #[test] + fn simplified_switch_preserves_distinct_arguments_on_the_jvm() { + use crate::classfile::{ + ClassAccessFlags, ClassFile, Method, MethodAccessFlags, Version, attributes::Attribute, + constant_pool::InternedConstantPool, + }; + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let mut b = Builder::new(&types, int); + let condition = b.parameter(b.current(), int); + let first = b.parameter(b.current(), int); + let second = b.parameter(b.current(), int); + let join = b.create_block(); + let result = b.parameter(join, int); + let invalid = b.create_block(); + let a = b.edge(join, vec![first]); + let c = b.edge(join, vec![second]); + let ub = b.edge(invalid, vec![]); + let invalid_case = b.edge(invalid, vec![]); + for (key, edge) in [(17, a), (23, c), (99, invalid_case)] { + b.body + .cases + .push((Scalar::integer(ScalarType::I32, key).unwrap(), edge)); + } + b.terminate(Terminator::Switch { + value: condition, + cases: List { start: 0, len: 3 }, + otherwise: ub, + }); + b.switch_to(join); + b.terminate(Terminator::Return(Some(result))); + b.switch_to(invalid); + b.terminate(Terminator::Unreachable); + let mut body = b.finish().unwrap(); + super::super::simplify_components(&mut body, &types); + let Some(Terminator::Switch { + cases, otherwise, .. + }) = body.blocks[body.entry.index()].terminator + else { + panic!("lost switch"); + }; + assert_eq!(cases.len, 1); + assert_eq!(otherwise, a); + assert_eq!(body.edges[a.index()].args, vec![first]); + assert_eq!(body.edges[c.index()].args, vec![second]); + assert!(!body.reachable()[invalid.index()]); + verify(&body, &types).unwrap(); + let mut cp = InternedConstantPool::default(); + let this_class = cp.add_class("SwitchUB").unwrap(); + let super_class = cp.add_class("java/lang/Object").unwrap(); + let code = crate::jvm::select::compile(&body, &types, &mut cp).unwrap(); + let method = Method { + access_flags: MethodAccessFlags::PUBLIC | MethodAccessFlags::STATIC, + name_index: cp.add_utf8("eval").unwrap(), + descriptor_index: cp.add_utf8("(III)I").unwrap(), + attributes: vec![Attribute::Code { + name_index: cp.add_utf8("Code").unwrap(), + max_stack: code.max_stack, + max_locals: code.max_locals, + code: code.instructions, + exception_table: code.exceptions, + attributes: code.attributes, + }], + }; + let class = ClassFile { + version: Version::Java8 { minor: 0 }, + constant_pool: cp.into_inner(), + access_flags: ClassAccessFlags::PUBLIC | ClassAccessFlags::SUPER, + this_class, + super_class, + methods: vec![method], + ..Default::default() + }; + let mut bytes = Vec::new(); + class.to_bytes(&mut bytes).unwrap(); + let dir = + std::env::temp_dir().join(format!("rcj-unreachable-switch-{}", std::process::id())); + std::fs::create_dir_all(&dir).unwrap(); + std::fs::write(dir.join("SwitchUB.class"), bytes).unwrap(); + std::fs::write( + dir.join("Run.java"), + r#" +public class Run { + public static void main(String[] args) { + for (int i = -100; i <= 100; ++i) { + if (SwitchUB.eval(17, i, i + 1) != i || SwitchUB.eval(23, i, i + 1) != i + 1) + throw new AssertionError("switch argument"); + } + } +}"#, + ) + .unwrap(); + for (program, args) in [ + ("javac", vec!["-cp", ".", "Run.java"]), + ("java", vec!["-Xverify:all", "-cp", ".", "Run"]), + ] { + let output = std::process::Command::new(program) + .args(args) + .current_dir(&dir) + .output() + .unwrap(); + assert!( + output.status.success(), + "{}", + String::from_utf8_lossy(&output.stderr) + ); + } + std::fs::remove_dir_all(dir).unwrap(); + } +} From e02f9e8def8fb6f40148a83b8ae8454e4da9588f Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 02:45:25 +1000 Subject: [PATCH 16/61] decompose borrowed views --- compiler-core/src/analysis/arrays_tests.rs | 240 ++++++++++++++++++++ compiler-core/src/analysis/mod.rs | 2 + compiler-core/src/jvm/select/view_tests.rs | 23 +- compiler-core/src/opt/mod.rs | 26 +++ compiler-core/src/opt/view_abi.rs | 165 ++++++++++++++ compiler-core/src/opt/views.rs | 249 +++++++++++++++++++++ compiler-core/src/opt/views_tests.rs | 215 ++++++++++++++++++ 7 files changed, 917 insertions(+), 3 deletions(-) create mode 100644 compiler-core/src/analysis/arrays_tests.rs create mode 100644 compiler-core/src/opt/view_abi.rs create mode 100644 compiler-core/src/opt/views.rs create mode 100644 compiler-core/src/opt/views_tests.rs diff --git a/compiler-core/src/analysis/arrays_tests.rs b/compiler-core/src/analysis/arrays_tests.rs new file mode 100644 index 00000000..d80b6c83 --- /dev/null +++ b/compiler-core/src/analysis/arrays_tests.rs @@ -0,0 +1,240 @@ +use super::native_array_accesses; +use crate::classfile::attributes::Instruction; +use crate::ir::*; +use crate::scalar::{BinaryOp, Scalar, ScalarType}; + +#[test] +fn initialization_is_native_only_until_the_first_escape() { + for unwind in [false, true] { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let array_ty = types.intern(Type::Array(int)); + let unit = types.intern(Type::Unit); + let mut b = Builder::new(&types, int); + let one = b.constant(int, Scalar::integer(ScalarType::I32, 1).unwrap()); + let zero = b.constant(int, Scalar::integer(ScalarType::I32, 0).unwrap()); + let array = if unwind { + let handler = b.create_block(); + let value = b + .invoke(Op::NewArray(one), Some(array_ty), handler) + .unwrap(); + let normal = b.current(); + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + b.switch_to(normal); + value + } else { + b.emit(Op::NewArray(one), Some(array_ty)).unwrap() + }; + let before = b.body.instructions.len(); + b.emit( + Op::ArraySet { + array, + index: zero, + value: one, + native: false, + }, + None, + ); + let method = b.method(MethodRef { + owner: "Escape".into(), + name: "registerAlias".into(), + params: vec![array_ty], + returns: unit, + interface: false, + }); + let args = b.args([array]); + b.emit( + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }, + None, + ); + let after = b.body.instructions.len(); + let value = b + .emit( + Op::ArrayGet { + array, + index: zero, + native: false, + }, + Some(int), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(value))); + let body = b.finish().unwrap(); + let native = native_array_accesses(&body, &types, &crate::opt::live(&body, &types)); + assert_eq!(native[before], Some(int)); + assert_eq!(native[after], None); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!(code.instructions.contains(&Instruction::Iastore)); + assert!(!code.instructions.contains(&Instruction::Iaload)); + } +} + +#[test] +fn private_slice_loop_uses_native_array_instructions() { + for escape in [false, true] { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let boolean = types.scalar(ScalarType::Bool); + let long = types.scalar(ScalarType::U64); + let array_ty = types.intern(Type::Array(int)); + let slice_ty = types.intern(Type::Slice(int)); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let unit = types.intern(Type::Unit); + let mut b = Builder::new(&types, int); + let limit = b.parameter(b.current(), int); + let one = b.constant(int, Scalar::integer(ScalarType::I32, 1).unwrap()); + let three = b.constant(int, Scalar::integer(ScalarType::I32, 3).unwrap()); + let array = b.emit(Op::NewArray(three), Some(array_ty)).unwrap(); + b.emit(Op::ArrayFill { array, value: one }, None); + let root = b.emit(Op::Reinterpret(array), Some(object)).unwrap(); + let length = b.constant(long, Scalar::integer(ScalarType::U64, 2).unwrap()); + let parts = b.args([root, one, length]); + let view = b.emit(Op::ViewPack(parts), Some(slice_ty)).unwrap(); + let index = b.variable(int); + b.define(index, one); + let header = b.create_block(); + let step = b.create_block(); + let done = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + let iteration = b.read(index); + let condition = b + .emit( + Op::Binary { + op: BinaryOp::Lt, + left: iteration, + right: limit, + }, + Some(boolean), + ) + .unwrap(); + b.branch(condition, step, done); + b.switch_to(step); + b.emit( + Op::ArraySet { + native: false, + array: view, + index: one, + value: iteration, + }, + None, + ); + if escape { + let method = b.method(MethodRef { + owner: "Escape".into(), + name: "registerAlias".into(), + params: vec![object], + returns: unit, + interface: false, + }); + let args = b.args([root]); + b.emit( + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }, + None, + ); + } + let loaded = b + .emit( + Op::ArrayGet { + native: false, + array: view, + index: one, + }, + Some(int), + ) + .unwrap(); + let next = b + .emit( + Op::Binary { + op: BinaryOp::Add, + left: loaded, + right: one, + }, + Some(int), + ) + .unwrap(); + b.define(index, next); + b.jump(header, vec![]); + b.switch_to(done); + b.terminate(Terminator::Return(Some(iteration))); + let mut body = b.finish().unwrap(); + crate::opt::decompose_views(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = crate::opt::live(&body, &types); + let native = native_array_accesses(&body, &types, &live); + for (id, inst) in body.instructions.iter().enumerate() { + if matches!(inst.op, Op::ViewGet(_) | Op::ViewSet { .. }) { + assert_eq!(native[id], (!escape).then_some(int)); + } + } + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!( + !code + .instructions + .iter() + .any(|i| matches!(i, Instruction::New(_))) + ); + assert_eq!(code.instructions.contains(&Instruction::Iaload), !escape); + assert_eq!(code.instructions.contains(&Instruction::Iastore), !escape); + assert!( + !code + .instructions + .iter() + .any(|i| matches!(i, Instruction::Getfield(_))) + ); + } +} + +#[test] +fn mixed_array_origins_escape_even_when_the_join_only_reads() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let boolean = types.scalar(ScalarType::Bool); + let array_ty = types.intern(Type::Array(int)); + let mut b = Builder::new(&types, int); + let condition = b.parameter(b.current(), boolean); + let input = b.parameter(b.current(), array_ty); + let one = b.constant(int, Scalar::integer(ScalarType::I32, 1).unwrap()); + let array = b.emit(Op::NewArray(one), Some(array_ty)).unwrap(); + let variable = b.variable(array_ty); + b.define(variable, array); + let other = b.create_block(); + let joined = b.create_block(); + b.branch(condition, other, joined); + b.switch_to(other); + b.define(variable, input); + b.jump(joined, vec![]); + b.switch_to(joined); + let alias = b.read(variable); + let zero = b.constant(int, Scalar::integer(ScalarType::I32, 0).unwrap()); + let loaded = b + .emit( + Op::ArrayGet { + native: false, + array: alias, + index: zero, + }, + Some(int), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(loaded))); + let body = b.finish().unwrap(); + let live = crate::opt::live(&body, &types); + assert!( + native_array_accesses(&body, &types, &live) + .iter() + .all(Option::is_none) + ); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!(!code.instructions.contains(&Instruction::Iaload)); +} diff --git a/compiler-core/src/analysis/mod.rs b/compiler-core/src/analysis/mod.rs index 2728e883..aac7fd87 100644 --- a/compiler-core/src/analysis/mod.rs +++ b/compiler-core/src/analysis/mod.rs @@ -5,3 +5,5 @@ mod users; pub(crate) use users::ValueUsers; mod arrays; pub(crate) use arrays::native_array_accesses; +#[cfg(test)] +mod arrays_tests; diff --git a/compiler-core/src/jvm/select/view_tests.rs b/compiler-core/src/jvm/select/view_tests.rs index 599bc4c8..30eb0a76 100644 --- a/compiler-core/src/jvm/select/view_tests.rs +++ b/compiler-core/src/jvm/select/view_tests.rs @@ -47,7 +47,17 @@ fn executes_full_width_view_metadata_and_typed_data_extraction() { let mut invalid = body.clone(); invalid.instructions[1].op = Op::Length(data); assert!(verify(&invalid, &types).is_err()); + let mut body = body; + crate::opt::decompose_views(&mut body, &mut types, None); + verify(&body, &types).unwrap(); let code = compile(&body, &types, &mut cp).unwrap(); + assert!( + !code + .instructions + .iter() + .any(|i| matches!(i, Instruction::New(_))), + "local view allocated a carrier" + ); methods.push(Method { access_flags: MethodAccessFlags::PUBLIC | MethodAccessFlags::STATIC, name_index: cp.add_utf8(name).unwrap(), @@ -131,11 +141,18 @@ public class Utf8View extends SliceView { package org.rustlang.runtime; public class Pointer { public static Pointer extracted; + public static long extractedLength; public static Pointer fromSlice(Object value, int size, String codec) { SliceView view = (SliceView)value; if (view.offset != 0 || codec != null || size != (view instanceof Utf8View ? 1 : 8)) throw new AssertionError("wrong extraction arguments"); - return extracted = (Pointer)view.array; + return fromSliceParts(view.array, view.offset, view.rustLength, size, codec); + } + public static Pointer fromSliceParts(Object data, int start, long length, int size, String codec) { + if (start != 0 || codec != null || (size != 1 && size != 8)) + throw new AssertionError("wrong component arguments"); + extractedLength = length; + return extracted = (Pointer)data; } }"#, ) @@ -146,9 +163,9 @@ public class ViewRun { public static void main(String[] args) { Pointer pointer = new Pointer(); for (long length : new long[] { 0, 1, Integer.MAX_VALUE, (1L << 40) + 23, Long.MAX_VALUE }) { - if (ViewSsa.slice(pointer, length) != length || Pointer.extracted != pointer) + if (ViewSsa.slice(pointer, length) != length || Pointer.extracted != pointer || Pointer.extractedLength != length) throw new AssertionError("slice " + length); - if (ViewSsa.string(pointer, length) != length || Pointer.extracted != pointer) + if (ViewSsa.string(pointer, length) != length || Pointer.extracted != pointer || Pointer.extractedLength != length) throw new AssertionError("string " + length); if (ViewSsa.read_slice(new SliceView(pointer, 3, length)) != length || ViewSsa.read_string(new Utf8View(pointer, 2, length)) != length) diff --git a/compiler-core/src/opt/mod.rs b/compiler-core/src/opt/mod.rs index 1e67a044..677c2573 100644 --- a/compiler-core/src/opt/mod.rs +++ b/compiler-core/src/opt/mod.rs @@ -10,10 +10,36 @@ mod cells; pub use cells::promote_cells; mod tagged; pub use tagged::decompose_tagged; +mod views; +pub use views::decompose_views; +#[cfg(test)] +mod views_tests; #[cfg(test)] mod cells_tests; +mod view_abi; +pub use view_abi::{component_argument_slots, lower_component_arguments}; + mod simplify; mod unreachable; pub use simplify::simplify_components; + +fn append_value( + body: &mut crate::ir::Body, + op: crate::ir::Op, + ty: crate::ir::TypeId, +) -> (crate::ir::InstId, crate::ir::ValueId) { + use crate::ir::{Inst, InstId, Value, ValueDef, ValueId}; + let inst = InstId::new(body.instructions.len()); + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty, + def: ValueDef::Inst(inst), + }); + body.instructions.push(Inst { + op, + result: Some(value), + }); + (inst, value) +} diff --git a/compiler-core/src/opt/view_abi.rs b/compiler-core/src/opt/view_abi.rs new file mode 100644 index 00000000..43947cd0 --- /dev/null +++ b/compiler-core/src/opt/view_abi.rs @@ -0,0 +1,165 @@ +//! Flatten borrowed-view arguments before instruction selection. +//! Keep boundary values as virtual packs. SSA liveness determines whether a +//! consumer needs a JVM carrier. Retain only one body's physical operands. +use super::append_value as append; +use crate::ir::*; +use crate::scalar::ScalarType; +use rustc_hash::FxHashMap; + +pub fn component_argument_slots(types: &Types, params: impl IntoIterator) -> usize { + params + .into_iter() + .map(|ty| match types.get(ty).unwrap() { + _ if ComponentShape::of(types, ty).is_some() => { + ComponentShape::of(types, ty).unwrap().slots() + } + Type::Scalar(ScalarType::I64 | ScalarType::U64 | ScalarType::F64) => 2, + Type::Unit | Type::Opaque(_) | Type::Layout(_) => 0, + _ => 1, + }) + .sum() +} + +pub fn lower_component_arguments( + body: &mut Body, + types: &mut Types, + entry: bool, + select: impl Fn(&MethodRef) -> bool, + debug: Option<&mut DebugInfo>, +) { + let shape = |ty| ComponentShape::of(types, ty); + let methods = body + .methods + .iter() + .map(|method| { + (select(method) + && component_argument_slots(types, method.params.iter().copied()) + + usize::from(ComponentShape::of(types, method.returns).is_some()) + <= 254) + .then(|| { + method + .params + .iter() + .map(|&ty| shape(ty)) + .collect::>() + }) + .filter(|p| p.iter().any(Option::is_some)) + }) + .collect::>(); + let entry = entry + && body.blocks[body.entry.index()] + .params + .iter() + .any(|&p| shape(body.value_type(p)).is_some()); + if !entry && methods.iter().all(Option::is_none) { + return; + } + let mut prologue = Vec::new(); + if entry { + let params = std::mem::take(&mut body.blocks[body.entry.index()].params); + for value in params { + let Some(shape) = ComponentShape::of(types, body.value_type(value)) else { + body.blocks[body.entry.index()].params.push(value); + continue; + }; + let args = shape + .parts(types) + .into_iter() + .map(|ty| { + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty, + def: ValueDef::Param(body.entry), + }); + body.blocks[body.entry.index()].params.push(value); + value + }) + .collect::>(); + let id = InstId::new(body.instructions.len()); + body.values[value.index()].def = ValueDef::Inst(id); + let args = List::append(&mut body.args, args); + body.instructions.push(Inst { + op: shape.pack(args), + result: Some(value), + }); + prologue.push(id); + } + } + let mut prefixes = FxHashMap::>::default(); + let count = body.instructions.len(); + for index in 0..count { + let Op::Call { method, kind, args } = body.instructions[index].op else { + continue; + }; + let Some(expand) = &methods[method.index()] else { + continue; + }; + let receiver = usize::from(matches!(kind, CallKind::Virtual | CallKind::Interface)); + let original = body.args[args.range()].to_vec(); + let mut arguments = original[..receiver].to_vec(); + let mut prefix = Vec::new(); + for (&value, &expand) in original[receiver..].iter().zip(expand) { + if let Some(shape) = expand { + for (index, ty) in shape.parts(types).into_iter().enumerate() { + let (inst, value) = append(body, shape.part(value, index as u8), ty); + prefix.push(inst); + arguments.push(value); + } + } else { + arguments.push(value); + } + } + let args = List::append(&mut body.args, arguments); + body.instructions[index].op = Op::Call { method, kind, args }; + prefixes.insert(InstId::new(index), prefix); + } + for (method, expand) in body.methods.iter_mut().zip(methods) { + if let Some(expand) = expand { + let original = std::mem::take(&mut method.params); + for (ty, expand) in original.into_iter().zip(expand) { + if let Some(shape) = expand { + method.params.extend(shape.parts(types)); + } else { + method.params.push(ty); + } + } + } + } + let mut positions = debug + .as_ref() + .map(|_| Vec::with_capacity(body.blocks.len())); + for (index, block) in body.blocks.iter_mut().enumerate() { + let previous = std::mem::take(&mut block.instructions); + if index == body.entry.index() { + block.instructions.append(&mut prologue); + } + let mut mapping = positions + .as_ref() + .map(|_| Vec::with_capacity(previous.len() + 1)); + for id in previous { + if let Some(mapping) = &mut mapping { + mapping.push(block.instructions.len() as u32); + } + if let Some(prefix) = prefixes.remove(&id) { + block.instructions.extend(prefix); + } + block.instructions.push(id); + } + if let Some(mapping) = &mut mapping { + mapping.push(block.instructions.len() as u32); + } + if let Some(Terminator::Invoke { inst, .. }) = block.terminator { + if let Some(prefix) = prefixes.remove(&inst) { + block.instructions.extend(prefix); + } + } + if let (Some(positions), Some(mapping)) = (&mut positions, mapping) { + positions.push(mapping); + } + } + if let (Some(debug), Some(positions)) = (debug, positions) { + for event in &mut debug.events { + event.position = positions[event.block.index()][event.position as usize]; + } + } +} diff --git a/compiler-core/src/opt/views.rs b/compiler-core/src/opt/views.rs new file mode 100644 index 00000000..13e6e4d7 --- /dev/null +++ b/compiler-core/src/opt/views.rs @@ -0,0 +1,249 @@ +//! Decompose local views through SSA joins. +//! Data and length consumers use scalars. Opaque boundaries allocate carriers only when needed. +use super::append_value as append; +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; +use rustc_hash::FxHashMap; + +pub fn decompose_views(body: &mut Body, types: &mut Types, debug: Option<&mut DebugInfo>) { + if !body.instructions.iter().any(|i| { + matches!( + i.op, + Op::View { .. } | Op::ViewPack(_) | Op::ViewPart { .. } + ) + }) { + return; + } + let count = body.values.len(); + let mut eligible = vec![false; count]; + let mut users = crate::analysis::ValueUsers::new(count); + let predecessors = body.predecessors(); + for index in 0..count { + let value = ValueId::new(index); + if !ComponentShape::View.accepts_annotation(types, body.value_type(value)) { + continue; + } + match body.values[index].def { + ValueDef::Inst(id) => match body.instructions[id.index()].op { + Op::View { .. } | Op::ViewPack(_) => eligible[index] = true, + Op::Constant(id) if matches!(body.constants[id.index()], Constant::Null(_)) => { + eligible[index] = true; + } + Op::Reinterpret(source) | Op::Adapt(source) => { + eligible[index] = true; + users.connect(source, index); + } + _ => {} + }, + ValueDef::Alias(source) => { + eligible[index] = true; + users.connect(source, index); + } + ValueDef::Param(block) if block != body.entry => { + eligible[index] = !predecessors[block.index()].is_empty(); + let position = body.blocks[block.index()] + .params + .iter() + .position(|&v| v == value) + .unwrap(); + for &(_, edge) in &predecessors[block.index()] { + users.connect(body.edges[edge.index()].args[position], index); + } + } + _ => {} + } + } + // Reject each node at most once, including cycles with unknown inputs. + // No repeated whole-body scan is needed. + users.close(&mut eligible); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let int = types.scalar(ScalarType::I32); + let length_type = types.scalar(ScalarType::U64); + let mut components = vec![None::<[ValueId; 3]>; count]; + let mut prefixes = FxHashMap::>::default(); + let mut joins = vec![Vec::new(); body.blocks.len()]; + for index in 0..count { + if !eligible[index] { + continue; + } + match body.values[index].def { + ValueDef::Param(block) => { + let mut parts = [ValueId::new(0); 3]; + for (part, ty) in parts.iter_mut().zip([object, int, length_type]) { + *part = ValueId::new(body.values.len()); + body.values.push(Value { + ty, + def: ValueDef::Param(block), + }); + body.blocks[block.index()].params.push(*part); + } + let position = body.blocks[block.index()] + .params + .iter() + .position(|&value| value.index() == index) + .unwrap(); + joins[block.index()].push((index, position)); + components[index] = Some(parts); + } + ValueDef::Inst(id) => { + // Partially initialized aggregates use null borrowed fields. + // Read their default component slots without dereferencing a view carrier. + if matches!(body.instructions[id.index()].op, Op::Constant(_)) { + let mut parts = [ValueId::new(0); 3]; + let mut prefix = Vec::new(); + for (index, (ty, constant)) in [ + (object, Constant::Null(object)), + ( + int, + Constant::Scalar(Scalar::integer(ScalarType::I32, 0).unwrap()), + ), + ( + length_type, + Constant::Scalar(Scalar::integer(ScalarType::U64, 0).unwrap()), + ), + ] + .into_iter() + .enumerate() + { + let constant_id = ConstId::new(body.constants.len()); + body.constants.push(constant); + let (instruction, value) = append(body, Op::Constant(constant_id), ty); + prefix.push(instruction); + parts[index] = value; + } + prefixes.insert(id, prefix); + components[index] = Some(parts); + } + if let Op::ViewPack(parts) = body.instructions[id.index()].op { + components[index] = Some(body.args[parts.range()].try_into().unwrap()); + } + if let Op::View { data, length } = body.instructions[id.index()].op { + let (cast, data) = append(body, Op::Reinterpret(data), object); + let constant = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar( + Scalar::integer(ScalarType::I32, 0).unwrap(), + )); + let (zero, start) = append(body, Op::Constant(constant), int); + prefixes.insert(id, vec![cast, zero]); + components[index] = Some([data, start, length]); + } + } + _ => {} + } + } + users.propagate(&eligible, &mut components); + for block in 0..body.blocks.len() { + // Visit parameters in the same order used when extending their target. + for &(_, edge) in &predecessors[block] { + for &(_, position) in &joins[block] { + let input = body.edges[edge.index()].args[position]; + body.edges[edge.index()] + .args + .extend(components[input.index()].unwrap()); + } + } + } + for inst in &mut body.instructions { + inst.op = match inst.op { + Op::Adapt(source) + if inst + .result + .and_then(|v| components.get(v.index()).copied().flatten()) + .is_some() => + { + let result = inst.result.unwrap(); + if matches!( + types.get(body.values[result.index()].ty), + Some(Type::Slice(_) | Type::Str) + ) { + Op::ViewPack(List::append( + &mut body.args, + components[result.index()].unwrap(), + )) + } else { + Op::Reinterpret(source) + } + } + Op::Length(view) if components.get(view.index()).is_some_and(Option::is_some) => { + Op::Reinterpret(components[view.index()].unwrap()[2]) + } + Op::ArrayLength(view) if components.get(view.index()).is_some_and(Option::is_some) => { + // SliceView.length contains the low 32 bits of the Rust length. + // Extract them without a view carrier. + Op::Cast(components[view.index()].unwrap()[2]) + } + Op::ViewPart { view, index } + if components.get(view.index()).is_some_and(Option::is_some) => + { + Op::Reinterpret(components[view.index()].unwrap()[index as usize]) + } + Op::ViewData { view, size, codec } + if components.get(view.index()).is_some_and(Option::is_some) => + { + let parts = List::append(&mut body.args, components[view.index()].unwrap()); + Op::ViewAddress { parts, size, codec } + } + Op::ArrayGet { + array, + index, + native: false, + } if components.get(array.index()).is_some_and(Option::is_some) => { + let [root, start, _] = components[array.index()].unwrap(); + Op::ViewGet(List::append(&mut body.args, [root, start, index])) + } + Op::ArraySet { + array, + index, + value, + native: false, + } if components.get(array.index()).is_some_and(Option::is_some) => { + let [root, start, _] = components[array.index()].unwrap(); + Op::ViewSet { + parts: List::append(&mut body.args, [root, start, index]), + value, + } + } + op => op, + }; + } + let mut positions = debug + .as_ref() + .map(|_| Vec::with_capacity(body.blocks.len())); + for block in &mut body.blocks { + let mut mapping = positions + .as_ref() + .map(|_| Vec::with_capacity(block.instructions.len() + 1)); + let previous = std::mem::take(&mut block.instructions); + for id in previous { + if let Some(mapping) = &mut mapping { + mapping.push(block.instructions.len() as u32); + } + if let Some(prefix) = prefixes.remove(&id) { + block.instructions.extend(prefix); + } + block.instructions.push(id); + } + if let Some(mapping) = &mut mapping { + mapping.push(block.instructions.len() as u32); + } + if let (Some(positions), Some(mapping)) = (&mut positions, mapping) { + positions.push(mapping); + } + if let Some(Terminator::Invoke { inst, normal, .. }) = block.terminator { + if let Some(prefix) = prefixes.remove(&inst) { + block.instructions.extend(prefix); + } + let op = body.instructions[inst.index()].op; + if matches!(op, Op::View { .. } | Op::ViewPack(_) | Op::Reinterpret(_)) { + block.instructions.push(inst); + block.terminator = Some(Terminator::Jump(normal)); + } + } + } + if let (Some(debug), Some(positions)) = (debug, positions) { + for event in &mut debug.events { + event.position = positions[event.block.index()][event.position as usize]; + } + } +} diff --git a/compiler-core/src/opt/views_tests.rs b/compiler-core/src/opt/views_tests.rs new file mode 100644 index 00000000..127b11e4 --- /dev/null +++ b/compiler-core/src/opt/views_tests.rs @@ -0,0 +1,215 @@ +use super::decompose_views; +use crate::ir::*; +use crate::scalar::{BinaryOp, Scalar, ScalarType}; + +#[test] +fn jvm_view_length_truncates_without_constructing_the_view() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let long = types.scalar(ScalarType::U64); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let slice = types.intern(Type::Slice(int)); + let mut b = Builder::new(&types, int); + let root = b.parameter(b.current(), object); + let start = b.constant(int, Scalar::integer(ScalarType::I32, 0).unwrap()); + let length = b.constant( + long, + Scalar::integer(ScalarType::U64, 0x1_0000_0007).unwrap(), + ); + let parts = b.args([root, start, length]); + let view = b.emit(Op::ViewPack(parts), Some(slice)).unwrap(); + let result = b.emit(Op::ArrayLength(view), Some(int)).unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + decompose_views(&mut body, &mut types, None); + super::simplify_components(&mut body, &types); + verify(&body, &types).unwrap(); + assert_eq!(body.scalar_value(result).unwrap().bits(), 7); + let live = super::live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!(inst.op, Op::ViewPack(_) | Op::ArrayLength(_))) + ); +} + +#[test] +fn view_components_cross_loops_without_a_carrier() { + let mut types = Types::default(); + let long = types.scalar(ScalarType::U64); + let boolean = types.scalar(ScalarType::Bool); + let pointer = types.intern(Type::Pointer(long)); + let slice = types.intern(Type::Slice(long)); + let runtime = types.symbol("org/rustlang/runtime/SliceView"); + let runtime = types.intern(Type::Class(runtime)); + let mut b = Builder::new(&types, pointer); + let data = b.parameter(b.current(), pointer); + let length = b.parameter(b.current(), long); + let initial = b.emit(Op::View { data, length }, Some(slice)).unwrap(); + let initial = b.emit(Op::Reinterpret(initial), Some(runtime)).unwrap(); + let initial = b.emit(Op::Adapt(initial), Some(slice)).unwrap(); + let value = b.variable(slice); + b.define(value, initial); + let header = b.create_block(); + let step = b.create_block(); + let done = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + let view = b.read(value); + let length = b.emit(Op::Length(view), Some(long)).unwrap(); + let zero = b.constant(long, Scalar::integer(ScalarType::U64, 0).unwrap()); + let more = b + .emit( + Op::Binary { + op: BinaryOp::Gt, + left: length, + right: zero, + }, + Some(boolean), + ) + .unwrap(); + b.branch(more, step, done); + b.switch_to(step); + let one = b.constant(long, Scalar::integer(ScalarType::U64, 1).unwrap()); + let next = b + .emit( + Op::Binary { + op: BinaryOp::Sub, + left: length, + right: one, + }, + Some(long), + ) + .unwrap(); + let view = b + .emit(Op::View { data, length: next }, Some(slice)) + .unwrap(); + b.define(value, view); + b.jump(header, vec![]); + b.switch_to(done); + let view = b.read(value); + let pointer_value = b + .emit( + Op::ViewData { + view, + size: 8, + codec: None, + }, + Some(pointer), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(pointer_value))); + let mut body = b.finish().unwrap(); + decompose_views(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + assert!( + body.instructions + .iter() + .enumerate() + .all(|(i, inst)| !live.instructions[i] + || !matches!( + inst.op, + Op::View { .. } + | Op::ViewPack(_) + | Op::Adapt(_) + | Op::Length(_) + | Op::ViewData { .. } + )) + ); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!( + !code + .instructions + .iter() + .any(|i| matches!(i, crate::classfile::attributes::Instruction::New(_))) + ); +} + +#[test] +fn unknown_input_keeps_a_mixed_join_conservative() { + let mut types = Types::default(); + let long = types.scalar(ScalarType::U64); + let boolean = types.scalar(ScalarType::Bool); + let pointer = types.intern(Type::Pointer(long)); + let slice = types.intern(Type::Slice(long)); + let mut b = Builder::new(&types, long); + let input = b.parameter(b.current(), slice); + let condition = b.parameter(b.current(), boolean); + let data = b.parameter(b.current(), pointer); + let value = b.variable(slice); + b.define(value, input); + let create = b.create_block(); + let join = b.create_block(); + b.branch(condition, create, join); + b.switch_to(create); + let length = b.constant(long, Scalar::integer(ScalarType::U64, 17).unwrap()); + let constructed = b.emit(Op::View { data, length }, Some(slice)).unwrap(); + b.define(value, constructed); + b.jump(join, vec![]); + b.switch_to(join); + let joined = b.read(value); + let length = b.emit(Op::Length(joined), Some(long)).unwrap(); + b.terminate(Terminator::Return(Some(length))); + let mut body = b.finish().unwrap(); + decompose_views(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::Length(_))) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +#[test] +fn component_calls_and_entries_need_no_view_objects() { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let length = types.scalar(ScalarType::U64); + let slice = types.intern(Type::Slice(byte)); + let mut b = Builder::new(&types, length); + let input = b.parameter(b.current(), slice); + let target = b.method(MethodRef { + owner: "Leaf".into(), + name: "len".into(), + params: vec![slice], + returns: length, + interface: false, + }); + let args = b.args([input]); + let result = b + .emit( + Op::Call { + method: target, + kind: CallKind::RustStatic, + args, + }, + Some(length), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + super::lower_component_arguments(&mut body, &mut types, true, |_| true, None); + decompose_views(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + assert_eq!(body.blocks[body.entry.index()].params.len(), 3); + assert_eq!(body.methods[target.index()].params.len(), 3); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + use crate::classfile::attributes::Instruction; + assert!(!code.instructions.iter().any(|i| matches!( + i, + Instruction::New(_) | Instruction::Getfield(_) | Instruction::Checkcast(_) + ))); + assert_eq!( + code.instructions + .iter() + .filter(|i| matches!(i, Instruction::Invokestatic(_))) + .count(), + 1 + ); +} From 339d739a09a5332fb40afb7454c7854b20aab5dc Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 04:40:33 +1000 Subject: [PATCH 17/61] flatten borrowed address ABI --- compiler-core/src/opt/address_parts.rs | 70 +++ compiler-core/src/opt/addresses.rs | 666 ++++++++++++++++++++++ compiler-core/src/opt/addresses_tests.rs | 544 ++++++++++++++++++ compiler-core/src/opt/field_abi.rs | 213 +++++++ compiler-core/src/opt/field_abi_tests.rs | 290 ++++++++++ compiler-core/src/opt/mod.rs | 18 + compiler-core/src/opt/return_abi.rs | 304 ++++++++++ compiler-core/src/opt/return_abi_tests.rs | 154 +++++ 8 files changed, 2259 insertions(+) create mode 100644 compiler-core/src/opt/address_parts.rs create mode 100644 compiler-core/src/opt/addresses.rs create mode 100644 compiler-core/src/opt/addresses_tests.rs create mode 100644 compiler-core/src/opt/field_abi.rs create mode 100644 compiler-core/src/opt/field_abi_tests.rs create mode 100644 compiler-core/src/opt/return_abi.rs create mode 100644 compiler-core/src/opt/return_abi_tests.rs diff --git a/compiler-core/src/opt/address_parts.rs b/compiler-core/src/opt/address_parts.rs new file mode 100644 index 00000000..41fb1c19 --- /dev/null +++ b/compiler-core/src/opt/address_parts.rs @@ -0,0 +1,70 @@ +//! Recover address components from representation lowering. +//! Read the defining pack's exact layout, even after JVM type erasure. +use crate::ir::*; + +pub(super) struct AddressParts { + pub value: ValueId, + pub parts: List, + pub pointee: Option, + pub layout: Option<(u32, Option)>, +} + +pub(super) fn address_parts(body: &Body, types: &Types, value: ValueId) -> Option { + find_parts(body, types, value, true) +} + +/// Scalar pointer casts replace the view width and codec. +/// They can ignore erased layout annotations, unlike ordinary ABI uses. +pub(super) fn scalar_cast_parts( + body: &Body, + types: &Types, + value: ValueId, +) -> Option { + find_parts(body, types, value, false) +} + +fn find_parts( + body: &Body, + types: &Types, + mut value: ValueId, + preserve_layout: bool, +) -> Option { + for _ in 0..32 { + value = body.resolve(value); + let ValueDef::Inst(id) = body.values[value.index()].def else { + return None; + }; + let ty = body.value_type(value); + let pointee = types.pointee(ty); + match body.instructions[id.index()].op { + Op::TypedAddressPack { parts, size, codec } => { + return Some(AddressParts { + value, + parts, + pointee, + layout: Some((size, codec)), + }); + } + Op::AddressPack(parts) => { + let layout = types + .address_layout(ty) + .or_else(|| StorageSlot::scalar(pointee?, types).map(|slot| (slot.size, None))); + return Some(AddressParts { + value, + parts, + pointee, + layout, + }); + } + Op::Refine(source) | Op::Reinterpret(source) + if !preserve_layout + || types.address_layout(ty) + == types.address_layout(body.value_type(source)) => + { + value = source + } + _ => return None, + } + } + None +} diff --git a/compiler-core/src/opt/addresses.rs b/compiler-core/src/opt/addresses.rs new file mode 100644 index 00000000..3955d0cb --- /dev/null +++ b/compiler-core/src/opt/addresses.rs @@ -0,0 +1,666 @@ +//! Keep scalar addresses as a storage root and byte displacement through loops +//! and calls. Unknown roots retain their full runtime provenance as one object. +use super::append_value as append; +use crate::ir::*; +use crate::scalar::{BinaryOp, Scalar, ScalarType}; +use rustc_hash::{FxHashMap, FxHashSet}; + +fn literal(body: &mut Body, ty: TypeId, bits: i64) -> (InstId, ValueId) { + let id = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar( + Scalar::integer(ScalarType::I64, bits as u128).unwrap(), + )); + append(body, Op::Constant(id), ty) +} +fn parts( + body: &mut Body, + types: &Types, + value: ValueId, + known: &[Option<[ValueId; 2]>], + object: TypeId, + long: TypeId, + prefix: &mut Vec, +) -> [ValueId; 2] { + if let Some(parts) = known.get(value.index()).copied().flatten() { + return parts; + } + // Scalar projections can use an aggregate location's components. + // Their borrowed shapes do not need to match. + let mut source = value; + for _ in 0..64 { + source = body.resolve(source); + let ValueDef::Inst(id) = body.values[source.index()].def else { + break; + }; + match body.instructions[id.index()].op { + Op::AddressPack(parts) => return body.args[parts.range()].try_into().unwrap(), + Op::Reinterpret(value) | Op::Refine(value) + if types.address_layout(body.value_type(source)) + == types.address_layout(body.value_type(value)) => + { + source = value + } + _ => break, + } + } + let (cast, root) = append(body, Op::Reinterpret(value), object); + let (zero, offset) = literal(body, long, 0); + prefix.extend([cast, zero]); + [root, offset] +} + +pub fn decompose_addresses(body: &mut Body, types: &mut Types, mut debug: Option<&mut DebugInfo>) { + for shape in [ComponentShape::Address, ComponentShape::StorageAddress] { + if body + .values + .iter() + .any(|v| ComponentShape::of(types, v.ty) == Some(shape)) + { + decompose(body, types, debug.as_deref_mut(), shape); + } + } +} + +fn decompose( + body: &mut Body, + types: &mut Types, + debug: Option<&mut DebugInfo>, + shape: ComponentShape, +) { + if !body.instructions.iter().any(|i| { + matches!( + i.op, + Op::AddressPack(_) + | Op::AddressOfSlot(_) + | Op::Offset { .. } + | Op::AddressPart { .. } + | Op::AddressTag(_) + | Op::AddressEqual { .. } + | Op::AddressCompare { .. } + | Op::ViewAddress { .. } + | Op::AddressViewPart { .. } + | Op::Project { .. } + ) + }) { + return; + } + let count = body.values.len(); + let predecessors = body.predecessors(); + let mut eligible = vec![false; count]; + let mut users = crate::analysis::ValueUsers::new(count); + for index in 0..count { + let ty = body.values[index].ty; + let address = ComponentShape::of(types, ty) == Some(shape); + if !shape.accepts_annotation(types, ty) { + continue; + } + match body.values[index].def { + ValueDef::Inst(id) => match body.instructions[id.index()].op { + Op::AddressPack(_) | Op::AddressOfSlot(_) | Op::Offset { .. } if address => { + eligible[index] = true + } + Op::Project { projection, .. } if address => { + let projection = &body.projections[projection.index()]; + let Type::Pointer(inner) = types.get(ty).unwrap() else { + unreachable!() + }; + eligible[index] = (projection.codec.is_none() + && StorageSlot::scalar(inner, types) + .is_some_and(|slot| u64::from(slot.size) == projection.size)) + || (ComponentShape::of(types, inner) + .is_some_and(ComponentShape::is_borrowed) + && matches!(projection.size, 8 | 16)); + } + Op::Constant(constant) + if address && matches!(body.constants[constant.index()], Constant::Null(_)) => + { + eligible[index] = true + } + Op::ViewAddress { + size, codec: None, .. + } if address && shape == ComponentShape::Address => { + let Some(Type::Pointer(inner)) = types.get(ty) else { + unreachable!() + }; + eligible[index] = + StorageSlot::scalar(inner, types).is_some_and(|slot| slot.size == size); + } + Op::Cast(source) + if address + && shape == ComponentShape::Address + && matches!(types.get(body.value_type(source)), Some(Type::Pointer(_))) => + { + eligible[index] = true + } + Op::Reinterpret(source) | Op::Adapt(source) + if types.address_layout(ty) + == types.address_layout(body.value_type(source)) => + { + eligible[index] = true; + users.connect(source, index); + } + _ => {} + }, + ValueDef::Alias(source) => { + eligible[index] = true; + users.connect(source, index); + } + ValueDef::Param(block) if block != body.entry => { + eligible[index] = !predecessors[block.index()].is_empty(); + let position = body.blocks[block.index()] + .params + .iter() + .position(|&p| p.index() == index) + .unwrap(); + for &(_, edge) in &predecessors[block.index()] { + users.connect(body.edges[edge.index()].args[position], index); + } + } + _ => {} + } + } + users.close(&mut eligible); + let mut component_types = shape.parts(types); + let (object, long) = ( + component_types.next().unwrap(), + component_types.next().unwrap(), + ); + let mut known = vec![None::<[ValueId; 2]>; count]; + let mut roots = Vec::new(); + let mut joins = vec![Vec::new(); body.blocks.len()]; + for index in 0..count { + if !eligible[index] { + continue; + } + match body.values[index].def { + ValueDef::Param(block) => { + let parts = [object, long].map(|ty| { + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty, + def: ValueDef::Param(block), + }); + body.blocks[block.index()].params.push(value); + value + }); + known[index] = Some(parts); + let position = body.blocks[block.index()] + .params + .iter() + .position(|&value| value.index() == index) + .unwrap(); + joins[block.index()].push((index, position)); + } + ValueDef::Inst(id) => match body.instructions[id.index()].op { + Op::AddressPack(args) => { + known[index] = Some(body.args[args.range()].try_into().unwrap()) + } + Op::Constant(_) + | Op::AddressOfSlot(_) + | Op::Offset { .. } + | Op::Cast(_) + | Op::Project { .. } + | Op::ViewAddress { .. } => { + let a = append(body, Op::Nop, object); + let b = append(body, Op::Nop, long); + known[index] = Some([a.1, b.1]); + roots.push((id, a.0, b.0)); + } + _ => {} + }, + _ => {} + } + } + users.propagate(&eligible, &mut known); + let mut prefixes = FxHashMap::>::default(); + for (id, root, offset) in roots { + let original = body.instructions[id.index()]; + let mut prefix = Vec::new(); + let (root_op, offset_op) = match original.op { + Op::Project { base, projection } => { + let source = parts(body, types, base, &known, object, long, &mut prefix); + let displacement = body.projections[projection.index()].offset; + let (constant, displacement) = literal(body, long, displacement as i64); + let (sum_inst, sum) = append( + body, + Op::Binary { + op: BinaryOp::Add, + left: source[1], + right: displacement, + }, + long, + ); + prefix.extend([constant, sum_inst]); + ( + Op::ProjectRoot { + address: List::append(&mut body.args, source), + projection, + }, + Op::ProjectOffset { + root: body.instructions[root.index()].result.unwrap(), + base: source[0], + offset: sum, + }, + ) + } + Op::Constant(_) => { + let null = ConstId::new(body.constants.len()); + body.constants.push(Constant::Null(object)); + let (zero, value) = literal(body, long, 0); + prefix.push(zero); + (Op::Constant(null), Op::Reinterpret(value)) + } + Op::ViewAddress { parts, size, .. } => { + let values = body.args[parts.range()].to_vec(); + let int = types.scalar(ScalarType::I32); + let constant = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar( + Scalar::integer(ScalarType::I32, size.into()).unwrap(), + )); + let (size_inst, size_arg) = append(body, Op::Constant(constant), int); + let (cast, start) = append(body, Op::Cast(values[1]), long); + let (stride_inst, stride) = literal(body, long, size as i64); + prefix.extend([size_inst, cast, stride_inst]); + let method = MethodId::new(body.methods.len()); + body.methods.push(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: "scalarSliceRoot".into(), + params: vec![object, int], + returns: object, + interface: false, + }); + ( + Op::Call { + method, + kind: CallKind::JvmStatic, + args: List::append(&mut body.args, [values[0], size_arg]), + }, + Op::Binary { + op: BinaryOp::Mul, + left: start, + right: stride, + }, + ) + } + Op::AddressOfSlot(slot) => { + let (zero, value) = literal(body, long, 0); + prefix.push(zero); + (Op::SlotRoot(slot), Op::Reinterpret(value)) + } + Op::Cast(value) => { + let source = super::address_parts::scalar_cast_parts(body, types, value) + .map(|address| body.args[address.parts.range()].try_into().unwrap()) + .unwrap_or_else(|| { + parts(body, types, value, &known, object, long, &mut prefix) + }); + (Op::Reinterpret(source[0]), Op::Reinterpret(source[1])) + } + Op::Offset { + pointer, + offset: delta, + bytes, + .. + } => { + let mut source = parts(body, types, pointer, &known, object, long, &mut prefix); + if let Some(Type::Pointer(inner)) = types.get(body.value_type(pointer)) + && ComponentShape::of(types, inner).is_some_and(ComponentShape::is_borrowed) + { + // Reference-to-borrow arithmetic needs the exact Rust stride. + // A fixed-array borrow uses one word. A slice borrow uses two. + // The enclosing storage size does not determine this stride. + let int = types.scalar(ScalarType::I32); + let plan = crate::jvm::abi::address_plan(types, inner); + let constant = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar( + Scalar::integer(ScalarType::I32, plan.into()).unwrap(), + )); + let (id, plan) = append(body, Op::Constant(constant), int); + prefix.push(id); + let method = MethodId::new(body.methods.len()); + let carrier = types.symbol("org/rustlang/runtime/Pointer"); + let carrier = types.intern(Type::Class(carrier)); + body.methods.push(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: "addressFromParts".into(), + params: vec![object, long, int], + returns: carrier, + interface: false, + }); + let args = List::append(&mut body.args, [source[0], source[1], plan]); + let (id, root) = append( + body, + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }, + carrier, + ); + let (erase, root) = append(body, Op::Reinterpret(root), object); + let (zero, displacement) = literal(body, long, 0); + prefix.extend([id, erase, zero]); + source = [root, displacement]; + } + let (cast, mut delta) = append(body, Op::Cast(delta), long); + prefix.push(cast); + if !bytes { + let Some(Type::Pointer(inner)) = types.get(body.value_type(pointer)) else { + unreachable!() + }; + let (constant, size) = + if let Some((size, _)) = types.address_layout(body.value_type(pointer)) { + literal(body, long, size as i64) + } else if let Some(slot) = StorageSlot::scalar(inner, types) { + literal(body, long, slot.size as i64) + } else if let ValueDef::Inst(id) = + body.values[body.resolve(pointer).index()].def + && let Op::RetypeAddress { size, .. } = body.instructions[id.index()].op + { + literal(body, long, size as i64) + } else { + let method = MethodId::new(body.methods.len()); + body.methods.push(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: "locationStride".into(), + params: vec![object], + returns: long, + interface: false, + }); + let args = List::append(&mut body.args, [source[0]]); + append( + body, + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }, + long, + ) + }; + let (mul, scaled) = append( + body, + Op::Binary { + op: BinaryOp::Mul, + left: delta, + right: size, + }, + long, + ); + prefix.extend([constant, mul]); + delta = scaled; + } + ( + Op::Reinterpret(source[0]), + Op::Binary { + op: BinaryOp::Add, + left: source[1], + right: delta, + }, + ) + } + _ => unreachable!(), + }; + body.instructions[root.index()].op = root_op; + body.instructions[offset.index()].op = offset_op; + prefix.extend([root, offset]); + let values = known[original.result.unwrap().index()].unwrap(); + body.instructions[id.index()].op = Op::AddressPack(List::append(&mut body.args, values)); + prefixes.insert(id, prefix); + } + for (block, incoming) in predecessors.iter().enumerate() { + for &(_, edge) in incoming { + for &(_, position) in &joins[block] { + let input = body.edges[edge.index()].args[position]; + body.edges[edge.index()] + .args + .extend(known[input.index()].unwrap()); + } + } + } + // A decoded load and its later commit must use the same carrier. + // Separate carriers would lose the mutable view binding. + let bound_views = if shape == ComponentShape::StorageAddress { + body.instructions + .iter() + .filter_map(|inst| { + if let Op::Commit(pointer) = inst.op { + return known.get(pointer.index()).copied().flatten(); + } + let Op::Call { + method, + kind: CallKind::Virtual, + args, + } = inst.op + else { + return None; + }; + let method = &body.methods[method.index()]; + if method.owner != "org/rustlang/runtime/Pointer" + || method.name != "commitMemoryView" + { + return None; + } + known + .get(body.args[args.start as usize].index()) + .copied() + .flatten() + }) + .collect::>() + } else { + FxHashSet::default() + }; + for index in 0..body.instructions.len() { + let original = body.instructions[index]; + if shape == ComponentShape::StorageAddress { + if let Op::LoadFieldCopy { base, projection } = original.op + && let Some(parts) = known.get(base.index()).copied().flatten() + { + body.instructions[index].op = Op::LoadStorageFieldCopy { + address: List::append(&mut body.args, parts), + projection, + }; + continue; + } + let stored = match original.op { + Op::StoreField { + base, + projection, + value, + } => Some((base, projection, vec![value], false)), + Op::StoreFieldParts { + base, + projection, + parts, + } => Some((base, projection, body.args[parts.range()].to_vec(), true)), + _ => None, + }; + if let Some((base, projection, values, split)) = stored + && let Some(parts) = known.get(base.index()).copied().flatten() + { + body.instructions[index].op = Op::StoreStorageField { + args: List::append(&mut body.args, parts.into_iter().chain(values)), + projection, + split, + }; + continue; + } + let field = match original.op { + Op::LoadField { base, projection } => Some((base, projection, None)), + Op::LoadFieldPart { + base, + projection, + index, + } => Some((base, projection, Some(index))), + _ => None, + }; + if let Some((base, projection, index_part)) = field + && let Some(parts) = known.get(base.index()).copied().flatten() + { + body.instructions[index].op = Op::LoadStorageField { + address: List::append(&mut body.args, parts), + projection, + index: index_part, + }; + continue; + } + } + if let Op::AddressViewPart { + address, + index: part, + } = original.op + { + // Opaque helpers can retain a layout that differs from the pointer type, + // such as a DST tail or fixed-array cast. Normalize only proven scalar locations. + if shape != ComponentShape::Address + || known.get(address.index()).copied().flatten().is_none() + { + continue; + } + let mut prefix = Vec::new(); + let source = parts(body, types, address, &known, object, long, &mut prefix); + let Some(Type::Pointer(inner)) = types.get(body.value_type(address)) else { + unreachable!() + }; + let size = StorageSlot::scalar(inner, types).unwrap().size; + let int = types.scalar(ScalarType::I32); + let constant = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar( + Scalar::integer(ScalarType::I32, size.into()).unwrap(), + )); + let (instruction, size) = append(body, Op::Constant(constant), int); + prefix.push(instruction); + let method = MethodId::new(body.methods.len()); + body.methods.push(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: if part == 0 { + "locationSliceBacking" + } else { + "locationSliceOffset" + } + .into(), + params: vec![object, long, int], + returns: body.value_type(original.result.unwrap()), + interface: false, + }); + body.instructions[index].op = Op::Call { + method, + kind: CallKind::JvmStatic, + args: List::append(&mut body.args, [source[0], source[1], size]), + }; + prefixes.insert(InstId::new(index), prefix); + continue; + } + if let Op::AddressTag(value) = original.op { + if ComponentShape::of(types, body.value_type(value)) != Some(shape) { + continue; + } + let mut prefix = Vec::new(); + let source = parts(body, types, value, &known, object, long, &mut prefix); + body.instructions[index].op = Op::LocationTag(List::append(&mut body.args, source)); + prefixes.insert(InstId::new(index), prefix); + continue; + } + if let Op::AddressEqual { left, right } | Op::AddressCompare { left, right } = original.op { + // Leave storage comparisons for the later storage-address pass. + // Boxing components here would hide their storage layout. + if [left, right] + .into_iter() + .any(|value| ComponentShape::of(types, body.value_type(value)) != Some(shape)) + { + continue; + } + let mut prefix = Vec::new(); + let left = parts(body, types, left, &known, object, long, &mut prefix); + let right = parts(body, types, right, &known, object, long, &mut prefix); + let parts = List::append(&mut body.args, left.into_iter().chain(right)); + body.instructions[index].op = if matches!(original.op, Op::AddressCompare { .. }) { + Op::LocationCompare(parts) + } else { + Op::LocationEqual(parts) + }; + prefixes.insert(InstId::new(index), prefix); + continue; + } + if let Op::Adapt(source) = original.op { + if let Some(result) = original.result + && let Some(parts) = known.get(result.index()).copied().flatten() + { + body.instructions[index].op = if shape == ComponentShape::StorageAddress { + // Keep the JVM type refinement even though the storage identity + // and layout are already known. + Op::Refine(source) + } else if ComponentShape::of(types, body.value_type(result)) == Some(shape) { + Op::AddressPack(List::append(&mut body.args, parts)) + } else { + Op::Reinterpret(source) + }; + continue; + } + } + let pointer = match original.op { + Op::Load(p) + | Op::LoadCopy(p) + | Op::Store { pointer: p, .. } + | Op::AddressPart { address: p, .. } => p, + _ => continue, + }; + let Some(parts) = known.get(pointer.index()).copied().flatten() else { + continue; + }; + body.instructions[index].op = match original.op { + Op::AddressPart { index, .. } => Op::Reinterpret(parts[index as usize]), + Op::LoadCopy(_) if shape == ComponentShape::StorageAddress => { + Op::LoadAddressCopy(List::append(&mut body.args, parts)) + } + Op::Load(_) + if (shape == ComponentShape::StorageAddress && !bound_views.contains(&parts)) + || StorageSlot::scalar(body.value_type(original.result.unwrap()), types) + .is_some() => + { + Op::LoadAddress(List::append(&mut body.args, parts)) + } + Op::Store { value, .. } + if shape == ComponentShape::StorageAddress + || StorageSlot::scalar(body.value_type(value), types).is_some() => + { + Op::StoreAddress { + parts: List::append(&mut body.args, parts), + value, + } + } + op => op, + }; + } + let mut positions = debug.as_ref().map(|_| Vec::new()); + for block in &mut body.blocks { + let previous = std::mem::take(&mut block.instructions); + let mut mapping = Vec::new(); + for id in previous { + if positions.is_some() { + mapping.push(block.instructions.len() as u32); + } + if let Some(prefix) = prefixes.remove(&id) { + block.instructions.extend(prefix); + } + block.instructions.push(id); + } + if let Some(positions) = &mut positions { + mapping.push(block.instructions.len() as u32); + positions.push(mapping); + } + if let Some(Terminator::Invoke { inst, normal, .. }) = block.terminator { + if let Some(prefix) = prefixes.remove(&inst) { + block.instructions.extend(prefix); + } + if matches!(body.instructions[inst.index()].op, Op::AddressPack(_)) { + block.instructions.push(inst); + block.terminator = Some(Terminator::Jump(normal)); + } + } + } + if let (Some(debug), Some(positions)) = (debug, positions) { + for event in &mut debug.events { + event.position = positions[event.block.index()][event.position as usize]; + } + } +} diff --git a/compiler-core/src/opt/addresses_tests.rs b/compiler-core/src/opt/addresses_tests.rs new file mode 100644 index 00000000..cebd131e --- /dev/null +++ b/compiler-core/src/opt/addresses_tests.rs @@ -0,0 +1,544 @@ +use super::{decompose_addresses, lower_component_arguments}; +use crate::ir::*; +use crate::scalar::{BinaryOp, Scalar, ScalarType}; + +#[test] +fn borrowed_scalar_fields_keep_their_aggregate_root_across_calls() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let owner = types.symbol("Pair"); + let owner = types.intern(Type::Class(owner)); + let aggregate = types.intern(Type::Pointer(owner)); + let scalar = types.intern(Type::Pointer(int)); + let mut b = Builder::new(&types, int); + let root = b.parameter(b.current(), aggregate); + let field = b.field(FieldRef { + owner, + name: "second".into(), + ty: int, + is_static: false, + }); + let projection = b.projection(PointerProjection { + field, + offset: 4, + size: 4, + codec: None, + }); + let pointer = b + .emit( + Op::Project { + base: root, + projection, + }, + Some(scalar), + ) + .unwrap(); + let method = b.method(MethodRef { + owner: "Kernel".into(), + name: "consume".into(), + params: vec![scalar], + returns: int, + interface: false, + }); + let args = b.args([pointer]); + let result = b + .emit( + Op::Call { + method, + kind: CallKind::RustStatic, + args, + }, + Some(int), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| true, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(id, inst)| live.instructions[id] + && matches!(inst.op, Op::AddressPack(_) | Op::Project { .. })) + ); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::ProjectRoot { .. })) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +#[test] +fn thin_pointer_comparisons_use_components() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let boolean = types.scalar(ScalarType::Bool); + let owner = types.symbol("Item"); + let owner = types.intern(Type::Class(owner)); + for (pointee, ordering) in [int, owner] + .into_iter() + .flat_map(|p| [(p, false), (p, true)]) + { + let result_type = if ordering { int } else { boolean }; + let pointer = types.intern(Type::Pointer(pointee)); + let mut b = Builder::new(&types, result_type); + let left = b.parameter(b.current(), pointer); + let right = b.parameter(b.current(), pointer); + let equal = b + .emit( + if ordering { + Op::AddressCompare { left, right } + } else { + Op::AddressEqual { left, right } + }, + Some(result_type), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(equal))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| false, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + assert!(body.instructions.iter().any(|i| if ordering { + matches!(i.op, Op::LocationCompare(_)) + } else { + matches!(i.op, Op::LocationEqual(_)) + })); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(index, i)| live.instructions[index] && matches!(i.op, Op::AddressPack(_))) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} + +#[test] +fn runtime_pointer_annotations_do_not_force_materialization() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let boolean = types.scalar(ScalarType::Bool); + let pointer = types.intern(Type::Pointer(int)); + let runtime = types.symbol("org/rustlang/runtime/Pointer"); + let runtime = types.intern(Type::Class(runtime)); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let mut b = Builder::new(&types, boolean); + let left = b.parameter(b.current(), pointer); + let right = b.parameter(b.current(), pointer); + // Runtime helper signatures and source bindings both erase annotations. + let erased = b.emit(Op::Reinterpret(left), Some(runtime)).unwrap(); + let erased = b.emit(Op::Reinterpret(erased), Some(object)).unwrap(); + let recovered = b.emit(Op::Adapt(erased), Some(pointer)).unwrap(); + let equal = b + .emit( + Op::AddressEqual { + left: recovered, + right, + }, + Some(boolean), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(equal))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| false, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!(inst.op, Op::AddressPack(_) | Op::Adapt(_))) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +#[test] +fn scalar_reference_loop_has_no_materialized_addresses() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let long = types.scalar(ScalarType::I64); + let boolean = types.scalar(ScalarType::Bool); + let pointer = types.intern(Type::Pointer(int)); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let mut b = Builder::new(&types, int); + let start = b.parameter(b.current(), pointer); + let end = b.parameter(b.current(), long); + let address = b.variable(object); + let erased = b.emit(Op::Reinterpret(start), Some(object)).unwrap(); + b.define(address, erased); + let index = b.variable(long); + let zero = b.constant(long, Scalar::integer(ScalarType::I64, 0).unwrap()); + b.define(index, zero); + let header = b.create_block(); + let step = b.create_block(); + let done = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + let i = b.read(index); + let current = b.read(address); + let current = b.emit(Op::Adapt(current), Some(pointer)).unwrap(); + let condition = b + .emit( + Op::Binary { + op: BinaryOp::Lt, + left: i, + right: end, + }, + Some(boolean), + ) + .unwrap(); + b.branch(condition, step, done); + b.switch_to(step); + let value = b.emit(Op::Load(current), Some(int)).unwrap(); + b.emit( + Op::Store { + pointer: current, + value, + }, + None, + ); + let one = b.constant(long, Scalar::integer(ScalarType::I64, 1).unwrap()); + let next = b + .emit( + Op::Offset { + pointer: current, + offset: one, + bytes: false, + wrapping: false, + }, + Some(pointer), + ) + .unwrap(); + let erased = b.emit(Op::Reinterpret(next), Some(object)).unwrap(); + b.define(address, erased); + let next = b + .emit( + Op::Binary { + op: BinaryOp::Add, + left: i, + right: one, + }, + Some(long), + ) + .unwrap(); + b.define(index, next); + b.jump(header, vec![]); + b.switch_to(done); + let value = b.emit(Op::Load(current), Some(int)).unwrap(); + b.terminate(Terminator::Return(Some(value))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| true, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + for (i, inst) in body.instructions.iter().enumerate() { + assert!( + !live.instructions[i] + || !matches!( + inst.op, + Op::AddressPack(_) | Op::Offset { .. } | Op::Load(_) | Op::Store { .. } + ) + ); + } + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!( + !code + .instructions + .iter() + .any(|i| matches!(i, crate::classfile::attributes::Instruction::New(_))) + ); +} + +#[test] +fn aggregate_addresses_keep_layout_roots_through_calls_and_field_access() { + let mut types = Types::default(); + let long = types.scalar(ScalarType::I64); + let owner = types.symbol("Pair"); + let owner = types.intern(Type::Class(owner)); + let pointer = types.intern(Type::Pointer(owner)); + let unit = types.intern(Type::Unit); + let mut b = Builder::new(&types, long); + let base = b.parameter(b.current(), pointer); + let offset = b.parameter(b.current(), long); + let field = b.field(FieldRef { + owner, + name: "value".into(), + ty: long, + is_static: false, + }); + let projection = b.projection(PointerProjection { + field, + offset: 0, + size: 8, + codec: None, + }); + let address = b + .emit( + Op::Offset { + pointer: base, + offset, + bytes: false, + wrapping: false, + }, + Some(pointer), + ) + .unwrap(); + let method = b.method(MethodRef { + owner: "test/Calls".into(), + name: "inspect".into(), + params: vec![pointer], + returns: unit, + interface: false, + }); + let args = b.args([address]); + b.emit( + Op::Call { + method, + kind: CallKind::RustStatic, + args, + }, + None, + ); + let value = b + .emit( + Op::LoadField { + base: address, + projection, + }, + Some(long), + ) + .unwrap(); + b.emit( + Op::StoreField { + base: address, + projection, + value, + }, + None, + ); + b.terminate(Terminator::Return(Some(value))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| true, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!( + inst.op, + Op::AddressPack(_) + | Op::AddressPart { .. } + | Op::Offset { .. } + | Op::LoadField { .. } + | Op::StoreField { .. } + )) + ); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::LoadStorageField { .. })) + ); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::StoreStorageField { .. })) + ); + assert!(body.methods.iter().any(|m| m.name == "locationStride")); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +#[test] +fn decoded_aggregate_writeback_retains_its_binding_owner() { + for (commit, owned) in [(false, false), (true, false), (false, true), (true, true)] { + let mut types = Types::default(); + let unit = types.intern(Type::Unit); + let long = types.scalar(ScalarType::I64); + let name = types.symbol("Pair"); + let object = types.intern(Type::Class(name)); + let pointer = types.intern(Type::Pointer(object)); + let name = types.symbol("java/lang/Object"); + let erased = types.intern(Type::Class(name)); + let mut b = Builder::new(&types, object); + let root = b.parameter(b.current(), pointer); + let displacement = b.parameter(b.current(), long); + let address = b + .emit( + Op::Offset { + pointer: root, + offset: displacement, + bytes: true, + wrapping: false, + }, + Some(pointer), + ) + .unwrap(); + let erased_address = b.emit(Op::Reinterpret(address), Some(erased)).unwrap(); + let load_address = b.emit(Op::Adapt(erased_address), Some(pointer)).unwrap(); + let value = b + .emit( + if owned { + Op::LoadCopy(load_address) + } else { + Op::Load(load_address) + }, + Some(object), + ) + .unwrap(); + if commit { + let method = b.method(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: "commitMemoryView".into(), + params: vec![], + returns: unit, + interface: false, + }); + let commit_address = b.emit(Op::Adapt(erased_address), Some(pointer)).unwrap(); + let args = b.args([commit_address]); + b.emit( + Op::Call { + method, + kind: CallKind::Virtual, + args, + }, + None, + ); + } + b.terminate(Terminator::Return(Some(value))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| false, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + assert_eq!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::Load(_))), + commit && !owned + ); + assert_eq!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::LoadAddress(_))), + !commit && !owned + ); + assert_eq!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::LoadAddressCopy(_))), + owned + ); + if commit && !owned { + let live = super::live(&body, &types); + assert_eq!( + body.instructions + .iter() + .enumerate() + .filter(|(i, inst)| { + live.instructions[*i] && matches!(inst.op, Op::AddressPack(_)) + }) + .count(), + 1, + "load and writeback must share one materialized owner" + ); + } + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} + +#[test] +fn nullable_discriminants_do_not_materialize_borrowed_addresses() { + let mut types = Types::default(); + let scalar = types.scalar(ScalarType::I32); + let owner = types.symbol("Item"); + let owner = types.intern(Type::Class(owner)); + let long = types.scalar(ScalarType::I64); + for pointee in [scalar, owner] { + let pointer = types.intern(Type::Pointer(pointee)); + let mut b = Builder::new(&types, long); + let value = b.parameter(b.current(), pointer); + let tag = b.emit(Op::AddressTag(value), Some(long)).unwrap(); + b.terminate(Terminator::Return(Some(tag))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| false, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::LocationTag(_))) + ); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!(inst.op, Op::AddressPack(_) | Op::AddressTag(_))) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} + +#[test] +fn nullable_return_joins_keep_address_components() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let boolean = types.scalar(ScalarType::Bool); + let owner = types.symbol("Item"); + let owner = types.intern(Type::Class(owner)); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + for pointee in [int, owner] { + let pointer = types.intern(Type::Pointer(pointee)); + let mut b = Builder::new(&types, pointer); + let condition = b.parameter(b.current(), boolean); + let value = b.parameter(b.current(), pointer); + let result = b.variable(object); + let some = b.create_block(); + let none = b.create_block(); + let done = b.create_block(); + b.branch(condition, some, none); + b.switch_to(some); + let erased = b.emit(Op::Reinterpret(value), Some(object)).unwrap(); + b.define(result, erased); + b.jump(done, vec![]); + b.switch_to(none); + let constant = ConstId::new(b.body.constants.len()); + b.body.constants.push(Constant::Null(pointer)); + let null = b.emit(Op::Constant(constant), Some(pointer)).unwrap(); + let erased = b.emit(Op::Reinterpret(null), Some(object)).unwrap(); + b.define(result, erased); + b.jump(done, vec![]); + b.switch_to(done); + let result = b.read(result); + let result = b.emit(Op::Adapt(result), Some(pointer)).unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| false, None); + super::lower_component_returns(&mut body, &mut types, true, |_| false, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + assert!(!body.instructions.iter().enumerate().any(|(i, inst)| { + live.instructions[i] && matches!(inst.op, Op::AddressPack(_) | Op::AddressPart { .. }) + })); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} diff --git a/compiler-core/src/opt/field_abi.rs b/compiler-core/src/opt/field_abi.rs new file mode 100644 index 00000000..7b3981ab --- /dev/null +++ b/compiler-core/src/opt/field_abi.rs @@ -0,0 +1,213 @@ +//! Split stored borrowed values using the same shapes as parameters and SSA joins. +use crate::ir::*; +use crate::jvm::abi::{address_displacement_name, view_field_names}; +use rustc_hash::FxHashMap; + +fn append( + body: &mut Body, + op: Op, + ty: Option, + prefix: &mut Vec, +) -> Option { + let id = InstId::new(body.instructions.len()); + let result = ty.map(|ty| { + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty, + def: ValueDef::Inst(id), + }); + value + }); + body.instructions.push(Inst { op, result }); + prefix.push(id); + result +} + +pub fn lower_borrowed_fields( + body: &mut Body, + types: &mut Types, + split: impl Fn(&str) -> bool, + debug: Option<&mut DebugInfo>, +) { + let constructors = body + .methods + .iter() + .map(|method| method.name == "" && split(&method.owner)) + .collect::>(); + let fields = (0..body.fields.len()) + .map(|index| { + let field = body.fields[index].clone(); + let Type::Class(owner) = types.get(field.owner)? else { + return None; + }; + if field.is_static || !split(types.symbol_name(owner)?) { + return None; + } + let shape = ComponentShape::of(types, field.ty)?; + let names = match types.get(field.ty)? { + Type::TaggedI64 => crate::jvm::abi::tagged_field_names(&field.name).to_vec(), + Type::Pointer(_) => vec![ + field.name.clone(), + address_displacement_name(types, field.ty, &field.name), + ], + ty => view_field_names(&field.name, matches!(ty, Type::Str)).to_vec(), + }; + let members = names + .into_iter() + .zip(shape.parts(types)) + .map(|(name, ty)| { + let id = MemberId::new(body.fields.len()); + body.fields.push(FieldRef { + name, + ty, + ..field.clone() + }); + id + }) + .collect::>(); + Some((shape, members)) + }) + .collect::>(); + if fields.iter().all(Option::is_none) && !constructors.iter().any(|&v| v) { + return; + } + let mut prefixes = FxHashMap::default(); + for index in 0..body.instructions.len() { + let mut prefix = Vec::new(); + let original = body.instructions[index].op; + let replacement = match original { + Op::GetField { .. } + | Op::SetField { .. } + | Op::LoadField { .. } + | Op::StoreField { .. } => { + // Projected fields carry a projection ID instead of a member ID. + let member = match original { + Op::LoadField { projection, .. } | Op::StoreField { projection, .. } => { + body.projections[projection.index()].field + } + Op::GetField { field, .. } | Op::SetField { field, .. } => field, + _ => unreachable!(), + }; + let Some((shape, physical)) = fields.get(member.index()).and_then(Option::as_ref) + else { + continue; + }; + let mut values = Vec::new(); + for (index, &field) in physical.iter().enumerate() { + let ty = body.fields[field.index()].ty; + let op = match original { + Op::GetField { object, .. } => Op::GetField { object, field }, + Op::LoadField { base, projection } => Op::LoadFieldPart { + base, + projection, + index: index as u8, + }, + Op::SetField { value, .. } | Op::StoreField { value, .. } => { + shape.part(value, index as u8) + } + _ => unreachable!(), + }; + let value = append(body, op, Some(ty), &mut prefix).unwrap(); + values.push(value); + if let Op::SetField { object, .. } = original { + append( + body, + Op::SetField { + object, + field, + value, + }, + None, + &mut prefix, + ); + } + } + match original { + Op::SetField { .. } => Op::Nop, + Op::StoreField { + base, projection, .. + } => Op::StoreFieldParts { + base, + projection, + parts: List::append(&mut body.args, values), + }, + _ => shape.pack(List::append(&mut body.args, values)), + } + } + Op::Call { + method, + kind: CallKind::Constructor, + args, + } if constructors[method.index()] => { + let original = body.args[args.range()].to_vec(); + let params = body.methods[method.index()].params.clone(); + let mut values = Vec::new(); + for (value, ty) in original.into_iter().zip(params) { + if let Some(shape) = ComponentShape::of(types, ty) { + for (index, ty) in shape.parts(types).enumerate() { + values.push( + append(body, shape.part(value, index as u8), Some(ty), &mut prefix) + .unwrap(), + ); + } + } else { + values.push(value); + } + } + Op::Call { + method, + kind: CallKind::Constructor, + args: List::append(&mut body.args, values), + } + } + _ => continue, + }; + body.instructions[index].op = replacement; + prefixes.insert(InstId::new(index), prefix); + } + for (method, split) in body.methods.iter_mut().zip(constructors) { + if split { + let previous = std::mem::take(&mut method.params); + for ty in previous { + if let Some(shape) = ComponentShape::of(types, ty) { + method.params.extend(shape.parts(types)); + } else { + method.params.push(ty); + } + } + } + } + let mut positions = debug.as_ref().map(|_| Vec::new()); + for block in &mut body.blocks { + let previous = std::mem::take(&mut block.instructions); + let mut mapping = Vec::new(); + for id in previous { + mapping.push(block.instructions.len() as u32); + if let Some(prefix) = prefixes.remove(&id) { + block.instructions.extend(prefix); + } + block.instructions.push(id); + } + mapping.push(block.instructions.len() as u32); + if let Some(positions) = &mut positions { + positions.push(mapping); + } + if let Some(Terminator::Invoke { inst, normal, .. }) = block.terminator { + if let Some(prefix) = prefixes.remove(&inst) { + block.instructions.extend(prefix); + } + if matches!( + body.instructions[inst.index()].op, + Op::Nop | Op::AddressPack(_) | Op::ViewPack(_) | Op::TaggedPack(_) + ) { + block.instructions.push(inst); + block.terminator = Some(Terminator::Jump(normal)); + } + } + } + if let (Some(debug), Some(positions)) = (debug, positions) { + for event in &mut debug.events { + event.position = positions[event.block.index()][event.position as usize]; + } + } +} diff --git a/compiler-core/src/opt/field_abi_tests.rs b/compiler-core/src/opt/field_abi_tests.rs new file mode 100644 index 00000000..d66bc8ea --- /dev/null +++ b/compiler-core/src/opt/field_abi_tests.rs @@ -0,0 +1,290 @@ +use crate::ir::*; +use crate::scalar::ScalarType; + +#[test] +fn uninitialized_view_fields_use_default_components() { + for utf8 in [false, true] { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let view = types.intern(if utf8 { Type::Str } else { Type::Slice(byte) }); + let name = types.symbol("test/Partial"); + let owner = types.intern(Type::Class(name)); + let void = types.intern(Type::Unit); + let mut b = Builder::new(&types, owner); + let null = ConstId::new(b.body.constants.len()); + b.body.constants.push(Constant::Null(view)); + let initial = b.emit(Op::Constant(null), Some(view)).unwrap(); + let constructor = b.method(MethodRef { + owner: "test/Partial".into(), + name: "".into(), + params: vec![view], + returns: void, + interface: false, + }); + let args = b.args([initial]); + let result = b + .emit( + Op::Call { + method: constructor, + kind: CallKind::Constructor, + args, + }, + Some(owner), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + super::lower_borrowed_fields(&mut body, &mut types, |_| true, None); + super::decompose_views(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!( + !code + .instructions + .iter() + .any(|i| matches!(i, crate::classfile::attributes::Instruction::Getfield(_))) + ); + } +} + +#[test] +fn stored_scalar_borrows_load_and_store_without_pointer_carriers() { + check_stored_borrow(false); + check_stored_borrow(true); +} + +#[test] +fn stored_views_have_no_live_view_carriers() { + for utf8 in [false, true] { + for indirect in [false, true] { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let len = types.scalar(ScalarType::U64); + let view = types.intern(if utf8 { Type::Str } else { Type::Slice(byte) }); + let name = types.symbol("test/Views"); + let owner = types.intern(Type::Class(name)); + let borrowed = types.intern(Type::Pointer(owner)); + let mut b = Builder::new(&types, len); + let object = b.parameter(b.current(), if indirect { borrowed } else { owner }); + let value = b.parameter(b.current(), view); + let field = b.field(FieldRef { + owner, + name: "value".into(), + ty: view, + is_static: false, + }); + let loaded = if indirect { + let projection = b.projection(PointerProjection { + field, + offset: 0, + size: 16, + codec: None, + }); + b.emit( + Op::StoreField { + base: object, + projection, + value, + }, + None, + ); + b.emit( + Op::LoadField { + base: object, + projection, + }, + Some(view), + ) + .unwrap() + } else { + b.emit( + Op::SetField { + object, + field, + value, + }, + None, + ); + b.emit(Op::GetField { object, field }, Some(view)).unwrap() + }; + let result = b.emit(Op::Length(loaded), Some(len)).unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + super::lower_component_arguments(&mut body, &mut types, true, |_| false, None); + super::lower_borrowed_fields(&mut body, &mut types, |_| true, None); + super::decompose_views(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(index, i)| live.instructions[index] + && matches!(i.op, Op::ViewPack(_) | Op::ViewPart { .. } | Op::Length(_))) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } + } +} + +fn check_stored_borrow(indirect: bool) { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let pointer = types.intern(Type::Pointer(int)); + let name = types.symbol("test/References"); + let owner = types.intern(Type::Class(name)); + let borrowed = types.intern(Type::Pointer(owner)); + let mut b = Builder::new(&types, int); + let object = b.parameter(b.current(), if indirect { borrowed } else { owner }); + let incoming = b.parameter(b.current(), pointer); + let field = b.field(FieldRef { + owner, + name: "value".into(), + ty: pointer, + is_static: false, + }); + let loaded = if indirect { + let projection = b.projection(PointerProjection { + field, + offset: 0, + size: 8, + codec: None, + }); + b.emit( + Op::StoreField { + base: object, + projection, + value: incoming, + }, + None, + ); + b.emit( + Op::LoadField { + base: object, + projection, + }, + Some(pointer), + ) + .unwrap() + } else { + b.emit( + Op::SetField { + object, + field, + value: incoming, + }, + None, + ); + b.emit(Op::GetField { object, field }, Some(pointer)) + .unwrap() + }; + let read = b.emit(Op::Load(loaded), Some(int)).unwrap(); + b.terminate(Terminator::Return(Some(read))); + let mut body = b.finish().unwrap(); + super::lower_component_arguments(&mut body, &mut types, true, |_| false, None); + super::lower_borrowed_fields( + &mut body, + &mut types, + |owner| owner == "test/References", + None, + ); + super::decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = super::live(&body, &types); + for (index, inst) in body.instructions.iter().enumerate() { + if live.instructions[index] { + assert!(!matches!( + inst.op, + Op::AddressPack(_) | Op::AddressPart { .. } | Op::Load(_) + )); + if let Op::GetField { field, .. } | Op::SetField { field, .. } = inst.op { + assert_ne!(body.fields[field.index()].ty, pointer); + } + } + } + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +#[test] +fn split_field_components_share_resolution_only_before_the_next_effect() { + use crate::classfile::{ + Constant, attributes::Instruction, constant_pool::InternedConstantPool, + }; + for intervening_call in [false, true] { + let mut types = Types::default(); + let unit = types.intern(Type::Unit); + let byte = types.scalar(ScalarType::U8); + let view = types.intern(Type::Slice(byte)); + let owner_name = types.symbol("test/Views"); + let owner = types.intern(Type::Class(owner_name)); + let pointer = types.intern(Type::Pointer(owner)); + let parts = ComponentShape::View.parts(&mut types).collect::>(); + let mut b = Builder::new(&types, unit); + let base = b.parameter(b.current(), pointer); + let field = b.field(FieldRef { + owner, + name: "view".into(), + ty: view, + is_static: false, + }); + let projection = b.projection(PointerProjection { + field, + offset: 0, + size: 16, + codec: None, + }); + for (index, &ty) in parts.iter().enumerate() { + if intervening_call && index == 1 { + let method = b.method(MethodRef { + owner: "test/External".into(), + name: "replaceStorage".into(), + params: vec![pointer], + returns: unit, + interface: false, + }); + let args = b.args([base]); + b.emit( + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }, + None, + ); + } + b.emit( + Op::LoadFieldPart { + base, + projection, + index: index as u8, + }, + Some(ty), + ); + } + b.terminate(Terminator::Return(None)); + let body = b.finish().unwrap(); + let mut cp = InternedConstantPool::default(); + let code = crate::jvm::select::compile(&body, &types, &mut cp).unwrap(); + let pool = cp.into_inner(); + let resolutions = code + .instructions + .iter() + .filter(|instruction| { + let Instruction::Invokevirtual(index) = instruction else { + return false; + }; + let Some(Constant::MethodRef { + name_and_type_index, + .. + }) = pool.get(*index) + else { + return false; + }; + let (name, _) = pool.try_get_name_and_type(*name_and_type_index).unwrap(); + pool.try_get_utf8(*name).unwrap() == "directAggregate" + }) + .count(); + assert_eq!(resolutions, if intervening_call { 2 } else { 1 }); + } +} diff --git a/compiler-core/src/opt/mod.rs b/compiler-core/src/opt/mod.rs index 677c2573..78c857fd 100644 --- a/compiler-core/src/opt/mod.rs +++ b/compiler-core/src/opt/mod.rs @@ -1,5 +1,9 @@ //! Analyses are built on demand and owned by one body compilation. +mod field_abi; mod fields; +pub use field_abi::lower_borrowed_fields; +#[cfg(test)] +mod field_abi_tests; mod live; pub use fields::promote_fields; #[cfg(test)] @@ -21,6 +25,20 @@ mod cells_tests; mod view_abi; pub use view_abi::{component_argument_slots, lower_component_arguments}; +mod addresses; +pub use addresses::decompose_addresses; + +#[cfg(test)] +mod addresses_tests; + +mod return_abi; +pub use return_abi::lower_component_returns; + +#[cfg(test)] +mod return_abi_tests; + +mod address_parts; + mod simplify; mod unreachable; pub use simplify::simplify_components; diff --git a/compiler-core/src/opt/return_abi.rs b/compiler-core/src/opt/return_abi.rs new file mode 100644 index 00000000..3b8ff9d0 --- /dev/null +++ b/compiler-core/src/opt/return_abi.rs @@ -0,0 +1,304 @@ +//! Return borrowed roots directly. Write scalar metadata into caller-owned scratch. +//! Each frame reuses one scratch array and captures successful results into SSA. +//! Other frames and Rust code cannot access this scratch. +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; +use rustc_hash::FxHashMap; + +fn emit(body: &mut Body, op: Op, ty: Option, into: &mut Vec) -> Option { + let inst = InstId::new(body.instructions.len()); + let result = ty.map(|ty| { + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty, + def: ValueDef::Inst(inst), + }); + value + }); + body.instructions.push(Inst { op, result }); + into.push(inst); + result +} +fn integer(body: &mut Body, ty: TypeId, bits: u128, into: &mut Vec) -> ValueId { + let constant = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar( + Scalar::integer(ScalarType::I32, bits).unwrap(), + )); + emit(body, Op::Constant(constant), Some(ty), into).unwrap() +} + +/// Bypass scratch copies only for the last call that writes metadata. +/// Aliases do not prove this condition because a later call can overwrite scratch. +fn forwards_last_call( + body: &Body, + predecessors: &[Vec<(BlockId, EdgeId)>], + mut block: usize, + wanted: InstId, + first: &[InstId], + calls: &FxHashMap, +) -> bool { + let mut budget = 64usize; + for step in 0..16 { + let current = &body.blocks[block]; + let invoke = match current.terminator { + Some(Terminator::Invoke { inst, .. }) => Some(inst), + _ => None, + }; + let instructions = if step == 0 { + first + } else { + ¤t.instructions + }; + for inst in invoke.into_iter().chain(instructions.iter().rev().copied()) { + if inst == wanted { + return true; + } + if calls.contains_key(&inst) || budget == 0 { + return false; + } + budget -= 1; + } + let [incoming] = predecessors[block].as_slice() else { + return false; + }; + block = incoming.0.index(); + } + false +} + +pub fn lower_component_returns( + body: &mut Body, + types: &mut Types, + entry: bool, + select: impl Fn(&MethodRef) -> bool, + debug: Option<&mut DebugInfo>, +) { + let entry = entry + .then(|| ComponentShape::of(types, body.return_type)) + .flatten(); + let methods = body + .methods + .iter() + .map(|method| { + (select(method) + && super::component_argument_slots(types, method.params.iter().copied()) < 254) + .then(|| ComponentShape::of(types, method.returns)) + .flatten() + }) + .collect::>(); + let has_calls = body.instructions.iter().any(|inst| { + matches!(inst.op, + Op::Call { method, .. } if methods[method.index()].is_some()) + }); + if entry.is_none() && methods.iter().all(Option::is_none) { + return; + } + + let long = types.scalar(ScalarType::I64); + let int = types.scalar(ScalarType::I32); + let metadata = types.intern(Type::Array(long)); + if entry.is_none() && !has_calls { + // Function addresses and handles use the same physical descriptor, + // including targets that this body never calls directly. + for (method, shape) in body.methods.iter_mut().zip(methods) { + if let Some(shape) = shape { + method.params.push(metadata); + method.returns = shape.parts(types).next().unwrap(); + } + } + return; + } + + let mut prologue = Vec::new(); + let scratch = if entry.is_some() { + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty: metadata, + def: ValueDef::Param(body.entry), + }); + body.blocks[body.entry.index()].params.push(value); + body.return_type = entry.unwrap().parts(types).next().unwrap(); + value + } else { + let length = integer(body, int, 2, &mut prologue); + emit(body, Op::NewArray(length), Some(metadata), &mut prologue).unwrap() + }; + let indices = [0, 1].map(|index| integer(body, int, index, &mut prologue)); + let count = body.instructions.len(); + let predecessors = body.predecessors(); + let mut call_results = FxHashMap::default(); + let mut calls = FxHashMap::default(); + let mut suffixes = FxHashMap::default(); + for index in 0..count { + let Inst { + op: Op::Call { method, kind, args }, + result, + } = body.instructions[index] + else { + continue; + }; + let Some(shape) = methods[method.index()] else { + continue; + }; + calls.insert(InstId::new(index), ()); + let mut arguments = body.args[args.range()].to_vec(); + arguments.push(scratch); + body.instructions[index].op = Op::Call { + method, + kind, + args: List::append(&mut body.args, arguments), + }; + if let Some(value) = result { + let root = ValueId::new(body.values.len()); + body.values.push(Value { + ty: shape.parts(types).next().unwrap(), + def: ValueDef::Inst(InstId::new(index)), + }); + body.instructions[index].result = Some(root); + call_results.insert(value, (InstId::new(index), root, shape)); + let mut suffix = Vec::new(); + let mut parts = vec![root]; + for (position, ty) in shape.parts(types).skip(1).enumerate() { + let value = emit( + body, + Op::ArrayGet { + native: true, + array: scratch, + index: indices[position], + }, + Some(long), + &mut suffix, + ) + .unwrap(); + let value = if ty == long { + value + } else { + emit(body, Op::Cast(value), Some(ty), &mut suffix).unwrap() + }; + parts.push(value); + } + let inst = InstId::new(body.instructions.len()); + body.values[value.index()].def = ValueDef::Inst(inst); + body.instructions.push(Inst { + op: shape.pack(List::append(&mut body.args, parts)), + result: Some(value), + }); + suffix.push(inst); + suffixes.insert(InstId::new(index), suffix); + } + } + for (method, shape) in body.methods.iter_mut().zip(methods) { + if let Some(shape) = shape { + method.params.push(metadata); + method.returns = shape.parts(types).next().unwrap(); + } + } + let original_blocks = body.blocks.len(); + let mut positions = debug.as_ref().map(|_| Vec::new()); + for index in 0..original_blocks { + let previous = std::mem::take(&mut body.blocks[index].instructions); + let mut instructions = Vec::new(); + if index == body.entry.index() { + instructions.append(&mut prologue); + } + let mut mapping = Vec::new(); + for id in previous { + mapping.push(instructions.len() as u32); + instructions.push(id); + if let Some(suffix) = suffixes.remove(&id) { + instructions.extend(suffix); + } + } + mapping.push(instructions.len() as u32); + if let Some(positions) = &mut positions { + positions.push(mapping); + } + match body.blocks[index].terminator { + Some(Terminator::Return(Some(value))) if entry.is_some() => { + let shape = entry.unwrap(); + let mut source = body.resolve(value); + for _ in 0..16 { + let ValueDef::Inst(inst) = body.values[source.index()].def else { + break; + }; + match body.instructions[inst.index()].op { + Op::Reinterpret(next) | Op::Refine(next) => source = body.resolve(next), + _ => break, + } + } + if let Some(&(call, root, returned)) = call_results.get(&source) + && shape == returned + && forwards_last_call(body, &predecessors, index, call, &instructions, &calls) + { + body.blocks[index].terminator = Some(Terminator::Return(Some(root))); + body.blocks[index].instructions = instructions; + continue; + } + let mut root = None; + for (position, ty) in shape.parts(types).enumerate() { + let part = emit( + body, + shape.part(value, position as u8), + Some(ty), + &mut instructions, + ) + .unwrap(); + if position == 0 { + root = Some(part); + continue; + } + let part = if ty == long { + part + } else { + emit(body, Op::Cast(part), Some(long), &mut instructions).unwrap() + }; + emit( + body, + Op::ArraySet { + native: true, + array: scratch, + index: indices[position - 1], + value: part, + }, + None, + &mut instructions, + ); + } + body.blocks[index].terminator = Some(Terminator::Return(root)); + } + Some(Terminator::Invoke { + inst, + normal, + unwind, + }) => { + if let Some(suffix) = suffixes.remove(&inst) { + // A throw leaves the caller's previous values intact. + // Only the successful edge can consume returned metadata. + let continuation = BlockId::new(body.blocks.len()); + body.blocks.push(Block { + params: Vec::new(), + instructions: suffix, + terminator: Some(Terminator::Jump(normal)), + }); + let edge = EdgeId::new(body.edges.len()); + body.edges.push(Edge { + target: continuation, + args: Vec::new(), + }); + body.blocks[index].terminator = Some(Terminator::Invoke { + inst, + normal: edge, + unwind, + }); + } + } + _ => {} + } + body.blocks[index].instructions = instructions; + } + if let (Some(debug), Some(positions)) = (debug, positions) { + for event in &mut debug.events { + event.position = positions[event.block.index()][event.position as usize]; + } + } +} diff --git a/compiler-core/src/opt/return_abi_tests.rs b/compiler-core/src/opt/return_abi_tests.rs new file mode 100644 index 00000000..54723f19 --- /dev/null +++ b/compiler-core/src/opt/return_abi_tests.rs @@ -0,0 +1,154 @@ +use super::*; +use crate::ir::*; +use crate::scalar::ScalarType; + +#[test] +fn borrowed_return_components_cross_calls_and_unwind_edges() { + for unwind in [false, true] { + for returns_borrow in [false, true] { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let long = types.scalar(ScalarType::U64); + let signed = types.scalar(ScalarType::I64); + let pointer = types.intern(Type::Pointer(int)); + let slice = types.intern(Type::Slice(int)); + let object = types.symbol("Pair"); + let object = types.intern(Type::Class(object)); + let storage = types.intern(Type::Pointer(object)); + let tagged = types.intern(Type::TaggedI64); + for borrowed in [pointer, slice, storage, tagged] { + let output = if returns_borrow { borrowed } else { long }; + let mut b = Builder::new(&types, output); + let input = b.parameter(b.current(), borrowed); + let method = b.method(MethodRef { + owner: "Leaf".into(), + name: "pick".into(), + params: vec![borrowed], + returns: borrowed, + interface: false, + }); + let args = b.args([input]); + let op = Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }; + let result = if unwind { + let handler = b.create_block(); + let value = b.invoke(op, Some(borrowed), handler).unwrap(); + let continuation = b.current(); + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + b.switch_to(continuation); + value + } else { + b.emit(op, Some(borrowed)).unwrap() + }; + let result = if returns_borrow { + result + } else if borrowed == tagged { + let tag = b + .emit( + Op::TaggedPart { + value: result, + index: 1, + }, + Some(signed), + ) + .unwrap(); + b.emit(Op::Cast(tag), Some(long)).unwrap() + } else if borrowed == slice { + b.emit(Op::Length(result), Some(long)).unwrap() + } else { + let tag = b.emit(Op::AddressTag(result), Some(signed)).unwrap(); + b.emit(Op::Cast(tag), Some(long)).unwrap() + }; + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| true, None); + lower_component_returns(&mut body, &mut types, true, |_| true, None); + verify(&body, &types).unwrap(); + decompose_tagged(&mut body, &types, None); + decompose_views(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + assert!(!body.instructions.iter().enumerate().any( + |(index, inst)| live.instructions[index] + && matches!( + inst.op, + Op::AddressPack(_) | Op::ViewPack(_) | Op::TaggedPack(_) + ) + )); + assert_eq!( + body.instructions + .iter() + .filter(|i| matches!(i.op, Op::NewArray(_))) + .count(), + usize::from(!returns_borrow) + ); + assert_eq!( + body.blocks + .iter() + .filter(|b| matches!(b.terminator, Some(Terminator::Invoke { .. }))) + .count(), + usize::from(unwind) + ); + let code = + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + use crate::classfile::attributes::Instruction; + assert_eq!( + code.instructions.contains(&Instruction::Laload), + !returns_borrow + ); + assert_eq!(code.instructions.contains(&Instruction::Lastore), false); + } + } + } +} + +#[test] +fn earlier_return_metadata_survives_a_later_component_call() { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let slice = types.intern(Type::Slice(byte)); + let mut b = Builder::new(&types, slice); + let input = b.parameter(b.current(), slice); + let method = b.method(MethodRef { + owner: "Leaf".into(), + name: "view".into(), + params: vec![slice], + returns: slice, + interface: false, + }); + let args = b.args([input]); + let op = Op::Call { + method, + kind: CallKind::RustStatic, + args, + }; + let first = b.emit(op, Some(slice)).unwrap(); + b.emit(op, Some(slice)); + b.terminate(Terminator::Return(Some(first))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| true, None); + lower_component_returns(&mut body, &mut types, true, |_| true, None); + decompose_views(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + use crate::classfile::attributes::Instruction; + assert_eq!( + code.instructions + .iter() + .filter(|i| **i == Instruction::Laload) + .count(), + 2 + ); + assert_eq!( + code.instructions + .iter() + .filter(|i| **i == Instruction::Lastore) + .count(), + 2 + ); +} From 7454ca0b00e90c8d0f7fdb367441eab01b5a146d Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 06:14:32 +1000 Subject: [PATCH 18/61] access aggregate fields directly --- compiler-core/src/opt/fields.rs | 58 ++++- compiler-core/src/opt/fields_tests.rs | 303 ++++++++++++++++++++++++++ 2 files changed, 354 insertions(+), 7 deletions(-) diff --git a/compiler-core/src/opt/fields.rs b/compiler-core/src/opt/fields.rs index 958e3dc5..1aa497d3 100644 --- a/compiler-core/src/opt/fields.rs +++ b/compiler-core/src/opt/fields.rs @@ -1,16 +1,34 @@ //! Promote exact typed field projections to direct JVM field accesses. use crate::ir::*; -fn projection(body: &Body, mut value: ValueId) -> Option<(ValueId, ProjectionId)> { - // Reinterpretations change only the annotation, never allocation identity. - // Casts, offsets and nontrivial joins must retain the general pointer path. +fn projection(body: &Body, types: &Types, mut value: ValueId) -> Option<(ValueId, ProjectionId)> { + // Annotations and exact-layout retypes preserve field storage identity. + // Casts, offsets and nontrivial joins require the general pointer path. + let mut layout: Option<(u32, Option)> = None; for _ in 0..64 { value = body.resolve(value); let ValueDef::Inst(id) = body.values[value.index()].def else { return None; }; match body.instructions[id.index()].op { - Op::Project { base, projection } => return Some((base, projection)), + Op::Project { base, projection } => { + let field = &body.projections[projection.index()]; + return layout + .is_none_or(|(size, codec)| { + field.size == u64::from(size) + && field.codec.as_deref() + == codec.and_then(|codec| types.symbol_name(codec)) + }) + .then_some((base, projection)); + } + Op::RetypeAddress { + pointer, + size: size @ 1.., + codec, + } if layout.is_none_or(|previous| previous == (size, codec)) => { + layout = Some((size, codec)); + value = pointer; + } Op::Reinterpret(source) => value = source, _ => return None, } @@ -50,15 +68,41 @@ pub fn promote_fields(body: &mut Body, types: &Types) { for index in 0..body.instructions.len() { let inst = body.instructions[index]; let (pointer, ty) = match inst.op { - Op::Load(pointer) => (pointer, body.value_type(inst.result.unwrap())), + Op::Load(pointer) | Op::LoadCopy(pointer) => { + (pointer, body.value_type(inst.result.unwrap())) + } Op::Store { pointer, value } => (pointer, body.value_type(value)), _ => continue, }; - let Some((base, projection)) = projection(body, pointer) else { + let Some((base, projection)) = projection(body, types, pointer) else { continue; }; let field = &body.fields[body.projections[projection.index()].field.index()]; - if field.ty != ty || !matches!(types.get(ty), Some(Type::Scalar(_))) { + if matches!(inst.op, Op::LoadCopy(_)) { + if field.ty == ty && matches!(types.get(ty), Some(Type::Class(_) | Type::Array(_))) { + body.instructions[index].op = Op::LoadFieldCopy { base, projection }; + } + continue; + } + if let Op::Store { value, .. } = inst.op + && field.ty == ty + && matches!(types.get(ty), Some(Type::Class(_) | Type::Array(_))) + { + body.instructions[index].op = Op::StoreField { + base, + projection, + value, + }; + continue; + } + // Pointer values are immutable carriers. + // Read and replace them through the same checked storage path. + if field.ty != ty + || !matches!( + types.get(ty), + Some(Type::Scalar(_) | Type::Pointer(_) | Type::Slice(_) | Type::Str) + ) + { continue; } body.instructions[index].op = match inst.op { diff --git a/compiler-core/src/opt/fields_tests.rs b/compiler-core/src/opt/fields_tests.rs index 6c832b3f..65942898 100644 --- a/compiler-core/src/opt/fields_tests.rs +++ b/compiler-core/src/opt/fields_tests.rs @@ -2,6 +2,92 @@ use super::fields::promote_fields; use crate::ir::*; use crate::scalar::{Scalar, ScalarType}; +#[test] +fn owned_field_copies_keep_exception_edges_and_eliminate_address_carriers() { + for components in [false, true] { + for protected in [false, true] { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let array = types.intern(Type::Array(byte)); + let array_pointer = types.intern(Type::Pointer(array)); + let owner = types.symbol("Pixel"); + let owner = types.intern(Type::Class(owner)); + let pointer = types.intern(Type::Pointer(owner)); + let mut b = Builder::new(&types, array); + let base = b.parameter(b.current(), pointer); + let field = b.field(FieldRef { + owner, + name: "rgba".into(), + ty: array, + is_static: false, + }); + let projection = b.projection(PointerProjection { + field, + offset: 4, + size: 4, + codec: Some("org/rustlang/runtime/ArrayMemoryCodec#array#[B#4".into()), + }); + let address = b + .emit(Op::Project { base, projection }, Some(array_pointer)) + .unwrap(); + let value = if protected { + let handler = b.create_block(); + let value = b + .invoke(Op::LoadCopy(address), Some(array), handler) + .unwrap(); + b.terminate(Terminator::Return(Some(value))); + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + value + } else { + let value = b.emit(Op::LoadCopy(address), Some(array)).unwrap(); + b.terminate(Terminator::Return(Some(value))); + value + }; + let mut body = b.finish().unwrap(); + promote_fields(&mut body, &types); + if components { + super::lower_component_arguments(&mut body, &mut types, true, |_| false, None); + super::decompose_addresses(&mut body, &mut types, None); + } + verify(&body, &types).unwrap(); + let ValueDef::Inst(copy) = body.values[value.index()].def else { + panic!() + }; + assert!(if components { + matches!( + body.instructions[copy.index()].op, + Op::LoadStorageFieldCopy { .. } + ) + } else { + matches!(body.instructions[copy.index()].op, Op::LoadFieldCopy { .. }) + }); + if protected { + assert!(body.blocks.iter().any(|b| matches!(b.terminator, + Some(Terminator::Invoke { inst, .. }) if inst == copy))); + } + let live = super::live(&body, &types); + assert!(!live.values[address.index()]); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!(inst.op, Op::Project { .. } | Op::AddressPack(_))) + ); + let mut pool = Default::default(); + let code = crate::jvm::select::compile(&body, &types, &mut pool).unwrap(); + let owner = pool.add_class("org/rustlang/runtime/Pointer").unwrap(); + let load = pool.add_method_ref(owner, "loadStorageFieldCopy", + "(Ljava/lang/Object;JLjava/lang/String;Ljava/lang/String;JJLjava/lang/String;Ljava/lang/String;)Ljava/lang/Object;").unwrap(); + assert!(code.instructions.contains( + &crate::classfile::attributes::Instruction::Invokestatic(load) + )); + } + } +} + fn layout(types: &mut Types) -> (TypeId, TypeId, TypeId) { let scalar = types.scalar(ScalarType::I64); let owner = types.symbol("Counter"); @@ -140,3 +226,220 @@ fn rejects_field_loads_with_a_different_view_type() { .any(|i| matches!(i.op, Op::LoadField { .. })) ); } + +#[test] +fn retyped_borrowed_fields_promote_only_with_the_original_exact_layout() { + for mismatch in [0, 1, 2, 3] { + for protected in [false, true] { + let mut types = Types::default(); + let (_, address, _) = layout(&mut types); + let slot = types.intern(Type::Pointer(address)); + let owner = types.symbol("Iterator"); + let owner = types.intern(Type::Class(owner)); + let base_ty = types.intern(Type::Pointer(owner)); + let codec = types.symbol("@raw-pointer\n8\n\n"); + let other_codec = types.symbol("@raw-pointer\n4\n\n"); + let mut b = Builder::new(&types, address); + let base = b.parameter(b.current(), base_ty); + let field = b.field(FieldRef { + owner, + name: "end".into(), + ty: address, + is_static: false, + }); + let projection = b.projection(PointerProjection { + field, + offset: 8, + size: 8, + codec: Some("@raw-pointer\n8\n\n".into()), + }); + let original = b + .emit(Op::Project { base, projection }, Some(slot)) + .unwrap(); + let retyped = b + .emit( + Op::RetypeAddress { + pointer: original, + size: if mismatch == 1 { + 4 + } else if mismatch == 3 { + 0 + } else { + 8 + }, + codec: Some(if mismatch == 2 { other_codec } else { codec }), + }, + Some(slot), + ) + .unwrap(); + let value = if protected { + let handler = b.create_block(); + let value = b.invoke(Op::Load(retyped), Some(address), handler).unwrap(); + let current = b.current(); + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + b.switch_to(current); + value + } else { + b.emit(Op::Load(retyped), Some(address)).unwrap() + }; + b.emit( + Op::Store { + pointer: retyped, + value, + }, + None, + ); + b.terminate(Terminator::Return(Some(value))); + let mut body = b.finish().unwrap(); + promote_fields(&mut body, &types); + verify(&body, &types).unwrap(); + let ValueDef::Inst(load) = body.values[value.index()].def else { + panic!() + }; + assert_eq!( + matches!(body.instructions[load.index()].op, Op::LoadField { .. }), + mismatch == 0 + ); + assert_eq!( + body.instructions + .iter() + .any(|inst| matches!(inst.op, Op::StoreField { .. })), + mismatch == 0 + ); + if protected { + assert!(body.blocks.iter().any(|block| matches!(block.terminator, + Some(Terminator::Invoke { inst, .. }) if inst == load))); + } + if mismatch == 0 { + assert!(!super::live(&body, &types).values[retyped.index()]); + } + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } + } +} + +#[test] +fn field_storage_dispatch_preserves_forwarded_roots() { + for store in [false, true] { + let mut types = Types::default(); + let (base_ty, _, scalar) = layout(&mut types); + let Type::Pointer(owner) = types.get(base_ty).unwrap() else { + unreachable!() + }; + let mut b = Builder::new(&types, scalar); + let input = b.parameter(b.current(), base_ty); + let field = b.field(FieldRef { + owner, + name: "value".into(), + ty: scalar, + is_static: false, + }); + let projection = b.projection(PointerProjection { + field, + offset: 8, + size: 8, + codec: None, + }); + let base = b.emit(Op::Opaque(input), Some(base_ty)).unwrap(); + let value = if store { + let value = b.constant(scalar, Scalar::integer(ScalarType::I64, 17).unwrap()); + b.emit( + Op::StoreField { + base, + projection, + value, + }, + None, + ); + value + } else { + b.emit(Op::LoadField { base, projection }, Some(scalar)) + .unwrap() + }; + b.terminate(Terminator::Return(Some(value))); + let body = b.finish().unwrap(); + let mut pool = crate::classfile::constant_pool::InternedConstantPool::default(); + let code = crate::jvm::select::compile(&body, &types, &mut pool).unwrap(); + let owner = pool.add_class("org/rustlang/runtime/Pointer").unwrap(); + let direct = pool + .add_method_ref( + owner, + if store { + "storeScalarField" + } else { + "loadScalarField" + }, + if store { + "(Ljava/lang/Object;JLjava/lang/String;Ljava/lang/String;JJI)V" + } else { + "(Ljava/lang/Object;JLjava/lang/String;Ljava/lang/String;JI)J" + }, + ) + .unwrap(); + let project = pool.add_method_ref(owner, "projectStructField", + "(Ljava/lang/String;Ljava/lang/String;JJLjava/lang/String;)Lorg/rustlang/runtime/Pointer;", + ).unwrap(); + use crate::classfile::attributes::Instruction; + let access = code + .instructions + .iter() + .position(|i| *i == Instruction::Invokestatic(direct)) + .expect("byte storage must have a direct scalar access"); + assert!( + !code.instructions[..access].contains(&Instruction::Invokevirtual(project)), + "byte storage must not first construct a projected field pointer" + ); + } +} + +#[test] +fn aggregate_field_stores_reuse_storage_components_and_keep_handlers() { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let array = types.intern(Type::Array(byte)); + let array_pointer = types.intern(Type::Pointer(array)); + let owner = types.symbol("Pixel"); + let owner = types.intern(Type::Class(owner)); + let pointer = types.intern(Type::Pointer(owner)); + let unit = types.intern(Type::Unit); + let mut b = Builder::new(&types, unit); + let base = b.parameter(b.current(), pointer); + let value = b.parameter(b.current(), array); + let field = b.field(FieldRef { + owner, + name: "rgba".into(), + ty: array, + is_static: false, + }); + let projection = b.projection(PointerProjection { + field, + offset: 4, + size: 4, + codec: Some("org/rustlang/runtime/ArrayMemoryCodec#array#[B#4".into()), + }); + let address = b + .emit(Op::Project { base, projection }, Some(array_pointer)) + .unwrap(); + let handler = b.create_block(); + b.invoke( + Op::Store { + pointer: address, + value, + }, + None, + handler, + ); + b.terminate(Terminator::Return(None)); + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + let mut body = b.finish().unwrap(); + promote_fields(&mut body, &types); + super::lower_component_arguments(&mut body, &mut types, true, |_| false, None); + super::decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + assert!(body.blocks.iter().any(|b| matches!(b.terminator, + Some(Terminator::Invoke { inst, .. }) if matches!(body.instructions[inst.index()].op, Op::StoreStorageField { .. })))); + assert!(!super::live(&body, &types).values[address.index()]); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} From 993c755136a4ff4a670f0e5a4a0ff89e93edb7da Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 08:03:25 +1000 Subject: [PATCH 19/61] preserve storage owners and borrows --- compiler-core/src/opt/borrowed_memory.rs | 122 ++++++++++ .../src/opt/borrowed_memory_tests.rs | 133 +++++++++++ compiler-core/src/opt/cells.rs | 95 ++++---- compiler-core/src/opt/cells_tests.rs | 213 ++++++++++++++++++ compiler-core/src/opt/mod.rs | 8 + compiler-core/src/opt/storage.rs | 203 +++++++++++++++++ compiler-core/src/opt/storage_tests.rs | 153 +++++++++++++ 7 files changed, 883 insertions(+), 44 deletions(-) create mode 100644 compiler-core/src/opt/borrowed_memory.rs create mode 100644 compiler-core/src/opt/borrowed_memory_tests.rs create mode 100644 compiler-core/src/opt/storage.rs create mode 100644 compiler-core/src/opt/storage_tests.rs diff --git a/compiler-core/src/opt/borrowed_memory.rs b/compiler-core/src/opt/borrowed_memory.rs new file mode 100644 index 00000000..724fddf8 --- /dev/null +++ b/compiler-core/src/opt/borrowed_memory.rs @@ -0,0 +1,122 @@ +//! Use the component ABI for stored borrows, calls, returns and fields. +//! Run after local promotion so eliminated cells gain no runtime calls. +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; +use rustc_hash::FxHashMap; + +const OWNER: &str = "org/rustlang/runtime/Pointer"; + +pub fn borrowed_memory_method(method: &MethodRef) -> bool { + method.owner == OWNER + && matches!( + method.name.as_str(), + "loadBorrowedView" + | "loadBorrowedAddress" + | "storeBorrowedView" + | "storeBorrowedUtf8" + | "storeBorrowedAddress" + ) +} + +pub fn lower_borrowed_memory(body: &mut Body, types: &mut Types, debug: Option<&mut DebugInfo>) { + let mut methods = FxHashMap::default(); + let int = types.scalar(ScalarType::I32); + let unit = types.intern(Type::Unit); + let mut sizes = FxHashMap::default(); + let mut prologue = Vec::new(); + for index in 0..body.instructions.len() { + let instruction = body.instructions[index]; + let (pointer, value, ty) = match instruction.op { + Op::Load(pointer) | Op::LoadCopy(pointer) => { + (pointer, None, body.value_type(instruction.result.unwrap())) + } + Op::Store { pointer, value } => (pointer, Some(value), body.value_type(value)), + _ => continue, + }; + let Some(shape) = ComponentShape::of(types, ty) else { + continue; + }; + if shape == ComponentShape::TaggedI64 { + continue; + } + let parameter = body.value_type(pointer); + let mut params = vec![parameter]; + let mut args = vec![pointer]; + let (name, returns) = if let Some(value) = value { + params.push(ty); + args.push(value); + if shape == ComponentShape::View { + ( + if types.get(ty) == Some(Type::Str) { + "storeBorrowedUtf8" + } else { + "storeBorrowedView" + }, + unit, + ) + } else { + let size = match types.get(ty) { + Some(Type::Pointer(inner)) => crate::jvm::abi::address_plan(types, inner), + _ => unreachable!(), + }; + let constant = *sizes.entry(size).or_insert_with(|| { + let id = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar( + Scalar::integer(ScalarType::I32, size.into()).unwrap(), + )); + let inst = InstId::new(body.instructions.len()); + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty: int, + def: ValueDef::Inst(inst), + }); + body.instructions.push(Inst { + op: Op::Constant(id), + result: Some(value), + }); + prologue.push(inst); + value + }); + params.push(int); + args.push(constant); + ("storeBorrowedAddress", unit) + } + } else { + ( + if shape == ComponentShape::View { + "loadBorrowedView" + } else { + "loadBorrowedAddress" + }, + ty, + ) + }; + let id = *methods + .entry((parameter, ty, value.is_some())) + .or_insert_with(|| { + body.methods.push(MethodRef { + owner: OWNER.into(), + name: name.into(), + params, + returns, + interface: false, + }); + body.methods.len() - 1 + }); + body.instructions[index].op = Op::Call { + method: MethodId::new(id), + kind: CallKind::JvmStatic, + args: List::append(&mut body.args, args), + }; + } + let count = prologue.len() as u32; + prologue.append(&mut body.blocks[body.entry.index()].instructions); + body.blocks[body.entry.index()].instructions = prologue; + if let Some(debug) = debug { + for event in &mut debug.events { + if event.block == body.entry { + event.position += count; + } + } + } +} diff --git a/compiler-core/src/opt/borrowed_memory_tests.rs b/compiler-core/src/opt/borrowed_memory_tests.rs new file mode 100644 index 00000000..72218afd --- /dev/null +++ b/compiler-core/src/opt/borrowed_memory_tests.rs @@ -0,0 +1,133 @@ +use super::*; +use crate::{ir::*, scalar::ScalarType}; + +#[test] +fn stored_borrows_cross_load_store_and_return_without_carriers() { + for view in 0..3 { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let borrowed = types.intern(match view { + 0 => Type::Pointer(byte), + 1 => Type::Slice(byte), + _ => Type::Str, + }); + let location = types.intern(Type::Pointer(borrowed)); + let mut builder = Builder::new(&types, borrowed); + let destination = builder.parameter(builder.current(), location); + let replacement = builder.parameter(builder.current(), borrowed); + let original = builder.emit(Op::Load(destination), Some(borrowed)).unwrap(); + builder.emit( + Op::Store { + pointer: destination, + value: replacement, + }, + None, + ); + builder.terminate(Terminator::Return(Some(original))); + let mut body = builder.finish().unwrap(); + lower_borrowed_memory(&mut body, &mut types, None); + lower_component_arguments(&mut body, &mut types, true, borrowed_memory_method, None); + lower_component_returns(&mut body, &mut types, true, borrowed_memory_method, None); + decompose_views(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + let mut calls = 0; + for (index, instruction) in body.instructions.iter().enumerate() { + if !live.instructions[index] { + continue; + } + assert!( + !matches!( + instruction.op, + Op::AddressPack(_) + | Op::ViewPack(_) + | Op::Load(_) + | Op::Store { .. } + | Op::NewArray(_) + ), + "{instruction:?}" + ); + if let Op::Call { method, .. } = instruction.op { + let method = &body.methods[method.index()]; + assert!( + method + .params + .iter() + .all(|&p| ComponentShape::of(&types, p).is_none()) + ); + assert!(ComponentShape::of(&types, method.returns).is_none()); + calls += 1; + } + } + assert_eq!(calls, 2); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} + +#[test] +fn references_to_stored_borrows_keep_their_enclosing_owner() { + for view in [false, true] { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let borrowed = types.intern(if view { + Type::Slice(byte) + } else { + Type::Pointer(byte) + }); + let location = types.intern(Type::Pointer(borrowed)); + let name = types.symbol("Slot"); + let owner = types.intern(Type::Class(name)); + let aggregate = types.intern(Type::Pointer(owner)); + let mut builder = Builder::new(&types, borrowed); + let root = builder.parameter(builder.current(), aggregate); + let field = builder.field(FieldRef { + owner, + name: "borrow".into(), + ty: borrowed, + is_static: false, + }); + let projection = builder.projection(PointerProjection { + field, + offset: 8, + size: if view { 16 } else { 8 }, + codec: Some("layout".into()), + }); + let address = builder + .emit( + Op::Project { + base: root, + projection, + }, + Some(location), + ) + .unwrap(); + let value = builder.emit(Op::Load(address), Some(borrowed)).unwrap(); + builder.terminate(Terminator::Return(Some(value))); + let mut body = builder.finish().unwrap(); + lower_borrowed_memory(&mut body, &mut types, None); + lower_component_arguments(&mut body, &mut types, true, borrowed_memory_method, None); + lower_component_returns(&mut body, &mut types, true, borrowed_memory_method, None); + decompose_views(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::ProjectRoot { .. })) + ); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, instruction)| live.instructions[i] + && matches!( + instruction.op, + Op::Project { .. } | Op::AddressPack(_) | Op::ViewPack(_) + )) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} diff --git a/compiler-core/src/opt/cells.rs b/compiler-core/src/opt/cells.rs index 8a10d715..356c2576 100644 --- a/compiler-core/src/opt/cells.rs +++ b/compiler-core/src/opt/cells.rs @@ -1,47 +1,12 @@ //! Escape analysis and SSA promotion for private, explicitly initialized cells. use crate::ir::*; -const NONE: u32 = u32::MAX; +use crate::analysis::{NO_ORIGIN as NONE, origins}; fn cell_index(index: u32) -> Option { (index != NONE).then_some(index as usize) } -fn origins(body: &Body, roots: &[u32]) -> Vec { - let mut origins = roots.to_vec(); - let mut known = roots.iter().map(|&index| index != NONE).collect::>(); - let mut path = Vec::new(); - for index in 0..body.values.len() { - if known[index] { - continue; - } - let mut value = ValueId::new(index); - let root = loop { - path.push(value); - value = body.resolve(value); - if known[value.index()] { - break origins[value.index()]; - } - // Memoize negative results too. Marking before following a use - // bounds the walk even for malformed reinterpretation cycles. - known[value.index()] = true; - path.push(value); - let ValueDef::Inst(inst) = body.values[value.index()].def else { - break NONE; - }; - let Op::Reinterpret(input) = body.instructions[inst.index()].op else { - break NONE; - }; - value = input; - }; - for value in path.drain(..) { - known[value.index()] = true; - origins[value.index()] = root; - } - } - origins -} - /// Each candidate is a fresh allocation and its typed initial contents. The /// producer guarantees that creating it has no effects beyond allocation. /// Keep cells whose address escapes, is retyped, or participates in a join of @@ -60,37 +25,66 @@ pub fn promote_cells( roots[cell.index()] = u32::try_from(index).expect("too many private cells"); } let origins = origins(&body, &roots); + // Field promotion can leave unused projections. + // They do not observe addresses and must not block SSA storage promotion. + let live = super::live(&body, types); let mut escaped = vec![false; cells.len()]; - for inst in &body.instructions { + for (id, inst) in body.instructions.iter().enumerate() { + if !live.instructions[id] { + continue; + } inst.op.visit_uses(&body.args, |value| { let Some(index) = cell_index(origins[value.index()]) else { return; }; let pointee = body.value_type(cells[index].1); let allowed = match inst.op { - Op::Reinterpret(_) => true, - Op::Load(pointer) => { + Op::Reinterpret(_) | Op::Refine(_) => true, + Op::Load(pointer) | Op::LoadCopy(pointer) => { pointer == value && inst.result.is_some_and(|v| body.value_type(v) == pointee) } Op::Store { pointer, value: stored, } => pointer == value && stored != value && body.value_type(stored) == pointee, + Op::LoadField { base, projection } => { + base == value + && body.fields[body.projections[projection.index()].field.index()].owner + == pointee + } + Op::StoreField { + base, + projection, + value: stored, + } => { + base == value + && stored != value + && body.fields[body.projections[projection.index()].field.index()].owner + == pointee + } _ => false, }; escaped[index] |= !allowed; }); } - // Only trivial address joins have been resolved. A nontrivial address phi - // may select another allocation; retain its storage rather than guessing. + // Equal-origin joins retain one allocation. + // Mixed joins must retain each allocation's addressable storage. let mut escape = |value: ValueId| { if let Some(index) = cell_index(origins[value.index()]) { escaped[index] = true; } }; for edge in &body.edges { - for &value in &edge.args { - escape(value); + for (&value, ¶m) in edge + .args + .iter() + .zip(&body.blocks[edge.target.index()].params) + { + if live.values[body.resolve(param).index()] + && origins[value.index()] != origins[param.index()] + { + escape(value); + } } } for block in &body.blocks { @@ -135,7 +129,8 @@ pub fn promote_cells( continue; } let pointer = match inst.op { - Op::Load(pointer) | Op::Store { pointer, .. } => pointer, + Op::Load(pointer) | Op::LoadCopy(pointer) | Op::Store { pointer, .. } => pointer, + Op::LoadField { base, .. } | Op::StoreField { base, .. } => base, _ => continue, }; let Some(variable) = cell_index(origins[pointer.index()]).and_then(|i| variables[i]) @@ -144,10 +139,22 @@ pub fn promote_cells( }; builder.body.instructions[id.index()].op = match inst.op { Op::Load(_) => Op::Reinterpret(builder.read(variable)), + Op::LoadCopy(_) => Op::CopyValue(builder.read(variable)), Op::Store { value, .. } => { builder.define(variable, value); Op::Nop } + Op::LoadField { projection, .. } => Op::GetField { + object: builder.read(variable), + field: builder.body.projections[projection.index()].field, + }, + Op::StoreField { + projection, value, .. + } => Op::SetField { + object: builder.read(variable), + field: builder.body.projections[projection.index()].field, + value, + }, _ => unreachable!(), }; } diff --git a/compiler-core/src/opt/cells_tests.rs b/compiler-core/src/opt/cells_tests.rs index fe7e615f..b210e062 100644 --- a/compiler-core/src/opt/cells_tests.rs +++ b/compiler-core/src/opt/cells_tests.rs @@ -100,6 +100,32 @@ fn keeps_storage_when_its_address_escapes() { assert_eq!(body, before); } +#[test] +fn private_array_cells_disappear_but_owned_loads_still_copy() { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let array = types.intern(Type::Array(byte)); + let pointer = types.intern(Type::Pointer(array)); + let mut b = Builder::new(&types, array); + let initial = b.parameter(b.current(), array); + let storage = cell(&mut b, array, pointer, initial); + let copied = b.emit(Op::LoadCopy(storage), Some(array)).unwrap(); + b.terminate(Terminator::Return(Some(copied))); + let body = promote_cells(b.finish().unwrap(), &types, &[(storage, initial)]).unwrap(); + verify(&body, &types).unwrap(); + assert!( + body.instructions + .iter() + .any(|i| i.op == Op::CopyValue(initial)) + ); + assert!( + !body + .instructions + .iter() + .any(|i| matches!(i.op, Op::Call { .. } | Op::LoadCopy(_))) + ); +} + #[test] fn promoted_contents_reach_the_unwind_edge_at_the_throwing_call() { let mut types = Types::default(); @@ -214,3 +240,190 @@ fn follows_long_alias_chains_for_both_loads_and_escapes() { } } } + +#[test] +fn projected_borrow_follows_replacement_of_private_aggregate_storage() { + let mut types = Types::default(); + types.intern(Type::Unit); + let int = types.scalar(ScalarType::I64); + let name = types.symbol("Counter"); + let object = types.intern(Type::Class(name)); + let pointer = types.intern(Type::Pointer(object)); + let field_pointer = types.intern(Type::Pointer(int)); + for escapes in [false, true] { + let mut b = Builder::new(&types, if escapes { field_pointer } else { int }); + let first = b.parameter(b.current(), object); + let second = b.parameter(b.current(), object); + let storage = cell(&mut b, object, pointer, first); + let field = b.field(FieldRef { + owner: object, + name: "count".into(), + ty: int, + is_static: false, + }); + let projection = b.projection(PointerProjection { + field, + offset: 0, + size: 8, + codec: None, + }); + let borrow = b + .emit( + Op::Project { + base: storage, + projection, + }, + Some(field_pointer), + ) + .unwrap(); + let before = b.emit(Op::Load(borrow), Some(int)).unwrap(); + b.emit( + Op::Store { + pointer: storage, + value: second, + }, + None, + ); + b.emit( + Op::Store { + pointer: borrow, + value: before, + }, + None, + ); + let after = b.emit(Op::Load(borrow), Some(int)).unwrap(); + b.terminate(Terminator::Return(Some(if escapes { + borrow + } else { + after + }))); + let mut body = b.finish().unwrap(); + super::promote_fields(&mut body, &types); + let original = body.clone(); + let body = promote_cells(body, &types, &[(storage, first)]).unwrap(); + verify(&body, &types).unwrap(); + if escapes { + assert_eq!(body, original); + continue; + } + let read = |value: ValueId| { + let ValueDef::Inst(inst) = body.values[value.index()].def else { + panic!() + }; + body.instructions[inst.index()].op + }; + assert_eq!( + read(before), + Op::GetField { + object: first, + field + } + ); + assert_eq!( + read(after), + Op::GetField { + object: second, + field + } + ); + assert!(body.instructions.iter().any(|i| i.op + == Op::SetField { + object: second, + field, + value: before, + })); + let live = super::live(&body, &types); + assert!(!live.values[storage.index()] && !live.values[borrow.index()]); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} + +#[test] +fn equal_allocation_origins_cross_control_flow_joins() { + for distinct in [false, true] { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I64); + let boolean = types.scalar(ScalarType::Bool); + let pointer = types.intern(Type::Pointer(int)); + let mut b = Builder::new(&types, int); + let condition = b.parameter(b.current(), boolean); + let initial = b.constant(int, Scalar::integer(ScalarType::I64, 7).unwrap()); + let first = cell(&mut b, int, pointer, initial); + let second = if distinct { + cell(&mut b, int, pointer, initial) + } else { + first + }; + let left = b.create_block(); + let right = b.create_block(); + let join = b.create_block(); + let selected = b.parameter(join, pointer); + b.branch(condition, left, right); + b.switch_to(left); + let alias = b.emit(Op::Reinterpret(first), Some(pointer)).unwrap(); + b.jump(join, vec![alias]); + b.switch_to(right); + let alias = b.emit(Op::Refine(second), Some(pointer)).unwrap(); + b.jump(join, vec![alias]); + b.switch_to(join); + let next = b.constant(int, Scalar::integer(ScalarType::I64, 11).unwrap()); + b.emit( + Op::Store { + pointer: selected, + value: next, + }, + None, + ); + let result = b.emit(Op::Load(first), Some(int)).unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut cells = vec![(first, initial)]; + if distinct { + cells.push((second, initial)); + } + let body = promote_cells(b.finish().unwrap(), &types, &cells).unwrap(); + verify(&body, &types).unwrap(); + let retains_storage = body + .instructions + .iter() + .any(|i| matches!(i.op, Op::Call { .. })); + assert_eq!(retains_storage, distinct); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} + +#[test] +fn nonentry_allocations_remain_conservative_at_joins() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I64); + let boolean = types.scalar(ScalarType::Bool); + let pointer = types.intern(Type::Pointer(int)); + let mut b = Builder::new(&types, int); + let again = b.parameter(b.current(), boolean); + let initial = b.constant(int, Scalar::integer(ScalarType::I64, 7).unwrap()); + let allocate = b.create_block(); + let join = b.create_block(); + let repeat = b.create_block(); + let done = b.create_block(); + let selected = b.parameter(join, pointer); + b.jump(allocate, vec![]); + b.switch_to(allocate); + let storage = cell(&mut b, int, pointer, initial); + b.jump(join, vec![storage]); + b.switch_to(join); + let value = b.emit(Op::Load(selected), Some(int)).unwrap(); + b.branch(again, repeat, done); + b.switch_to(repeat); + let alias = b.emit(Op::Reinterpret(storage), Some(pointer)).unwrap(); + b.jump(join, vec![alias]); + b.switch_to(done); + b.terminate(Terminator::Return(Some(value))); + let body = promote_cells(b.finish().unwrap(), &types, &[(storage, initial)]).unwrap(); + // Only the entry region has a single-execution proof. + // This pass must retain allocations elsewhere, even if another proof could remove them. + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::Call { .. })) + ); + verify(&body, &types).unwrap(); +} diff --git a/compiler-core/src/opt/mod.rs b/compiler-core/src/opt/mod.rs index 78c857fd..67997224 100644 --- a/compiler-core/src/opt/mod.rs +++ b/compiler-core/src/opt/mod.rs @@ -1,5 +1,9 @@ //! Analyses are built on demand and owned by one body compilation. +mod borrowed_memory; mod field_abi; +pub use borrowed_memory::{borrowed_memory_method, lower_borrowed_memory}; +#[cfg(test)] +mod borrowed_memory_tests; mod fields; pub use field_abi::lower_borrowed_fields; #[cfg(test)] @@ -12,6 +16,10 @@ pub use live::{Live, live, live_with_roots}; mod cells; pub use cells::promote_cells; +mod storage; +pub use storage::lower_typed_storage; +#[cfg(test)] +mod storage_tests; mod tagged; pub use tagged::decompose_tagged; mod views; diff --git a/compiler-core/src/opt/storage.rs b/compiler-core/src/opt/storage.rs new file mode 100644 index 00000000..09a6c90e --- /dev/null +++ b/compiler-core/src/opt/storage.rs @@ -0,0 +1,203 @@ +//! Keep typed allocations as storage owners instead of eager address wrappers. +//! Private cells already use SSA. Escaping locations retain one owner, +//! including when a store replaces their contents. +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; +use rustc_hash::FxHashMap; + +/// The frontend checks positive size, alignment and exact initial contents. +/// It also requires a direct carrier with no DST metadata. +pub fn lower_typed_storage( + body: &mut Body, + types: &mut Types, + cells: &[(ValueId, ValueId)], + debug: Option<&mut DebugInfo>, +) { + if cells.is_empty() { + return; + } + let mut parts = ComponentShape::StorageAddress.parts(types); + let object = parts.next().unwrap(); + let long = parts.next().unwrap(); + let mut prefixes = FxHashMap::default(); + let mut suffixes = FxHashMap::default(); + for &(cell, _) in cells { + let ValueDef::Inst(id) = body.values[cell.index()].def else { + continue; + }; + let mut suffix = Vec::new(); + let (op, root_ty, initial) = match body.instructions[id.index()].op { + Op::ScalarCell(initial) => { + let int = types.scalar(ScalarType::I32); + let (one_inst, one) = integer(body, int, ScalarType::I32, 1); + prefixes.insert(id, one_inst); + ( + Op::NewArray(one), + types.intern(Type::Array(body.value_type(initial))), + Some(initial), + ) + } + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + } => { + let mut method = body.methods[method.index()].clone(); + if method.owner != "org/rustlang/runtime/Pointer" { + continue; + } + let borrowed = match types.get(body.value_type(cell)) { + Some(Type::Pointer(inner)) => match types.get(inner) { + Some(Type::Pointer(_)) => Some(0), + Some(Type::Slice(_)) => Some(1), + Some(Type::Str) => Some(2), + _ => None, + }, + _ => None, + }; + method.name = match (method.name.as_str(), borrowed.is_some()) { + ("cell", false) => "storage", + ("cellAligned", false) => "storageAligned", + ("cell", true) => "borrowedStorage", + ("cellAligned", true) => "borrowedStorageAligned", + _ => continue, + } + .into(); + let args = if let Some(kind) = borrowed { + let int = types.scalar(ScalarType::I32); + let (inst, kind) = integer(body, int, ScalarType::I32, kind); + prefixes.insert(id, inst); + method.params.push(int); + let mut values = body.args[args.range()].to_vec(); + values.push(kind); + List::append(&mut body.args, values) + } else { + args + }; + method.returns = object; + let method_id = MethodId::new(body.methods.len()); + body.methods.push(method); + ( + Op::Call { + method: method_id, + kind: CallKind::JvmStatic, + args, + }, + object, + None, + ) + } + _ => continue, + }; + let root = ValueId::new(body.values.len()); + body.values.push(Value { + ty: root_ty, + def: ValueDef::Inst(id), + }); + body.instructions[id.index()] = Inst { + op, + result: Some(root), + }; + let root = if let Some(initial) = initial { + let int = types.scalar(ScalarType::I32); + let (index_inst, index) = integer(body, int, ScalarType::I32, 0); + suffix.push(index_inst); + suffix.push(InstId::new(body.instructions.len())); + body.instructions.push(Inst { + op: Op::ArraySet { + native: false, + array: root, + index, + value: initial, + }, + result: None, + }); + let cast = InstId::new(body.instructions.len()); + let erased = ValueId::new(body.values.len()); + body.values.push(Value { + ty: object, + def: ValueDef::Inst(cast), + }); + body.instructions.push(Inst { + op: Op::Reinterpret(root), + result: Some(erased), + }); + suffix.push(cast); + erased + } else { + root + }; + let (zero_inst, zero) = integer(body, long, ScalarType::I64, 0); + let pack = InstId::new(body.instructions.len()); + body.values[cell.index()].def = ValueDef::Inst(pack); + body.instructions.push(Inst { + op: Op::AddressPack(List::append(&mut body.args, [root, zero])), + result: Some(cell), + }); + suffix.extend([zero_inst, pack]); + suffixes.insert(id, suffix); + } + if suffixes.is_empty() { + return; + } + let mut positions = debug.as_ref().map(|_| Vec::new()); + for index in 0..body.blocks.len() { + let previous = std::mem::take(&mut body.blocks[index].instructions); + let mut mapping = Vec::new(); + for id in previous { + mapping.push(body.blocks[index].instructions.len() as u32); + if let Some(prefix) = prefixes.remove(&id) { + body.blocks[index].instructions.push(prefix); + } + body.blocks[index].instructions.push(id); + if let Some(suffix) = suffixes.remove(&id) { + body.blocks[index].instructions.extend(suffix); + } + } + mapping.push(body.blocks[index].instructions.len() as u32); + if let Some(positions) = &mut positions { + positions.push(mapping); + } + if let Some(Terminator::Invoke { + inst, + normal, + unwind, + }) = body.blocks[index].terminator + && let Some(suffix) = suffixes.remove(&inst) + { + if let Some(prefix) = prefixes.remove(&inst) { + body.blocks[index].instructions.push(prefix); + } + // Keep allocation in its original unwind region. + // Pack its result on the normal edge before copies consume the address. + let continuation = BlockId::new(body.blocks.len()); + body.blocks.push(Block { + params: Vec::new(), + instructions: suffix, + terminator: Some(Terminator::Jump(normal)), + }); + let edge = EdgeId::new(body.edges.len()); + body.edges.push(Edge { + target: continuation, + args: Vec::new(), + }); + body.blocks[index].terminator = Some(Terminator::Invoke { + inst, + normal: edge, + unwind, + }); + } + } + if let (Some(debug), Some(positions)) = (debug, positions) { + for event in &mut debug.events { + event.position = positions[event.block.index()][event.position as usize]; + } + } +} + +fn integer(body: &mut Body, ty: TypeId, scalar: ScalarType, value: u128) -> (InstId, ValueId) { + let constant = ConstId::new(body.constants.len()); + body.constants + .push(Constant::Scalar(Scalar::integer(scalar, value).unwrap())); + super::append_value(body, Op::Constant(constant), ty) +} diff --git a/compiler-core/src/opt/storage_tests.rs b/compiler-core/src/opt/storage_tests.rs new file mode 100644 index 00000000..84dd7330 --- /dev/null +++ b/compiler-core/src/opt/storage_tests.rs @@ -0,0 +1,153 @@ +use super::{decompose_addresses, live, lower_typed_storage}; +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; + +#[test] +fn scalar_storage_uses_one_primitive_array_without_boxing() { + for scalar in [ + ScalarType::Bool, + ScalarType::U8, + ScalarType::I16, + ScalarType::U16, + ScalarType::F16, + ScalarType::I32, + ScalarType::U32, + ScalarType::I64, + ScalarType::U64, + ScalarType::F32, + ScalarType::F64, + ] { + for unwind in [false, true] { + let mut types = Types::default(); + let scalar_ty = types.scalar(scalar); + let pointer = types.intern(Type::Pointer(scalar_ty)); + let mut b = Builder::new(&types, scalar_ty); + let initial = b.parameter(b.current(), scalar_ty); + let cell = if unwind { + let handler = b.create_block(); + let cell = b + .invoke(Op::ScalarCell(initial), Some(pointer), handler) + .unwrap(); + let continuation = b.current(); + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + b.switch_to(continuation); + cell + } else { + b.emit(Op::ScalarCell(initial), Some(pointer)).unwrap() + }; + let result = b.emit(Op::Load(cell), Some(scalar_ty)).unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + let promoted = super::promote_cells(body.clone(), &types, &[(cell, initial)]).unwrap(); + let private = live(&promoted, &types); + assert!( + !promoted + .instructions + .iter() + .enumerate() + .any(|(id, i)| private.instructions[id] + && matches!(i.op, Op::ScalarCell(_) | Op::NewArray(_))) + ); + lower_typed_storage(&mut body, &mut types, &[(cell, initial)], None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let retained = live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(id, i)| retained.instructions[id] + && matches!( + i.op, + Op::ScalarCell(_) | Op::AddressPack(_) | Op::Adapt(_) | Op::Call { .. } + )) + ); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert_eq!( + code.instructions + .iter() + .filter(|i| matches!(i, crate::classfile::attributes::Instruction::Newarray(_))) + .count(), + 1 + ); + } + } +} + +#[test] +fn typed_storage_keeps_its_owner_and_allocation_unwind_edge() { + for unwind in [false, true] { + let mut types = Types::default(); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let string = types.symbol("java/lang/String"); + let string = types.intern(Type::Class(string)); + let pointer = types.intern(Type::Pointer(object)); + let int = types.scalar(ScalarType::I32); + let mut b = Builder::new(&types, object); + let initial = b.parameter(b.current(), object); + let size = b.constant(int, Scalar::integer(ScalarType::I32, 8).unwrap()); + let codec = ConstId::new(b.body.constants.len()); + b.body.constants.push(Constant::Null(string)); + let codec = b.emit(Op::Constant(codec), Some(string)).unwrap(); + let method = b.method(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: "cellAligned".into(), + params: vec![object, int, string, int], + returns: pointer, + interface: false, + }); + let args = b.args([initial, size, codec, size]); + let op = Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }; + let cell = if unwind { + let handler = b.create_block(); + let cell = b.invoke(op, Some(pointer), handler).unwrap(); + let continuation = b.current(); + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + b.switch_to(continuation); + cell + } else { + b.emit(op, Some(pointer)).unwrap() + }; + let result = b.emit(Op::Load(cell), Some(object)).unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + let mut debug = DebugInfo::default(); + debug.events.push(DebugEvent { + block: body.entry, + position: 2, + change: DebugChange::Scope(0), + line: Some(1), + }); + lower_typed_storage(&mut body, &mut types, &[(cell, initial)], Some(&mut debug)); + verify(&body, &types).unwrap(); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(id, inst)| live.instructions[id] + && matches!(inst.op, Op::AddressPack(_) | Op::Load(_))) + ); + assert_eq!( + body.blocks + .iter() + .filter(|b| matches!(b.terminator, Some(Terminator::Invoke { .. }))) + .count(), + usize::from(unwind) + ); + assert!(body.instructions.iter().any(|i| matches!(i.op, + Op::Call { method, .. } if body.methods[method.index()].name == "storageAligned"))); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} From ac0d3254bdfa13dd1d9e4f17031a5cefca2d83a7 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 09:26:30 +1000 Subject: [PATCH 20/61] retain exact address layouts --- compiler-core/src/opt/array_view_tests.rs | 166 +++++++ compiler-core/src/opt/mod.rs | 13 + compiler-core/src/opt/typed_addresses.rs | 470 ++++++++++++++++++ .../src/opt/typed_addresses_tests.rs | 452 +++++++++++++++++ compiler-core/src/opt/typed_loads.rs | 103 ++++ compiler-core/src/opt/typed_loads_tests.rs | 266 ++++++++++ 6 files changed, 1470 insertions(+) create mode 100644 compiler-core/src/opt/array_view_tests.rs create mode 100644 compiler-core/src/opt/typed_addresses.rs create mode 100644 compiler-core/src/opt/typed_addresses_tests.rs create mode 100644 compiler-core/src/opt/typed_loads.rs create mode 100644 compiler-core/src/opt/typed_loads_tests.rs diff --git a/compiler-core/src/opt/array_view_tests.rs b/compiler-core/src/opt/array_view_tests.rs new file mode 100644 index 00000000..cfd2ffd1 --- /dev/null +++ b/compiler-core/src/opt/array_view_tests.rs @@ -0,0 +1,166 @@ +use super::*; +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; + +#[test] +fn zero_sized_and_unsized_views_keep_their_metadata_carrier() { + for zero_sized in [false, true] { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let length = types.scalar(ScalarType::U64); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let pointee = if zero_sized { + types.intern(Type::Opaque(0)) + } else { + let name = types.symbol("test/Dst"); + types.intern(Type::Interface(name)) + }; + let pointer = types.intern(Type::Pointer(pointee)); + let mut b = Builder::new(&types, pointer); + let backing = b.parameter(b.current(), object); + let start = b.parameter(b.current(), int); + let count = b.parameter(b.current(), length); + let parts = b.args([backing, start, count]); + let address = b + .emit( + Op::ViewAddress { + parts, + size: if zero_sized { 0 } else { 1 }, + codec: None, + }, + Some(pointer), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(address))); + let mut body = b.finish().unwrap(); + lower_typed_addresses(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + assert!( + body.instructions + .iter() + .any(|inst| matches!(inst.op, Op::ViewAddress { .. })) + ); + assert!(live(&body, &types).values[count.index()]); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} + +#[test] +fn slice_to_fixed_array_to_scalar_slice_keeps_the_source_byte_displacement() { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let int = types.scalar(ScalarType::I32); + let length = types.scalar(ScalarType::U64); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let byte_pointer = types.intern(Type::Pointer(byte)); + let array = types.intern(Type::Array(byte)); + let codec = types.symbol("org/rustlang/runtime/ArrayMemoryCodec#array#[B#8"); + let layout = types.layout(AddressLayout { + value: array, + size: 8, + codec: Some(codec), + }); + let fixed_pointer = types.intern(Type::Pointer(layout)); + let erased_fixed = types.intern(Type::Pointer(array)); + let mut b = Builder::new(&types, byte); + let backing = b.parameter(b.current(), object); + let start = b.parameter(b.current(), int); + let count = b.constant(length, Scalar::integer(ScalarType::U64, 8).unwrap()); + let parts = b.args([backing, start, count]); + let data = b + .emit( + Op::ViewAddress { + parts, + size: 1, + codec: None, + }, + Some(byte_pointer), + ) + .unwrap(); + let fixed = b + .emit( + Op::RetypeAddress { + pointer: data, + size: 8, + codec: Some(codec), + }, + Some(fixed_pointer), + ) + .unwrap(); + let erased = b.emit(Op::Reinterpret(fixed), Some(erased_fixed)).unwrap(); + let bytes = b.emit(Op::Cast(erased), Some(byte_pointer)).unwrap(); + let root = b + .emit( + Op::AddressViewPart { + address: bytes, + index: 0, + }, + Some(object), + ) + .unwrap(); + let offset = b + .emit( + Op::AddressViewPart { + address: bytes, + index: 1, + }, + Some(int), + ) + .unwrap(); + let zero = b.constant(int, Scalar::integer(ScalarType::I32, 0).unwrap()); + let parts = b.args([root, offset, zero]); + let first = b.emit(Op::ViewGet(parts), Some(byte)).unwrap(); + b.terminate(Terminator::Return(Some(first))); + let mut body = b.finish().unwrap(); + lower_typed_addresses(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + simplify_components(&mut body, &types); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!( + inst.op, + Op::AddressPack(_) + | Op::TypedAddressPack { .. } + | Op::RetypeAddress { .. } + | Op::ViewAddress { .. } + | Op::TypedAddressViewPart { .. } + )) + ); + assert!(body.instructions.iter().any(|inst| matches!( + inst.op, + Op::ViewRoot { + size: 1, + codec: None, + .. + } + ))); + // The eight-byte array layout must not scale the slice start by eight. + // Its folded byte displacement must equal i64(start). + let Op::Call { args, .. } = body + .instructions + .iter() + .find(|inst| { + matches!(inst.op, + Op::Call { method, .. } if body.methods[method.index()].name == "locationSliceOffset") + }) + .unwrap() + .op + else { + unreachable!() + }; + let displacement = body.resolve(body.args[args.start as usize + 1]); + let ValueDef::Inst(id) = body.values[displacement.index()].def else { + panic!() + }; + assert!(matches!(body.instructions[id.index()].op, Op::Cast(value) if value == start)); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} diff --git a/compiler-core/src/opt/mod.rs b/compiler-core/src/opt/mod.rs index 67997224..e4e1df64 100644 --- a/compiler-core/src/opt/mod.rs +++ b/compiler-core/src/opt/mod.rs @@ -47,6 +47,19 @@ mod return_abi_tests; mod address_parts; +mod typed_addresses; +pub use typed_addresses::lower_typed_addresses; +#[cfg(test)] +mod typed_addresses_tests; + +mod typed_loads; +pub use typed_loads::lower_typed_loads; +#[cfg(test)] +mod typed_loads_tests; + +#[cfg(test)] +mod array_view_tests; + mod simplify; mod unreachable; pub use simplify::simplify_components; diff --git a/compiler-core/src/opt/typed_addresses.rs b/compiler-core/src/opt/typed_addresses.rs new file mode 100644 index 00000000..d44e6c78 --- /dev/null +++ b/compiler-core/src/opt/typed_addresses.rs @@ -0,0 +1,470 @@ +//! Use an exact source-language view layout without allocating a cast address. +//! Facts are attached to operations, never inferred from shared JVM classes. +use crate::ir::*; +use crate::scalar::{BinaryOp, Scalar, ScalarType}; +use rustc_hash::FxHashMap; + +fn append(body: &mut Body, prefix: &mut Vec, op: Op, ty: TypeId) -> ValueId { + let (id, value) = super::append_value(body, op, ty); + prefix.push(id); + value +} +fn literal(body: &mut Body, prefix: &mut Vec, ty: TypeId, n: u32) -> ValueId { + let id = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar( + Scalar::integer(ScalarType::I64, n.into()).unwrap(), + )); + append(body, prefix, Op::Constant(id), ty) +} + +pub fn lower_typed_addresses(body: &mut Body, types: &mut Types, debug: Option<&mut DebugInfo>) { + if !body.instructions.iter().any(|i| { + matches!( + i.op, + Op::RetypeAddress { size: 1.., .. } | Op::ViewAddress { size: 1.., .. } + ) + }) && !body + .values + .iter() + .any(|v| types.address_layout(v.ty).is_some()) + { + return; + } + // Join exact layouts across aliases and CFG edges. Each value starts unseen, + // gains one layout, can lose its root guarantee, then becomes unknown. + // Process backedges without whole-body scans or origin sets. + #[derive(Clone, Copy, PartialEq, Eq)] + enum Fact { + Unseen, + // A rooted value retains the layout in its root. + // The component ABI can pass it without an exact type annotation. + Known { + size: u32, + codec: Option, + rooted: bool, + }, + Unknown, + } + let count = body.values.len(); + let predecessors = body.predecessors(); + let mut facts = vec![Fact::Unknown; count]; + let mut users = crate::analysis::ValueUsers::new(count); + for (index, value) in body.values.iter().enumerate() { + if !matches!(types.get(value.ty), Some(Type::Pointer(_))) { + continue; + } + match value.def { + ValueDef::Inst(id) => match body.instructions[id.index()].op { + Op::AddressPack(_) if types.address_layout(value.ty).is_some() => { + let (size, codec) = types.address_layout(value.ty).unwrap(); + facts[index] = Fact::Known { + size, + codec, + rooted: false, + }; + } + Op::RetypeAddress { + pointer, + size: size @ 1.., + codec, + } => { + facts[index] = Fact::Known { + size, + codec, + rooted: true, + }; + users.connect(pointer, index); + } + Op::ViewAddress { + size: size @ 1.., + codec, + .. + } if codec.is_some() + || types + .pointee(value.ty) + .and_then(|ty| StorageSlot::scalar(ty, types)) + .is_some_and(|slot| slot.size == size) => + { + facts[index] = Fact::Known { + size, + codec, + rooted: true, + }; + } + Op::Offset { + pointer: source, .. + } + | Op::Refine(source) + | Op::Reinterpret(source) => { + facts[index] = Fact::Unseen; + users.connect(source, index); + } + _ => {} + }, + ValueDef::Alias(source) => { + facts[index] = Fact::Unseen; + users.connect(source, index); + } + ValueDef::Param(block) + if block != body.entry && !predecessors[block.index()].is_empty() => + { + facts[index] = Fact::Unseen; + let position = body.blocks[block.index()] + .params + .iter() + .position(|p| p.index() == index) + .unwrap(); + for &(_, edge) in &predecessors[block.index()] { + users.connect(body.edges[edge.index()].args[position], index); + } + } + _ => {} + } + } + let mut pending = (0..count) + .filter(|&i| facts[i] != Fact::Unseen) + .collect::>(); + for phase in 0..2 { + while let Some(source) = pending.pop() { + for target in users.users(source) { + let previous = facts[target]; + let incoming = facts[source]; + let retype = match body.values[target].def { + ValueDef::Inst(id) => { + matches!(body.instructions[id.index()].op, Op::RetypeAddress { .. }) + } + _ => false, + }; + let merged = match (previous, incoming) { + ( + Fact::Known { + size, + codec, + rooted, + }, + fact, + ) if retype => Fact::Known { + size, + codec, + rooted: rooted + && matches!(fact, Fact::Known { size: s, codec: c, rooted: true } if s == size && c == codec), + }, + (Fact::Unseen, fact) => fact, + ( + Fact::Known { + size: a, + codec: ac, + rooted: ar, + }, + Fact::Known { + size: b, + codec: bc, + rooted: br, + }, + ) if a == b && ac == bc => Fact::Known { + size: a, + codec: ac, + rooted: ar && br, + }, + _ => Fact::Unknown, + }; + if previous != merged { + facts[target] = merged; + pending.push(target); + } + } + } + if phase == 0 { + for (index, fact) in facts.iter_mut().enumerate() { + if *fact == Fact::Unseen { + *fact = Fact::Unknown; + pending.push(index); + } + } + } + } + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let long = types.scalar(ScalarType::I64); + let mut components = vec![None::<[ValueId; 2]>; count]; + let mut roots = Vec::new(); + let mut joins = vec![Vec::new(); body.blocks.len()]; + for index in 0..count { + if !matches!(facts[index], Fact::Known { .. }) { + continue; + } + match body.values[index].def { + ValueDef::Param(block) if block != body.entry => { + let position = body.blocks[block.index()] + .params + .iter() + .position(|p| p.index() == index) + .unwrap(); + components[index] = Some([object, long].map(|ty| { + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty, + def: ValueDef::Param(block), + }); + body.blocks[block.index()].params.push(value); + value + })); + joins[block.index()].push((index, position)); + } + ValueDef::Inst(id) + if matches!(body.instructions[id.index()].op, Op::AddressPack(_)) => + { + let Op::AddressPack(parts) = body.instructions[id.index()].op else { + unreachable!() + }; + components[index] = Some(body.args[parts.range()].try_into().unwrap()); + } + ValueDef::Inst(id) + if matches!( + body.instructions[id.index()].op, + Op::RetypeAddress { .. } | Op::Offset { .. } | Op::ViewAddress { .. } + ) => + { + let mut pending = Vec::new(); + components[index] = + Some([object, long].map(|ty| append(body, &mut pending, Op::Nop, ty))); + roots.push((id, pending)); + } + _ => {} + } + } + let eligible = facts + .iter() + .map(|f| matches!(f, Fact::Known { .. })) + .collect::>(); + users.propagate(&eligible, &mut components); + let mut prefixes = FxHashMap::default(); + for (id, created) in roots { + let original = body.instructions[id.index()]; + let result = original.result.unwrap(); + let Fact::Known { + size, + codec, + rooted, + } = facts[result.index()] + else { + unreachable!(); + }; + let mut prefix = Vec::new(); + let source_parts = if let Op::ViewAddress { parts, .. } = original.op { + let values = &body.args[parts.range()]; + let (backing, start) = (values[0], values[1]); + let root = append( + body, + &mut prefix, + Op::ViewRoot { + backing, + size, + codec, + }, + object, + ); + let start = append(body, &mut prefix, Op::Cast(start), long); + let stride = literal(body, &mut prefix, long, size); + let offset = append( + body, + &mut prefix, + Op::Binary { + op: BinaryOp::Mul, + left: start, + right: stride, + }, + long, + ); + [root, offset] + } else { + let source = match original.op { + Op::RetypeAddress { pointer, .. } | Op::Offset { pointer, .. } => pointer, + Op::Refine(pointer) | Op::Reinterpret(pointer) => pointer, + _ => unreachable!(), + }; + // Repeated retypes can carry ZST or source-layout metadata. + // Keep the boundary unless the exact layouts match. + let source = body.resolve(source); + let source_parts = if matches!(facts.get(source.index()), + Some(&Fact::Known { size: s, codec: c, .. }) if s == size && c == codec) + { + components[source.index()] + } else { + None + }; + source_parts + .or_else(|| { + let parts = super::address_parts::address_parts(body, types, source)?; + let ValueDef::Inst(inst) = body.values[parts.value.index()].def else { + unreachable!() + }; + matches!(body.instructions[inst.index()].op, Op::AddressPack(_)) + .then(|| body.args[parts.parts.range()].try_into().unwrap()) + }) + .unwrap_or_else(|| { + [ + append(body, &mut prefix, Op::Reinterpret(source), object), + literal(body, &mut prefix, long, 0), + ] + }) + }; + let [root, mut offset] = source_parts; + if let Op::Offset { + offset: delta, + bytes, + .. + } = original.op + { + let mut delta = append(body, &mut prefix, Op::Cast(delta), long); + if !bytes { + let stride = literal(body, &mut prefix, long, size); + delta = append( + body, + &mut prefix, + Op::Binary { + op: BinaryOp::Mul, + left: delta, + right: stride, + }, + long, + ); + } + offset = append( + body, + &mut prefix, + Op::Binary { + op: BinaryOp::Add, + left: offset, + right: delta, + }, + long, + ); + } + body.instructions[created[0].index()].op = Op::Reinterpret(root); + body.instructions[created[1].index()].op = Op::Reinterpret(offset); + prefix.extend(created); + let parts = List::append(&mut body.args, components[result.index()].unwrap()); + body.instructions[id.index()].op = + if rooted || types.address_layout(body.value_type(result)) == Some((size, codec)) { + Op::AddressPack(parts) + } else { + Op::TypedAddressPack { parts, size, codec } + }; + prefixes.insert(id, prefix); + } + for (block, incoming) in predecessors.iter().enumerate() { + for &(_, edge) in incoming { + for &(_, position) in &joins[block] { + let input = body.edges[edge.index()].args[position]; + body.edges[edge.index()] + .args + .extend(components[input.index()].unwrap()); + } + } + } + for index in 0..body.instructions.len() { + let op = body.instructions[index].op; + if let Op::AddressViewPart { + address, + index: part, + } = op + && let Some(&Fact::Known { size, codec, .. }) = facts.get(address.index()) + && let Some(source) = components.get(address.index()).copied().flatten() + { + // Use ordinary scalar normalization for scalar components. + // It can retain a primitive backing array without a carrier. + if codec.is_none() + && ComponentShape::of(types, body.value_type(address)) + == Some(ComponentShape::Address) + { + continue; + } + body.instructions[index].op = Op::TypedAddressViewPart { + parts: List::append(&mut body.args, source), + size, + codec, + index: part, + }; + continue; + } + if let Op::AddressPart { + address, + index: part, + } = op + && (types.address_layout(body.value_type(address)).is_some() + || matches!( + facts.get(address.index()), + Some(Fact::Known { rooted: true, .. }) + )) + && let Some(parts) = components.get(address.index()).copied().flatten() + { + body.instructions[index].op = Op::Reinterpret(parts[part as usize]); + continue; + } + let pointer = match op { + Op::LoadCopy(pointer) => pointer, + Op::Store { pointer, value } + if types + .get(body.value_type(value)) + .is_some_and(|t| t.carrier() == 5) => + { + pointer + } + _ => continue, + }; + if let Some(&Fact::Known { size, codec, .. }) = facts.get(pointer.index()) + && let Some(parts) = components[pointer.index()] + { + body.instructions[index].op = match op { + Op::LoadCopy(_) => Op::LoadTypedCopy { + parts: List::append(&mut body.args, parts), + size, + codec, + }, + Op::Store { value, .. } => Op::StoreTyped { + parts: List::append(&mut body.args, [parts[0], parts[1], value]), + size, + codec, + }, + _ => unreachable!(), + }; + } + } + + let mut positions = debug.as_ref().map(|_| Vec::new()); + for block in &mut body.blocks { + let previous = std::mem::take(&mut block.instructions); + let mut mapping = Vec::new(); + for id in previous { + if positions.is_some() { + mapping.push(block.instructions.len() as u32); + } + if let Some(prefix) = prefixes.remove(&id) { + block.instructions.extend(prefix); + } + block.instructions.push(id); + } + if let Some(positions) = &mut positions { + mapping.push(block.instructions.len() as u32); + positions.push(mapping); + } + if let Some(Terminator::Invoke { inst, normal, .. }) = block.terminator { + if let Some(prefix) = prefixes.remove(&inst) { + block.instructions.extend(prefix); + } + if matches!( + body.instructions[inst.index()].op, + Op::TypedAddressPack { .. } | Op::AddressPack(_) + ) { + block.instructions.push(inst); + block.terminator = Some(Terminator::Jump(normal)); + } + } + } + if let (Some(debug), Some(positions)) = (debug, positions) { + for event in &mut debug.events { + event.position = positions[event.block.index()][event.position as usize]; + } + } +} diff --git a/compiler-core/src/opt/typed_addresses_tests.rs b/compiler-core/src/opt/typed_addresses_tests.rs new file mode 100644 index 00000000..bfacfb62 --- /dev/null +++ b/compiler-core/src/opt/typed_addresses_tests.rs @@ -0,0 +1,452 @@ +use super::*; +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; + +#[test] +fn aggregate_slice_addresses_cross_calls_returns_and_fields_as_components() { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let long = types.scalar(ScalarType::I64); + let pair_name = types.symbol("Pair"); + let pair = types.intern(Type::Class(pair_name)); + let pointer = types.intern(Type::Pointer(pair)); + let slice = types.intern(Type::Slice(pair)); + let codec = types.symbol("PairCodec#pair#LPair;"); + let object_name = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object_name)); + let component_types = ComponentShape::View.parts(&mut types); + let mut b = Builder::new(&types, pointer); + let input = b.parameter(b.current(), slice); + let parts = component_types + .enumerate() + .map(|(index, ty)| { + b.emit( + Op::ViewPart { + view: input, + index: index as u8, + }, + Some(ty), + ) + .unwrap() + }) + .collect::>(); + let args = b.args(parts); + let first = b + .emit( + Op::ViewAddress { + parts: args, + size: 8, + codec: Some(codec), + }, + Some(pointer), + ) + .unwrap(); + let one = b.constant(long, Scalar::integer(ScalarType::I64, 1).unwrap()); + let next = b + .emit( + Op::Offset { + pointer: first, + offset: one, + bytes: false, + wrapping: false, + }, + Some(pointer), + ) + .unwrap(); + let next = b + .emit( + Op::RetypeAddress { + pointer: next, + size: 8, + codec: Some(codec), + }, + Some(pointer), + ) + .unwrap(); + let view_root = b + .emit( + Op::AddressViewPart { + address: next, + index: 0, + }, + Some(object), + ) + .unwrap(); + let view_start = b + .emit( + Op::AddressViewPart { + address: next, + index: 1, + }, + Some(int), + ) + .unwrap(); + let field = b.field(FieldRef { + owner: pair, + name: "second".into(), + ty: int, + is_static: false, + }); + let projection = b.projection(PointerProjection { + field, + offset: 4, + size: 4, + codec: None, + }); + let value = b + .emit( + Op::LoadField { + base: next, + projection, + }, + Some(int), + ) + .unwrap(); + let method = b.method(MethodRef { + owner: "Kernel".into(), + name: "consume".into(), + params: vec![pointer, int, object, int], + returns: pointer, + interface: false, + }); + let args = b.args([next, value, view_root, view_start]); + let result = b + .emit( + Op::Call { + method, + kind: CallKind::RustStatic, + args, + }, + Some(pointer), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| true, None); + lower_component_returns(&mut body, &mut types, true, |_| true, None); + decompose_views(&mut body, &mut types, None); + lower_typed_addresses(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + simplify_components(&mut body, &types); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!( + inst.op, + Op::ViewAddress { .. } + | Op::Offset { .. } + | Op::AddressPack(_) + | Op::TypedAddressPack { .. } + | Op::LoadField { .. } + )) + ); + assert_eq!( + body.instructions + .iter() + .filter(|i| matches!(i.op, Op::ViewRoot { .. })) + .count(), + 1 + ); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::LoadStorageField { .. })) + ); + assert!(!body.methods.iter().any(|m| m.name == "locationStride")); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +#[test] +fn aggregate_copy_uses_static_layout_without_cast_or_offset_carriers() { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let bytes = types.intern(Type::Pointer(byte)); + let name = types.symbol("Pair"); + let pair = types.intern(Type::Class(name)); + let pointer = types.intern(Type::Pointer(pair)); + let long = types.scalar(ScalarType::I64); + let codec = types.symbol("PairCodec#pair#LPair;"); + let mut b = Builder::new(&types, pair); + let source = b.parameter(b.current(), bytes); + let cast = b + .emit( + Op::RetypeAddress { + pointer: source, + size: 8, + codec: Some(codec), + }, + Some(pointer), + ) + .unwrap(); + let one = b.constant(long, Scalar::integer(ScalarType::I64, 1).unwrap()); + let element = b + .emit( + Op::Offset { + pointer: cast, + offset: one, + bytes: false, + wrapping: false, + }, + Some(pointer), + ) + .unwrap(); + let replacement = b.parameter(b.current(), pair); + b.emit( + Op::Store { + pointer: element, + value: replacement, + }, + None, + ); + let copy = b.emit(Op::LoadCopy(element), Some(pair)).unwrap(); + b.terminate(Terminator::Return(Some(copy))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| false, None); + lower_typed_addresses(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!( + inst.op, + Op::RetypeAddress { .. } | Op::Offset { .. } | Op::AddressPack(_) + )) + ); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::LoadTypedCopy { size: 8, .. })) + ); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::StoreTyped { size: 8, .. })) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +#[test] +fn aggregate_cursor_keeps_components_through_backedges() { + use crate::scalar::BinaryOp; + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let bytes = types.intern(Type::Pointer(byte)); + let pair_name = types.symbol("Pair"); + let pair = types.intern(Type::Class(pair_name)); + let pointer = types.intern(Type::Pointer(pair)); + let long = types.scalar(ScalarType::I64); + let boolean = types.scalar(ScalarType::Bool); + let codec = types.symbol("PairCodec#pair#LPair;"); + let mut b = Builder::new(&types, pair); + let source = b.parameter(b.current(), bytes); + let end = b.parameter(b.current(), long); + let cast = b + .emit( + Op::RetypeAddress { + pointer: source, + size: 8, + codec: Some(codec), + }, + Some(pointer), + ) + .unwrap(); + let cursor = b.variable(pointer); + let position = b.variable(long); + b.define(cursor, cast); + let zero = b.constant(long, Scalar::integer(ScalarType::I64, 0).unwrap()); + b.define(position, zero); + let header = b.create_block(); + let step = b.create_block(); + let done = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + let current = b.read(cursor); + let index = b.read(position); + let test = b + .emit( + Op::Binary { + op: BinaryOp::Lt, + left: index, + right: end, + }, + Some(boolean), + ) + .unwrap(); + b.branch(test, step, done); + b.switch_to(step); + b.emit(Op::LoadCopy(current), Some(pair)); + let one = b.constant(long, Scalar::integer(ScalarType::I64, 1).unwrap()); + let next = b + .emit( + Op::Offset { + pointer: current, + offset: one, + bytes: false, + wrapping: true, + }, + Some(pointer), + ) + .unwrap(); + b.define(cursor, next); + let next = b + .emit( + Op::Binary { + op: BinaryOp::Add, + left: index, + right: one, + }, + Some(long), + ) + .unwrap(); + b.define(position, next); + b.jump(header, vec![]); + b.switch_to(done); + let result = b.emit(Op::LoadCopy(current), Some(pair)).unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| false, None); + lower_typed_addresses(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!( + inst.op, + Op::RetypeAddress { .. } + | Op::Offset { .. } + | Op::AddressPack(_) + | Op::TypedAddressPack { .. } + | Op::LoadCopy(_) + )) + ); + assert_eq!( + body.instructions + .iter() + .filter(|i| matches!(i.op, Op::LoadTypedCopy { .. })) + .count(), + 2 + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +#[test] +fn conflicting_pointee_layouts_keep_a_boundary_at_the_join() { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let bytes = types.intern(Type::Pointer(byte)); + let name = types.symbol("Pair"); + let pair = types.intern(Type::Class(name)); + let pointer = types.intern(Type::Pointer(pair)); + let boolean = types.scalar(ScalarType::Bool); + let mut b = Builder::new(&types, pair); + let source = b.parameter(b.current(), bytes); + let condition = b.parameter(b.current(), boolean); + let result = b.variable(pointer); + let left = b.create_block(); + let right = b.create_block(); + let join = b.create_block(); + b.branch(condition, left, right); + for (block, size) in [(left, 8), (right, 16)] { + b.switch_to(block); + let value = b + .emit( + Op::RetypeAddress { + pointer: source, + size, + codec: None, + }, + Some(pointer), + ) + .unwrap(); + b.define(result, value); + b.jump(join, vec![]); + } + b.switch_to(join); + let pointer = b.read(result); + let value = b.emit(Op::LoadCopy(pointer), Some(pair)).unwrap(); + b.terminate(Terminator::Return(Some(value))); + let mut body = b.finish().unwrap(); + lower_typed_addresses(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + assert!( + !body + .instructions + .iter() + .any(|i| matches!(i.op, Op::LoadTypedCopy { .. })) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +#[test] +fn exact_array_layout_survives_call_parameters_and_returns() { + let mut types = Types::default(); + let byte = types.scalar(ScalarType::U8); + let array = types.intern(Type::Array(byte)); + let codec = types.symbol("ArrayCodec#eight#[B#8"); + let layout = types.layout(AddressLayout { + value: array, + size: 8, + codec: Some(codec), + }); + let pointer = types.intern(Type::Pointer(layout)); + let long = types.scalar(ScalarType::I64); + let mut b = Builder::new(&types, pointer); + let input = b.parameter(b.current(), pointer); + let one = b.constant(long, Scalar::integer(ScalarType::I64, 1).unwrap()); + let next = b + .emit( + Op::Offset { + pointer: input, + offset: one, + bytes: false, + wrapping: false, + }, + Some(pointer), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(next))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| true, None); + lower_component_returns(&mut body, &mut types, true, |_| true, None); + lower_typed_addresses(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + super::simplify_components(&mut body, &types); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, instruction)| live.instructions[i] + && matches!( + instruction.op, + Op::AddressPack(_) + | Op::TypedAddressPack { .. } + | Op::RetypeAddress { .. } + | Op::Offset { .. } + )) + ); + assert!(!body.methods.iter().any(|m| m.name == "locationStride")); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!( + !code + .instructions + .contains(&crate::classfile::attributes::Instruction::Lmul) + ); +} diff --git a/compiler-core/src/opt/typed_loads.rs b/compiler-core/src/opt/typed_loads.rs new file mode 100644 index 00000000..462bde6d --- /dev/null +++ b/compiler-core/src/opt/typed_loads.rs @@ -0,0 +1,103 @@ +//! Remove exact-layout carriers used only for one borrowed object read. +//! Shared carriers retain their binding for later commits. +use super::address_parts::address_parts; +use crate::ir::*; + +fn sole_view_consumer(body: &Body, uses: &[u8], mut value: ValueId) -> bool { + for _ in 0..32 { + value = body.resolve(value); + if uses[value.index()] != 1 { + return false; + } + let ValueDef::Inst(id) = body.values[value.index()].def else { + return false; + }; + match body.instructions[id.index()].op { + Op::AddressPack(_) | Op::TypedAddressPack { .. } => return true, + Op::Refine(source) | Op::Reinterpret(source) => value = source, + _ => return false, + } + } + false +} + +pub fn lower_typed_loads(body: &mut Body, types: &Types) { + let candidates = body + .instructions + .iter() + .enumerate() + .filter_map(|(index, inst)| { + let result = body.value_type(inst.result?); + if let Op::Load(pointer) = inst.op { + let eligible = match types.get(result) { + Some(Type::Class(name) | Type::Interface(name)) => { + types.symbol_name(name) != Some("java/lang/Object") + } + Some(Type::Array(_)) => true, + _ => false, + }; + return eligible.then_some((index, pointer, None)); + } + let Op::Call { + method, + kind: CallKind::Virtual, + args, + } = inst.op + else { + return None; + }; + let method = &body.methods[method.index()]; + let named = method.name == "getObjectAs" && args.len == 2; + // getObject's null target differs from getObjectAs(Object). + // Preserve this distinction when inferring load targets. + let untyped = method.name == "getObject" + && args.len == 1 + && matches!(types.get(result), Some(Type::Class(name)) + if types.symbol_name(name) == Some("java/lang/Object")); + (method.owner == "org/rustlang/runtime/Pointer" && (named || untyped)).then(|| { + ( + index, + body.args[args.start as usize], + named.then(|| body.args[args.start as usize + 1]), + ) + }) + }) + .collect::>(); + if candidates.is_empty() { + return; + } + let mut uses = vec![0_u8; body.values.len()]; + let mut use_value = |value: ValueId| { + let count = &mut uses[body.resolve(value).index()]; + *count = count.saturating_add(1); + }; + for inst in &body.instructions { + inst.op.visit_uses(&body.args, &mut use_value); + } + for block in &body.blocks { + block.terminator.unwrap().visit_uses(&mut use_value); + } + for edge in &body.edges { + for &value in &edge.args { + use_value(value); + } + } + for (index, receiver, target) in candidates { + if !sole_view_consumer(body, &uses, receiver) { + continue; + } + let Some(address) = address_parts(body, types, receiver) else { + continue; + }; + let Some((size, codec)) = address.layout else { + continue; + }; + let mut parts = body.args[address.parts.range()].to_vec(); + parts.extend(target); + body.instructions[index].op = Op::LoadTyped { + parts: List::append(&mut body.args, parts), + size, + codec, + }; + } +} diff --git a/compiler-core/src/opt/typed_loads_tests.rs b/compiler-core/src/opt/typed_loads_tests.rs new file mode 100644 index 00000000..550cb236 --- /dev/null +++ b/compiler-core/src/opt/typed_loads_tests.rs @@ -0,0 +1,266 @@ +use super::*; +use crate::ir::*; +use crate::scalar::ScalarType; + +#[test] +fn typed_borrow_loads_preserve_unwind_edges_and_shared_bindings() { + assert!(std::mem::size_of::() <= 20); + for (named, commit) in [(false, false), (true, false), (true, true)] { + for protected in [false, true] { + let (mut body, types, packed) = borrowed_load(named, protected, commit); + lower_typed_loads(&mut body, &types); + verify(&body, &types).unwrap(); + assert_eq!(live(&body, &types).values[packed.index()], commit); + let lowered = body + .instructions + .iter() + .enumerate() + .find(|(_, inst)| matches!(inst.op, Op::LoadTyped { .. })); + if commit { + assert!(lowered.is_none()); + } else { + let (id, inst) = lowered.unwrap(); + let Op::LoadTyped { parts, size, codec } = inst.op else { + unreachable!() + }; + assert_eq!(parts.len, if named { 3 } else { 2 }); + assert_eq!(size, 16); + assert_eq!( + types.symbol_name(codec.unwrap()), + Some("test/EnumCodec#enum#Ltest/Enum;") + ); + if protected { + assert!(body.blocks.iter().any(|block| matches!(block.terminator, + Some(Terminator::Invoke { inst, .. }) if inst.index() == id))); + } + } + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } + } +} + +#[test] +fn semantic_object_borrows_fuse_but_copies_and_shared_views_do_not() { + use crate::classfile::attributes::Instruction; + + for kind in 0..6 { + for (copied, shared) in [(false, false), (false, true), (true, false)] { + if copied && kind == 5 { + continue; + } + for protected in [false, true] { + let mut types = Types::default(); + let long = types.scalar(ScalarType::I64); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let name = types.symbol("test/Variant"); + let payload = match kind { + 0 => types.intern(Type::Class(name)), + 1 => types.intern(Type::Interface(name)), + 2 => types.intern(Type::Array(long)), + 3 => object, + 4 => types.intern(Type::Pointer(long)), + _ => long, + }; + let pointer = types.intern(Type::Pointer(payload)); + let codec = types.symbol("test/EnumCodec#enum#Ltest/Enum;"); + let mut b = Builder::new(&types, payload); + let root = b.parameter(b.current(), object); + let offset = b.parameter(b.current(), long); + let parts = b.args([root, offset]); + let packed = b + .emit( + Op::TypedAddressPack { + parts, + size: 16, + codec: Some(codec), + }, + Some(pointer), + ) + .unwrap(); + let op = if copied { + Op::LoadCopy(packed) + } else { + Op::Load(packed) + }; + let handler = protected.then(|| b.create_block()); + let value = match handler { + Some(handler) => b.invoke(op, Some(payload), handler), + None => b.emit(op, Some(payload)), + } + .unwrap(); + if shared { + b.emit(Op::Commit(packed), None); + } + b.terminate(Terminator::Return(Some(value))); + if let Some(handler) = handler { + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + } + let mut body = b.finish().unwrap(); + lower_typed_loads(&mut body, &types); + verify(&body, &types).unwrap(); + let expected = kind < 3 && !copied && !shared; + assert_eq!(!live(&body, &types).values[packed.index()], expected); + let ValueDef::Inst(id) = body.values[value.index()].def else { + panic!() + }; + assert_eq!( + matches!(body.instructions[id.index()].op, Op::LoadTyped { .. }), + expected + ); + if protected { + assert!(body.blocks.iter().any(|block| matches!(block.terminator, + Some(Terminator::Invoke { inst, .. }) if inst == id))); + } + let mut pool = Default::default(); + let code = crate::jvm::select::compile(&body, &types, &mut pool).unwrap(); + let owner = pool.add_class("org/rustlang/runtime/Pointer").unwrap(); + let load = pool.add_method_ref(owner, "loadTypedStorage", + "(Ljava/lang/Object;JILjava/lang/String;Ljava/lang/String;)Ljava/lang/Object;").unwrap(); + assert_eq!( + code.instructions.contains(&Instruction::Invokestatic(load)), + expected + ); + if expected && kind < 2 { + let name = pool.add_name_string("test/Variant").unwrap(); + assert!(code.instructions.contains(&Instruction::Ldc_w(name))); + } + } + } + } +} + +#[test] +fn enum_downcasts_keep_the_entry_components_through_erased_annotations() { + let mut types = Types::default(); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let string = types.symbol("java/lang/String"); + let string = types.intern(Type::Class(string)); + let variant = types.symbol("test/Enum"); + let variant = types.intern(Type::Interface(variant)); + let pointer = types.intern(Type::Pointer(variant)); + let runtime = types.symbol("org/rustlang/runtime/Pointer"); + let runtime = types.intern(Type::Class(runtime)); + let codec = types.symbol("test/EnumCodec#enum#Ltest/Enum;"); + let mut b = Builder::new(&types, object); + let source = b.parameter(b.current(), pointer); + let target = b.parameter(b.current(), string); + let erased = b.emit(Op::Reinterpret(source), Some(runtime)).unwrap(); + let recovered = b.emit(Op::Reinterpret(erased), Some(pointer)).unwrap(); + let retyped = b + .emit( + Op::RetypeAddress { + pointer: recovered, + size: 16, + codec: Some(codec), + }, + Some(pointer), + ) + .unwrap(); + let method = b.method(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: "getObjectAs".into(), + params: vec![string], + returns: object, + interface: false, + }); + let args = b.args([retyped, target]); + let result = b + .emit( + Op::Call { + method, + kind: CallKind::Virtual, + args, + }, + Some(object), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + lower_component_arguments(&mut body, &mut types, true, |_| false, None); + lower_typed_addresses(&mut body, &mut types, None); + decompose_addresses(&mut body, &mut types, None); + lower_typed_loads(&mut body, &types); + verify(&body, &types).unwrap(); + let live = live(&body, &types); + assert!( + !body + .instructions + .iter() + .enumerate() + .any(|(i, inst)| live.instructions[i] + && matches!( + inst.op, + Op::AddressPack(_) | Op::TypedAddressPack { .. } | Op::RetypeAddress { .. } + )) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} + +fn borrowed_load(named: bool, protected: bool, commit: bool) -> (Body, Types, ValueId) { + let mut types = Types::default(); + let long = types.scalar(ScalarType::I64); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let string = types.symbol("java/lang/String"); + let string = types.intern(Type::Class(string)); + let variant = types.symbol("test/Enum"); + let variant = types.intern(Type::Interface(variant)); + let pointer = types.intern(Type::Pointer(variant)); + let runtime = types.symbol("org/rustlang/runtime/Pointer"); + let runtime = types.intern(Type::Class(runtime)); + let codec = types.symbol("test/EnumCodec#enum#Ltest/Enum;"); + let mut b = Builder::new(&types, object); + let root = b.parameter(b.current(), object); + let offset = b.parameter(b.current(), long); + let target = b.parameter(b.current(), string); + let parts = b.args([root, offset]); + let packed = b + .emit( + Op::TypedAddressPack { + parts, + size: 16, + codec: Some(codec), + }, + Some(pointer), + ) + .unwrap(); + let erased = b.emit(Op::Reinterpret(packed), Some(runtime)).unwrap(); + let method = b.method(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: if named { "getObjectAs" } else { "getObject" }.into(), + params: if named { vec![string] } else { vec![] }, + returns: object, + interface: false, + }); + let args = b.args(if named { + vec![erased, target] + } else { + vec![erased] + }); + let call = Op::Call { + method, + kind: CallKind::Virtual, + args, + }; + let (result, handler) = if protected { + let handler = b.create_block(); + ( + b.invoke(call, Some(object), handler).unwrap(), + Some(handler), + ) + } else { + (b.emit(call, Some(object)).unwrap(), None) + }; + if commit { + b.emit(Op::Commit(packed), None); + } + b.terminate(Terminator::Return(Some(result))); + if let Some(handler) = handler { + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + } + (b.finish().unwrap(), types, packed) +} From 2f3214a7418087add223bec55137a8e2ea6f29af Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 11:37:45 +1000 Subject: [PATCH 21/61] promote private aggregate values --- compiler-core/src/opt/aggregate_joins.rs | 226 ++++++++ compiler-core/src/opt/aggregates.rs | 193 +++++++ compiler-core/src/opt/aggregates_tests.rs | 623 ++++++++++++++++++++++ compiler-core/src/opt/mod.rs | 7 + compiler-core/src/opt/value_copies.rs | 290 ++++++++++ 5 files changed, 1339 insertions(+) create mode 100644 compiler-core/src/opt/aggregate_joins.rs create mode 100644 compiler-core/src/opt/aggregates.rs create mode 100644 compiler-core/src/opt/aggregates_tests.rs create mode 100644 compiler-core/src/opt/value_copies.rs diff --git a/compiler-core/src/opt/aggregate_joins.rs b/compiler-core/src/opt/aggregate_joins.rs new file mode 100644 index 00000000..b419d60e --- /dev/null +++ b/compiler-core/src/opt/aggregate_joins.rs @@ -0,0 +1,226 @@ +//! Promote read-only carriers at joins as values, without allocation identity. +//! Join aliases and CFG inputs with union-find. Unknown producers, writes +//! and escaping uses reject the connected group. +use super::aggregates::Candidate; +use crate::ir::*; + +struct Groups { + parents: Vec, + ranks: Vec, +} +impl Groups { + fn root(&mut self, mut value: usize) -> usize { + while self.parents[value] != value { + self.parents[value] = self.parents[self.parents[value]]; + value = self.parents[value]; + } + value + } + fn join(&mut self, a: usize, b: usize) { + let a = self.root(a); + let b = self.root(b); + if a == b { + return; + } + let (a, b) = if self.ranks[a] < self.ranks[b] { + (b, a) + } else { + (a, b) + }; + self.parents[b] = a; + if self.ranks[a] == self.ranks[b] { + self.ranks[a] += 1; + } + } +} + +pub(super) fn promote( + body: &mut Body, + roots: &mut Vec, + candidates: &[Candidate], +) -> Vec { + let mut removed = vec![false; candidates.len()]; + if !body + .blocks + .iter() + .enumerate() + .any(|(index, block)| BlockId::new(index) != body.entry && !block.params.is_empty()) + { + return removed; + } + let count = body.values.len(); + let mut groups = Groups { + parents: (0..count).collect(), + ranks: vec![0; count], + }; + let mut users = crate::analysis::ValueUsers::new(count); + let mut known = vec![false; count]; + for (index, value) in body.values.iter().enumerate() { + if roots[index] != u32::MAX { + known[index] = true; + continue; + } + let source = match value.def { + ValueDef::Alias(source) => Some(source), + ValueDef::Inst(id) => match body.instructions[id.index()].op { + Op::Reinterpret(source) | Op::Refine(source) => Some(source), + _ => None, + }, + ValueDef::Param(block) if block != body.entry => { + known[index] = true; + None + } + _ => None, + }; + if let Some(source) = source { + groups.join(index, source.index()); + users.connect(source, index); + known[index] = true; + } + } + for edge in &body.edges { + for (&source, &target) in edge + .args + .iter() + .zip(&body.blocks[edge.target.index()].params) + { + groups.join(source.index(), target.index()); + } + } + let group = (0..count).map(|i| groups.root(i)).collect::>(); + let mut layout = vec![None::; count]; + let mut valid = vec![true; count]; + for index in 0..count { + let group = group[index]; + valid[group] &= known[index]; + let root = roots[index]; + if root == u32::MAX { + continue; + } + let candidate = root as usize; + if candidates[candidate].fields.len() > 8 { + valid[group] = false; + } + if let Some(previous) = layout[group] { + valid[group] &= candidates[previous].fields == candidates[candidate].fields; + } else { + layout[group] = Some(candidate); + } + } + let position = |value: ValueId, field: MemberId| { + layout[group[value.index()]].and_then(|index| { + candidates[index] + .fields + .iter() + .position(|f| f == &body.fields[field.index()]) + }) + }; + for inst in &body.instructions { + inst.op.visit_uses(&body.args, |value| { + let allowed = match inst.op { + Op::Reinterpret(_) | Op::Refine(_) => true, + Op::GetField { object, field } => { + object == value && position(value, field).is_some() + } + _ => false, + }; + if !allowed { + valid[group[value.index()]] = false; + } + }); + } + for block in &body.blocks { + block + .terminator + .unwrap() + .visit_uses(|v| valid[group[v.index()]] = false); + } + let eligible = group + .iter() + .map(|&g| valid[g] && layout[g].is_some()) + .collect::>(); + if !eligible.iter().any(|&yes| yes) { + return removed; + } + let mut components = vec![None; count]; + let mut joins = Vec::new(); + for index in 0..count { + if !eligible[index] { + continue; + } + let root = roots[index]; + if root != u32::MAX { + components[index] = Some(candidates[root as usize].initial); + removed[root as usize] = true; + continue; + } + if let ValueDef::Param(block) = body.values[index].def { + let position = body.blocks[block.index()] + .params + .iter() + .position(|p| p.index() == index) + .unwrap(); + let fields = &candidates[layout[group[index]].unwrap()].fields; + let mut parts = Vec::with_capacity(fields.len()); + for field in fields { + let part = ValueId::new(body.values.len()); + body.values.push(Value { + ty: field.ty, + def: ValueDef::Param(block), + }); + body.blocks[block.index()].params.push(part); + parts.push(part); + } + components[index] = Some(List::append(&mut body.args, parts)); + joins.push((block, position)); + } + } + users.propagate(&eligible, &mut components); + let predecessors = body.predecessors(); + for (block, position) in joins { + for &(_, edge) in &predecessors[block.index()] { + let source = body.edges[edge.index()].args[position]; + let parts = components[source.index()].unwrap(); + body.edges[edge.index()] + .args + .extend_from_slice(&body.args[parts.range()]); + } + } + for inst in &mut body.instructions { + if let Some(value) = inst.result { + let root = roots[value.index()]; + if root != u32::MAX && removed[root as usize] { + let id = ConstId::new(body.constants.len()); + body.constants + .push(Constant::Null(body.values[value.index()].ty)); + inst.op = Op::Constant(id); + roots[value.index()] = u32::MAX; + continue; + } + } + if let Op::GetField { object, field } = inst.op + && eligible[object.index()] + { + let fields = &candidates[layout[group[object.index()]].unwrap()].fields; + let position = fields + .iter() + .position(|f| f == &body.fields[field.index()]) + .unwrap(); + let parts = components[object.index()].unwrap(); + inst.op = Op::Reinterpret(body.args[parts.range()][position]); + } + } + for block in &mut body.blocks { + if let Some(Terminator::Invoke { inst, normal, .. }) = block.terminator + && matches!( + body.instructions[inst.index()].op, + Op::Constant(_) | Op::Reinterpret(_) + ) + { + block.instructions.push(inst); + block.terminator = Some(Terminator::Jump(normal)); + } + } + roots.resize(body.values.len(), u32::MAX); + removed +} diff --git a/compiler-core/src/opt/aggregates.rs b/compiler-core/src/opt/aggregates.rs new file mode 100644 index 00000000..ad72c854 --- /dev/null +++ b/compiler-core/src/opt/aggregates.rs @@ -0,0 +1,193 @@ +//! Promote private generated carriers to SSA fields. +//! Require supplied constructor layouts. Arbitrary JVM constructors can run code. +use crate::ir::*; + +const NONE: u32 = u32::MAX; + +pub(super) struct Candidate { + pub(super) fields: Vec, + pub(super) initial: List, +} + +pub fn promote_aggregates( + mut body: Body, + types: &Types, + mut layout: impl FnMut(&MethodRef) -> Option>, +) -> Result { + super::value_copies::expand_flat_copies(&mut body, types, &mut layout); + if !body.instructions.iter().any(|inst| { + matches!( + inst.op, + Op::Call { + kind: CallKind::Constructor, + .. + } + ) + }) { + return Ok(body); + } + let mut roots = vec![NONE; body.values.len()]; + let mut candidates = Vec::new(); + for inst in &body.instructions { + let Op::Call { + method, + kind: CallKind::Constructor, + args, + } = inst.op + else { + continue; + }; + let Some(fields) = layout(&body.methods[method.index()]) else { + continue; + }; + if fields.len() != args.len as usize + || fields + .iter() + .zip(&body.args[args.range()]) + .any(|(f, &v)| f.ty != body.value_type(v)) + { + continue; + } + roots[inst.result.unwrap().index()] = candidates.len() as u32; + candidates.push(Candidate { + fields, + initial: args, + }); + } + if candidates.is_empty() { + return Ok(body); + } + let mut escaped = super::aggregate_joins::promote(&mut body, &mut roots, &candidates); + let origins = crate::analysis::origins(&body, &roots); + let position = |index: usize, field: MemberId| { + candidates[index] + .fields + .iter() + .position(|f| f == &body.fields[field.index()]) + }; + for inst in &body.instructions { + inst.op.visit_uses(&body.args, |value| { + let index = origins[value.index()]; + if index == NONE { + return; + } + let index = index as usize; + let allowed = match inst.op { + Op::Reinterpret(_) | Op::Refine(_) => true, + Op::GetField { object, field } => { + object == value && position(index, field).is_some() + } + Op::SetField { + object, + field, + value: stored, + } => object == value && stored != value && position(index, field).is_some(), + _ => false, + }; + escaped[index] |= !allowed; + }); + } + let mut escape = |value: ValueId| { + let index = origins[value.index()]; + if index != NONE { + escaped[index as usize] = true; + } + }; + for edge in &body.edges { + for (&value, ¶m) in edge + .args + .iter() + .zip(&body.blocks[edge.target.index()].params) + { + if origins[value.index()] != origins[param.index()] { + escape(value); + } + } + } + for block in &body.blocks { + block.terminator.unwrap().visit_uses(&mut escape); + } + if escaped.iter().all(|&e| e) { + return Ok(body); + } + let mut builder = Builder::from_body(body, types); + let variables: Vec<_> = candidates + .iter() + .enumerate() + .map(|(index, candidate)| { + (!escaped[index]).then(|| { + candidate + .fields + .iter() + .map(|field| builder.variable(field.ty)) + .collect::>() + }) + }) + .collect(); + for block_index in 0..builder.body.blocks.len() { + builder.switch_to(BlockId::new(block_index)); + let instructions = builder.body.blocks[block_index].instructions.clone(); + let invoke = match builder.body.blocks[block_index].terminator.unwrap() { + Terminator::Invoke { inst, .. } => Some(inst), + _ => None, + }; + for id in instructions.into_iter().chain(invoke) { + let inst = builder.body.instructions[id.index()]; + if let Some(result) = inst.result { + let index = roots[result.index()]; + if index != NONE + && let Some(vars) = &variables[index as usize] + { + let initial = candidates[index as usize].initial; + for (&var, pos) in vars.iter().zip(initial.range()) { + builder.define(var, builder.body.args[pos]); + } + let constant = ConstId::new(builder.body.constants.len()); + builder + .body + .constants + .push(Constant::Null(builder.body.value_type(result))); + builder.body.instructions[id.index()].op = Op::Constant(constant); + continue; + } + } + let (object, field) = match inst.op { + Op::GetField { object, field } | Op::SetField { object, field, .. } => { + (object, field) + } + _ => continue, + }; + let index = origins[object.index()]; + if index == NONE { + continue; + } + let Some(vars) = &variables[index as usize] else { + continue; + }; + let position = candidates[index as usize] + .fields + .iter() + .position(|f| f == &builder.body.fields[field.index()]) + .unwrap(); + let variable = vars[position]; + builder.body.instructions[id.index()].op = match inst.op { + Op::GetField { .. } => Op::Reinterpret(builder.read(variable)), + Op::SetField { value, .. } => { + builder.define(variable, value); + Op::Nop + } + _ => unreachable!(), + }; + } + if let Some(Terminator::Invoke { inst, normal, .. }) = + builder.body.blocks[block_index].terminator + && !builder.body.instructions[inst.index()] + .op + .may_throw(&builder.body, types) + { + builder.body.blocks[block_index].instructions.push(inst); + builder.body.blocks[block_index].terminator = Some(Terminator::Jump(normal)); + } + } + builder.finish() +} diff --git a/compiler-core/src/opt/aggregates_tests.rs b/compiler-core/src/opt/aggregates_tests.rs new file mode 100644 index 00000000..3c8be3eb --- /dev/null +++ b/compiler-core/src/opt/aggregates_tests.rs @@ -0,0 +1,623 @@ +use super::promote_aggregates; +use crate::ir::*; +use crate::scalar::{BinaryOp, Scalar, ScalarType}; + +#[test] +fn fresh_arrays_transfer_after_initialization_but_not_with_later_or_escaping_uses() { + for mode in ["last", "later", "escape", "loop", "invoke"] { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let boolean = types.scalar(ScalarType::Bool); + let array = types.intern(Type::Array(int)); + let mut b = Builder::new(&types, array); + let repeat = b.parameter(b.current(), boolean); + let length = b.constant(int, Scalar::integer(ScalarType::I32, 2).unwrap()); + let zero = b.constant(int, Scalar::integer(ScalarType::I32, 0).unwrap()); + let source = b.emit(Op::NewArray(length), Some(array)).unwrap(); + b.emit( + Op::ArraySet { + array: source, + index: zero, + value: length, + native: true, + }, + None, + ); + let value = b + .emit( + Op::ArrayGet { + array: source, + index: zero, + native: true, + }, + Some(int), + ) + .unwrap(); + if mode == "escape" { + let method = b.method(MethodRef { + owner: "Consumer".into(), + name: "retain".into(), + params: vec![array], + returns: int, + interface: false, + }); + let args = b.args([source]); + b.emit( + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }, + Some(int), + ); + } + if mode == "loop" { + let header = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + } + let handler = (mode == "invoke").then(|| b.create_block()); + let copy = if let Some(handler) = handler { + b.invoke(Op::CopyValue(source), Some(array), handler) + .unwrap() + } else { + b.emit(Op::CopyValue(source), Some(array)).unwrap() + }; + if mode == "later" { + // Mutating the transferred result must not change the old source. + b.emit( + Op::ArraySet { + array: copy, + index: zero, + value: zero, + native: true, + }, + None, + ); + let later = b + .emit( + Op::ArrayGet { + array: source, + index: zero, + native: true, + }, + Some(int), + ) + .unwrap(); + b.emit( + Op::ArraySet { + array: copy, + index: zero, + value: later, + native: true, + }, + None, + ); + } else { + b.emit( + Op::ArraySet { + array: copy, + index: zero, + value, + native: true, + }, + None, + ); + } + if mode == "loop" { + let header = b.current(); + let done = b.create_block(); + b.branch(repeat, header, done); + b.switch_to(done); + } + b.terminate(Terminator::Return(Some(copy))); + if let Some(handler) = handler { + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + } + let body = promote_aggregates(b.finish().unwrap(), &types, |_| None).unwrap(); + verify(&body, &types).unwrap(); + let copies = body + .instructions + .iter() + .filter(|i| matches!(i.op, Op::CopyValue(_))) + .count(); + assert_eq!(copies, usize::from(mode != "last"), "{mode}"); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} + +#[test] +fn redundant_array_copies_keep_escaping_and_loop_snapshots_distinct() { + for mode in ["local", "escape", "loop", "object"] { + let mut types = Types::default(); + let int = types.scalar(ScalarType::I32); + let boolean = types.scalar(ScalarType::Bool); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let element = if mode == "object" { object } else { int }; + let array = types.intern(Type::Array(element)); + let mut b = Builder::new(&types, array); + let input = b.parameter(b.current(), array); + let repeat = b.parameter(b.current(), boolean); + let first = b.emit(Op::CopyValue(input), Some(array)).unwrap(); + if mode == "escape" { + let method = b.method(MethodRef { + owner: "Consumer".into(), + name: "retain".into(), + params: vec![array], + returns: int, + interface: false, + }); + let args = b.args([first]); + b.emit( + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }, + Some(int), + ); + } + if mode == "loop" { + let header = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + } + let second = b.emit(Op::CopyValue(first), Some(array)).unwrap(); + if mode == "loop" { + let header = b.current(); + let done = b.create_block(); + b.branch(repeat, header, done); + b.switch_to(done); + } + b.terminate(Terminator::Return(Some(second))); + let body = promote_aggregates(b.finish().unwrap(), &types, |_| None).unwrap(); + verify(&body, &types).unwrap(); + let copies = body + .instructions + .iter() + .filter(|i| matches!(i.op, Op::CopyValue(_))) + .count(); + assert_eq!(copies, if mode == "local" { 1 } else { 2 }, "{mode}"); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} + +fn construct( + b: &mut Builder<'_>, + owner: TypeId, + fields: &[FieldRef], + values: &[ValueId], +) -> ValueId { + let method = b.method(MethodRef { + owner: "Pair".into(), + name: "".into(), + params: fields.iter().map(|f| f.ty).collect(), + returns: TypeId::new(0), + interface: false, + }); + let args = b.args(values.iter().copied()); + b.emit( + Op::Call { + method, + kind: CallKind::Constructor, + args, + }, + Some(owner), + ) + .unwrap() +} + +#[test] +fn private_mutable_fields_become_loop_parameters() { + let mut types = Types::default(); + types.intern(Type::Unit); + let int = types.scalar(ScalarType::I64); + let boolean = types.scalar(ScalarType::Bool); + let symbol = types.symbol("Pair"); + let owner = types.intern(Type::Class(symbol)); + let fields = ["count", "sum"].map(|name| FieldRef { + owner, + name: name.into(), + ty: int, + is_static: false, + }); + let mut b = Builder::new(&types, int); + let limit = b.parameter(b.current(), int); + let zero = b.constant(int, Scalar::integer(ScalarType::I64, 0).unwrap()); + let pair = construct(&mut b, owner, &fields, &[zero, zero]); + let count = b.field(fields[0].clone()); + let sum = b.field(fields[1].clone()); + let header = b.create_block(); + let step = b.create_block(); + let done = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + let n = b + .emit( + Op::GetField { + object: pair, + field: count, + }, + Some(int), + ) + .unwrap(); + let test = b + .emit( + Op::Binary { + op: BinaryOp::Lt, + left: n, + right: limit, + }, + Some(boolean), + ) + .unwrap(); + b.branch(test, step, done); + b.switch_to(step); + let previous = b + .emit( + Op::GetField { + object: pair, + field: sum, + }, + Some(int), + ) + .unwrap(); + let next = b + .emit( + Op::Binary { + op: BinaryOp::Add, + left: n, + right: previous, + }, + Some(int), + ) + .unwrap(); + b.emit( + Op::SetField { + object: pair, + field: sum, + value: next, + }, + None, + ); + let one = b.constant(int, Scalar::integer(ScalarType::I64, 1).unwrap()); + let n = b + .emit( + Op::Binary { + op: BinaryOp::Add, + left: n, + right: one, + }, + Some(int), + ) + .unwrap(); + b.emit( + Op::SetField { + object: pair, + field: count, + value: n, + }, + None, + ); + b.jump(header, vec![]); + b.switch_to(done); + let result = b + .emit( + Op::GetField { + object: pair, + field: sum, + }, + Some(int), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let body = promote_aggregates(b.finish().unwrap(), &types, |_| Some(fields.to_vec())).unwrap(); + verify(&body, &types).unwrap(); + assert_eq!(body.blocks[header.index()].params.len(), 2); + assert!(!body.instructions.iter().any(|inst| matches!( + inst.op, + Op::Call { .. } | Op::GetField { .. } | Op::SetField { .. } + ))); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!( + !code + .instructions + .iter() + .any(|op| matches!(op, ristretto_classfile::attributes::Instruction::New(_))) + ); +} + +#[test] +fn keeps_returned_carriers_and_unknown_constructors() { + let mut types = Types::default(); + types.intern(Type::Unit); + let int = types.scalar(ScalarType::I64); + let symbol = types.symbol("Pair"); + let owner = types.intern(Type::Class(symbol)); + let field = FieldRef { + owner, + name: "value".into(), + ty: int, + is_static: false, + }; + let mut b = Builder::new(&types, owner); + let value = b.parameter(b.current(), int); + let pair = construct(&mut b, owner, &[field.clone()], &[value]); + b.terminate(Terminator::Return(Some(pair))); + let body = b.finish().unwrap(); + assert_eq!( + body, + promote_aggregates(body.clone(), &types, |_| Some(vec![field.clone()])).unwrap() + ); + assert_eq!( + body, + promote_aggregates(body.clone(), &types, |_| None).unwrap() + ); +} + +#[test] +fn distinct_readonly_values_join_as_fields() { + let mut types = Types::default(); + types.intern(Type::Unit); + let int = types.scalar(ScalarType::I64); + let boolean = types.scalar(ScalarType::Bool); + let name = types.symbol("Pair"); + let pair = types.intern(Type::Class(name)); + let fields = ["a", "b"].map(|name| FieldRef { + owner: pair, + name: name.into(), + ty: int, + is_static: false, + }); + for escaping in [false, true] { + let mut b = Builder::new(&types, int); + let condition = b.parameter(b.current(), boolean); + let unknown = escaping.then(|| b.parameter(b.current(), pair)); + let left = b.create_block(); + let right = b.create_block(); + let join = b.create_block(); + let merged = b.parameter(join, pair); + b.branch(condition, left, right); + b.switch_to(left); + let a = b.constant(int, Scalar::integer(ScalarType::I64, 13).unwrap()); + let c = b.constant(int, Scalar::integer(ScalarType::I64, 17).unwrap()); + let value = construct(&mut b, pair, &fields, &[a, c]); + b.jump(join, vec![value]); + b.switch_to(right); + let a = b.constant(int, Scalar::integer(ScalarType::I64, 19).unwrap()); + let c = b.constant(int, Scalar::integer(ScalarType::I64, 23).unwrap()); + let value = unknown.unwrap_or_else(|| construct(&mut b, pair, &fields, &[a, c])); + b.jump(join, vec![value]); + b.switch_to(join); + let field = b.field(fields[1].clone()); + let result = b + .emit( + Op::GetField { + object: merged, + field, + }, + Some(int), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let body = + promote_aggregates(b.finish().unwrap(), &types, |_| Some(fields.to_vec())).unwrap(); + verify(&body, &types).unwrap(); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + let constructed = code + .instructions + .iter() + .any(|i| matches!(i, ristretto_classfile::attributes::Instruction::New(_))); + assert_eq!( + constructed, escaping, + "an unknown input must keep its identity" + ); + if !escaping { + assert_eq!(body.blocks[join.index()].params.len(), 3); + assert!( + !body + .instructions + .iter() + .any(|i| matches!(i.op, Op::GetField { .. } | Op::Call { .. })) + ); + } + } +} + +#[test] +fn new_loop_values_are_not_confused_with_previous_iterations() { + let mut types = Types::default(); + types.intern(Type::Unit); + let int = types.scalar(ScalarType::I64); + let boolean = types.scalar(ScalarType::Bool); + let name = types.symbol("Pair"); + let pair = types.intern(Type::Class(name)); + let fields = ["a", "b"].map(|name| FieldRef { + owner: pair, + name: name.into(), + ty: int, + is_static: false, + }); + for mutate in [false, true] { + let mut b = Builder::new(&types, int); + let again = b.parameter(b.current(), boolean); + let zero = b.constant(int, Scalar::integer(ScalarType::I64, 0).unwrap()); + let initial = construct(&mut b, pair, &fields, &[zero, zero]); + let cursor = b.variable(pair); + b.define(cursor, initial); + let header = b.create_block(); + let step = b.create_block(); + let done = b.create_block(); + b.jump(header, vec![]); + b.switch_to(header); + let previous = b.read(cursor); + b.branch(again, step, done); + b.switch_to(step); + let field = b.field(fields[1].clone()); + let value = b + .emit( + Op::GetField { + object: previous, + field, + }, + Some(int), + ) + .unwrap(); + let one = b.constant(int, Scalar::integer(ScalarType::I64, 1).unwrap()); + let next = b + .emit( + Op::Binary { + op: BinaryOp::Add, + left: value, + right: one, + }, + Some(int), + ) + .unwrap(); + let new = construct(&mut b, pair, &fields, &[value, next]); + if mutate { + b.emit( + Op::SetField { + object: previous, + field, + value: one, + }, + None, + ); + } + b.define(cursor, new); + b.jump(header, vec![]); + b.switch_to(done); + let result = b + .emit( + Op::GetField { + object: previous, + field, + }, + Some(int), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let body = + promote_aggregates(b.finish().unwrap(), &types, |_| Some(fields.to_vec())).unwrap(); + verify(&body, &types).unwrap(); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + let constructed = code + .instructions + .iter() + .any(|i| matches!(i, ristretto_classfile::attributes::Instruction::New(_))); + assert_eq!( + constructed, mutate, + "aliased mutable identities must not become snapshots" + ); + } +} + +#[test] +fn flat_copy_snapshots_survive_source_mutation_without_carriers() { + let mut types = Types::default(); + types.intern(Type::Unit); + let int = types.scalar(ScalarType::I64); + let symbol = types.symbol("Pair"); + let pair = types.intern(Type::Class(symbol)); + let fields = ["a", "b"].map(|name| FieldRef { + owner: pair, + name: name.into(), + ty: int, + is_static: false, + }); + let mut b = Builder::new(&types, int); + let seven = b.constant(int, Scalar::integer(ScalarType::I64, 7).unwrap()); + let nine = b.constant(int, Scalar::integer(ScalarType::I64, 9).unwrap()); + let original = construct(&mut b, pair, &fields, &[seven, nine]); + let copy = b.emit(Op::CopyValue(original), Some(pair)).unwrap(); + let second = b.emit(Op::CopyValue(copy), Some(pair)).unwrap(); + let field = b.field(fields[0].clone()); + b.emit( + Op::SetField { + object: original, + field, + value: nine, + }, + None, + ); + b.emit( + Op::SetField { + object: copy, + field, + value: nine, + }, + None, + ); + let result = b + .emit( + Op::GetField { + object: second, + field, + }, + Some(int), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let body = promote_aggregates(b.finish().unwrap(), &types, |_| Some(fields.to_vec())).unwrap(); + verify(&body, &types).unwrap(); + assert!(!body.instructions.iter().any(|i| matches!( + i.op, + Op::CopyValue(_) | Op::Call { .. } | Op::GetField { .. } | Op::SetField { .. } + ))); + let Some(Terminator::Return(Some(result))) = body.blocks[body.entry.index()].terminator else { + panic!("return"); + }; + let mut snapshot = body.resolve(result); + while let ValueDef::Inst(id) = body.values[snapshot.index()].def { + let Op::Reinterpret(source) = body.instructions[id.index()].op else { + break; + }; + snapshot = body.resolve(source); + } + assert_eq!(snapshot, body.resolve(seven)); + let code = crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + assert!( + !code + .instructions + .iter() + .any(|i| matches!(i, ristretto_classfile::attributes::Instruction::New(_))) + ); +} + +#[test] +fn nullable_or_nested_owned_copies_keep_the_general_operation() { + let mut types = Types::default(); + types.intern(Type::Unit); + let int = types.scalar(ScalarType::I64); + let symbol = types.symbol("Pair"); + let pair = types.intern(Type::Class(symbol)); + for nested in [false, true] { + let field = FieldRef { + owner: pair, + name: "a".into(), + ty: if nested { pair } else { int }, + is_static: false, + }; + let mut b = Builder::new(&types, pair); + let input = b.parameter(b.current(), pair); + let source = if nested { + construct(&mut b, pair, &[field.clone()], &[input]) + } else { + input + }; + let copy = b.emit(Op::CopyValue(source), Some(pair)).unwrap(); + b.terminate(Terminator::Return(Some(copy))); + let body = + promote_aggregates(b.finish().unwrap(), &types, |_| Some(vec![field.clone()])).unwrap(); + verify(&body, &types).unwrap(); + assert!( + body.instructions + .iter() + .any(|i| matches!(i.op, Op::CopyValue(_))) + ); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } +} diff --git a/compiler-core/src/opt/mod.rs b/compiler-core/src/opt/mod.rs index e4e1df64..1146ba85 100644 --- a/compiler-core/src/opt/mod.rs +++ b/compiler-core/src/opt/mod.rs @@ -39,6 +39,13 @@ pub use addresses::decompose_addresses; #[cfg(test)] mod addresses_tests; +mod aggregate_joins; +mod aggregates; +mod value_copies; +pub use aggregates::promote_aggregates; +#[cfg(test)] +mod aggregates_tests; + mod return_abi; pub use return_abi::lower_component_returns; diff --git a/compiler-core/src/opt/value_copies.rs b/compiler-core/src/opt/value_copies.rs new file mode 100644 index 00000000..3e1e6369 --- /dev/null +++ b/compiler-core/src/opt/value_copies.rs @@ -0,0 +1,290 @@ +//! Expose bounded scalar snapshots before aggregate promotion. +//! Supplied constructor schemas prove non-null values. Nullable boundaries +//! retain the original copy operation. +use crate::ir::*; +use rustc_hash::FxHashMap; + +/// Transfer a fresh primitive array at its last local use. +/// Earlier reads and writes are allowed. Escaping aliases and later uses are not. +/// Require allocation and transfer in the same block to prevent snapshot reuse across loop iterations. +fn reuse_array_snapshots(body: &mut Body, types: &Types) { + let mut candidates = Vec::new(); + for (index, inst) in body.instructions.iter().enumerate() { + let Op::CopyValue(source) = inst.op else { + continue; + }; + let source = body.resolve(source); + let ty = body.value_type(source); + let Some(Type::Array(element)) = types.get(ty) else { + continue; + }; + if body.value_type(inst.result.unwrap()) != ty + || !types + .get(element) + .is_some_and(|t| matches!(t, Type::Scalar(_)) && t.carrier() != 5) + { + continue; + } + let ValueDef::Inst(def) = body.values[source.index()].def else { + continue; + }; + if !matches!( + body.instructions[def.index()].op, + Op::NewArray(_) + | Op::CopyValue(_) + | Op::LoadCopy(_) + | Op::LoadFieldCopy { .. } + | Op::LoadStorageFieldCopy { .. } + | Op::LoadAddressCopy(_) + | Op::LoadTypedCopy { .. } + ) { + continue; + } + candidates.push((index, source, def)); + } + if candidates.is_empty() { + return; + } + let mut positions = vec![None; body.instructions.len()]; + for (block, data) in body.blocks.iter().enumerate() { + for (position, id) in data.instructions.iter().enumerate() { + positions[id.index()] = Some((block as u32, position as u32)); + } + } + let mut transfers = vec![None; body.values.len()]; + let mut invalid = vec![false; body.values.len()]; + for (index, source, def) in candidates { + let (Some((owner, _)), Some((block, _))) = (positions[def.index()], positions[index]) + else { + continue; + }; + if owner != block { + continue; + } + if transfers[source.index()] + .replace(InstId::new(index)) + .is_some() + { + invalid[source.index()] = true; + } + } + if transfers.iter().all(Option::is_none) { + return; + } + for (index, inst) in body.instructions.iter().enumerate() { + inst.op.visit_uses(&body.args, |value| { + let value = body.resolve(value); + let Some(copy) = transfers[value.index()] else { + return; + }; + if copy.index() == index { + return; + } + let local_array = match inst.op { + Op::ArrayGet { array, .. } + | Op::ArraySet { array, .. } + | Op::ArrayFill { array, .. } + | Op::ArrayLength(array) => body.resolve(array) == value, + _ => false, + }; + let before = match (positions[index], positions[copy.index()]) { + (Some((a, i)), Some((b, j))) => a == b && i < j, + _ => false, + }; + invalid[value.index()] |= !local_array || !before; + }); + } + let mut escape = |value| { + invalid[body.resolve(value).index()] = true; + }; + for block in &body.blocks { + if let Some(term) = block.terminator { + term.visit_uses(&mut escape); + } + } + for edge in &body.edges { + for &value in &edge.args { + escape(value); + } + } + for (source, copy) in transfers.into_iter().enumerate() { + if let Some(copy) = copy + && !invalid[source] + { + body.instructions[copy.index()].op = Op::Reinterpret(ValueId::new(source)); + } + } +} + +pub(super) fn expand_flat_copies( + body: &mut Body, + types: &Types, + layout: &mut impl FnMut(&MethodRef) -> Option>, +) { + if !body + .instructions + .iter() + .any(|i| matches!(i.op, Op::CopyValue(_))) + { + return; + } + reuse_array_snapshots(body, types); + let mut schemas = FxHashMap::)>::default(); + for inst in &body.instructions { + let Op::Call { + method, + kind: CallKind::Constructor, + .. + } = inst.op + else { + continue; + }; + let ty = body.value_type(inst.result.unwrap()); + if schemas.contains_key(&ty) { + continue; + } + let method_ref = &body.methods[method.index()]; + let Some(fields) = layout(method_ref) else { + continue; + }; + if fields.len() > 8 + || fields.len() != method_ref.params.len() + || !fields.iter().zip(&method_ref.params).all(|(field, &ty)| { + field.ty == ty && !field.is_static && matches!(types.get(ty), Some(Type::Scalar(_))) + }) + { + continue; + } + schemas.insert(ty, (method, fields)); + } + if schemas.is_empty() { + return; + } + // Join non-null facts across all incoming values. + // Each value changes from unseen to non-null to unknown at most once. + let count = body.values.len(); + let mut facts = vec![2u8; count]; + let mut users = crate::analysis::ValueUsers::new(count); + let predecessors = body.predecessors(); + for (index, value) in body.values.iter().enumerate() { + if !schemas.contains_key(&value.ty) { + continue; + } + match value.def { + ValueDef::Inst(id) => match body.instructions[id.index()].op { + Op::Call { + kind: CallKind::Constructor, + .. + } => facts[index] = 1, + Op::CopyValue(source) | Op::Reinterpret(source) | Op::Refine(source) + if body.value_type(source) == value.ty => + { + facts[index] = 0; + users.connect(source, index); + } + _ => {} + }, + ValueDef::Alias(source) => { + facts[index] = 0; + users.connect(source, index); + } + ValueDef::Param(block) + if block != body.entry && !predecessors[block.index()].is_empty() => + { + facts[index] = 0; + let position = body.blocks[block.index()] + .params + .iter() + .position(|v| v.index() == index) + .unwrap(); + for &(_, edge) in &predecessors[block.index()] { + users.connect(body.edges[edge.index()].args[position], index); + } + } + _ => {} + } + } + let mut pending = (0..count).filter(|&i| facts[i] != 0).collect::>(); + for phase in 0..2 { + while let Some(source) = pending.pop() { + for target in users.users(source) { + let next = facts[source].max(facts[target]); + if next != facts[target] { + facts[target] = next; + pending.push(target); + } + } + } + if phase == 0 { + for (i, fact) in facts.iter_mut().enumerate() { + if *fact == 0 { + *fact = 2; + pending.push(i); + } + } + } + } + let mut prefixes = FxHashMap::>::default(); + for index in 0..body.instructions.len() { + let Op::CopyValue(source) = body.instructions[index].op else { + continue; + }; + if facts[source.index()] != 1 { + continue; + } + let Some((method, fields)) = schemas.get(&body.value_type(source)) else { + continue; + }; + let mut prefix = Vec::with_capacity(fields.len()); + let args = fields + .iter() + .map(|field| { + let member = body + .fields + .iter() + .position(|f| f == field) + .map(MemberId::new) + .unwrap_or_else(|| { + let id = MemberId::new(body.fields.len()); + body.fields.push(field.clone()); + id + }); + let inst = InstId::new(body.instructions.len()); + let value = ValueId::new(body.values.len()); + body.values.push(Value { + ty: field.ty, + def: ValueDef::Inst(inst), + }); + body.instructions.push(Inst { + op: Op::GetField { + object: source, + field: member, + }, + result: Some(value), + }); + prefix.push(inst); + value + }) + .collect::>(); + body.instructions[index].op = Op::Call { + method: *method, + kind: CallKind::Constructor, + args: List::append(&mut body.args, args), + }; + prefixes.insert(InstId::new(index), prefix); + } + for block in &mut body.blocks { + let previous = std::mem::take(&mut block.instructions); + for id in previous { + if let Some(prefix) = prefixes.remove(&id) { + block.instructions.extend(prefix); + } + block.instructions.push(id); + } + if let Some(Terminator::Invoke { inst, .. }) = block.terminator { + if let Some(prefix) = prefixes.remove(&inst) { + block.instructions.extend(prefix); + } + } + } +} From 355b8ad5e1da09a5769eafabdf922378d254dddd Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 12:54:57 +1000 Subject: [PATCH 22/61] lower intrinsics to component addresses --- compiler-core/src/opt/address_intrinsics.rs | 68 ++++++++ .../src/opt/address_intrinsics_tests.rs | 132 ++++++++++++++ compiler-core/src/opt/address_observers.rs | 135 ++++++++++++++ .../src/opt/address_observers_tests.rs | 164 ++++++++++++++++++ compiler-core/src/opt/memory_copies.rs | 53 ++++++ compiler-core/src/opt/memory_copies_tests.rs | 100 +++++++++++ compiler-core/src/opt/mod.rs | 14 ++ 7 files changed, 666 insertions(+) create mode 100644 compiler-core/src/opt/address_intrinsics.rs create mode 100644 compiler-core/src/opt/address_intrinsics_tests.rs create mode 100644 compiler-core/src/opt/address_observers.rs create mode 100644 compiler-core/src/opt/address_observers_tests.rs create mode 100644 compiler-core/src/opt/memory_copies.rs create mode 100644 compiler-core/src/opt/memory_copies_tests.rs diff --git a/compiler-core/src/opt/address_intrinsics.rs b/compiler-core/src/opt/address_intrinsics.rs new file mode 100644 index 00000000..e4ea560b --- /dev/null +++ b/compiler-core/src/opt/address_intrinsics.rs @@ -0,0 +1,68 @@ +//! Use address components for intrinsics with explicit byte widths. +//! Runtime helpers retain synchronization and memory tracking. +use super::address_parts::address_parts; +use crate::ir::*; +use crate::scalar::ScalarType; + +fn arity(name: &str) -> Option { + Some(match name { + "atomicLoad" | "writeBytes" => 3, + "atomicStore" | "atomicExchange" | "atomicAdd" | "atomicSubtract" | "atomicAnd" + | "atomicNand" | "atomicOr" | "atomicXor" | "atomicMax" | "atomicMin" + | "atomicUnsignedMax" | "atomicUnsignedMin" => 4, + "atomicCompareExchange" => 6, + _ => return None, + }) +} + +pub fn lower_address_intrinsics(body: &mut Body, types: &mut Types) { + let mut component_types = None; + for index in 0..body.instructions.len() { + let Op::Call { + method, + kind: CallKind::JvmStatic, + args, + } = body.instructions[index].op + else { + continue; + }; + let target = &body.methods[method.index()]; + if target.owner != "org/rustlang/runtime/Pointer" + || arity(&target.name) != Some(args.len as usize) + || target.params.len() != args.len as usize + || !matches!(types.get(target.params[0]), Some(Type::Pointer(_))) + { + continue; + } + if target.name == "writeBytes" + && !matches!( + types.get(target.params[2]), + Some(Type::Scalar(ScalarType::I64 | ScalarType::U64)) + ) + { + continue; + } + let values = &body.args[args.range()]; + let Some(address) = address_parts(body, types, values[0]) else { + continue; + }; + let mut new_args = body.args[address.parts.range()].to_vec(); + new_args.extend_from_slice(&values[1..]); + let mut target = target.clone(); + let [object, long] = *component_types.get_or_insert_with(|| { + let object = types.symbol("java/lang/Object"); + [ + types.intern(Type::Class(object)), + types.scalar(ScalarType::I64), + ] + }); + target.params.splice(..1, [object, long]); + let method = MethodId::new(body.methods.len()); + body.methods.push(target); + body.instructions[index].op = Op::Call { + method, + kind: CallKind::JvmStatic, + args: List::append(&mut body.args, new_args), + }; + } +} diff --git a/compiler-core/src/opt/address_intrinsics_tests.rs b/compiler-core/src/opt/address_intrinsics_tests.rs new file mode 100644 index 00000000..a7e253e5 --- /dev/null +++ b/compiler-core/src/opt/address_intrinsics_tests.rs @@ -0,0 +1,132 @@ +use super::*; +use crate::{ir::*, scalar::ScalarType}; + +#[test] +fn intrinsic_calls_keep_address_components_and_unwind_edges() { + for name in [ + "atomicLoad", + "atomicStore", + "atomicExchange", + "atomicAdd", + "atomicSubtract", + "atomicAnd", + "atomicNand", + "atomicOr", + "atomicXor", + "atomicMax", + "atomicMin", + "atomicUnsignedMax", + "atomicUnsignedMin", + "atomicCompareExchange", + "writeBytes", + ] { + for decomposed in [false, true] { + for typed in [false, true] { + let mut types = Types::default(); + let long = types.scalar(ScalarType::I64); + let int = types.scalar(ScalarType::I32); + let void = types.intern(Type::Unit); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let pointer = types.intern(Type::Pointer(long)); + let result = if matches!(name, "atomicStore" | "writeBytes") { + void + } else { + long + }; + let mut b = Builder::new(&types, result); + let address = if decomposed { + let root = b.parameter(b.current(), object); + let offset = b.parameter(b.current(), long); + let parts = b.args([root, offset]); + b.emit( + if typed { + Op::TypedAddressPack { + parts, + size: 8, + codec: None, + } + } else { + Op::AddressPack(parts) + }, + Some(pointer), + ) + .unwrap() + } else { + b.parameter(b.current(), pointer) + }; + let mut params = vec![pointer]; + let mut values = vec![address]; + if name == "writeBytes" { + params.extend([int, long]); + values.extend([ + b.parameter(b.current(), int), + b.parameter(b.current(), long), + ]); + } else { + if name != "atomicLoad" { + params.push(long); + values.push(b.parameter(b.current(), long)); + } + if name == "atomicCompareExchange" { + params.push(long); + values.push(b.parameter(b.current(), long)); + } + for _ in 0..if name == "atomicCompareExchange" { + 3 + } else { + 2 + } { + params.push(int); + values.push(b.parameter(b.current(), int)); + } + } + let original_arity = params.len(); + let method = b.method(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: name.into(), + params, + returns: result, + interface: false, + }); + let args = b.args(values); + let handler = b.create_block(); + let result = b.invoke( + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }, + (result != void).then_some(result), + handler, + ); + b.terminate(Terminator::Return(result)); + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + let mut body = b.finish().unwrap(); + lower_address_intrinsics(&mut body, &mut types); + verify(&body, &types).unwrap(); + let (id, method, args) = body + .instructions + .iter() + .enumerate() + .find_map(|(id, inst)| { + let Op::Call { method, args, .. } = inst.op else { + return None; + }; + Some((id, &body.methods[method.index()], args)) + }) + .unwrap(); + assert_eq!(method.name, name); + assert_eq!(args.len as usize, original_arity + usize::from(decomposed)); + if decomposed { + assert!(!live(&body, &types).values[address.index()]); + assert_eq!(&method.params[..2], &[object, long]); + } + assert!(body.blocks.iter().any(|block| matches!(block.terminator, + Some(Terminator::Invoke { inst, .. }) if inst.index() == id))); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } + } + } +} diff --git a/compiler-core/src/opt/address_observers.rs b/compiler-core/src/opt/address_observers.rs new file mode 100644 index 00000000..8148714c --- /dev/null +++ b/compiler-core/src/opt/address_observers.rs @@ -0,0 +1,135 @@ +//! Observe address components without a Pointer carrier. +//! Keep calls in place to retain provenance and alignment checks under their original handlers. +use super::address_parts::{AddressParts, address_parts}; +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; +use rustc_hash::FxHashMap; + +const OWNER: &str = "org/rustlang/runtime/Pointer"; + +fn stride(address: &AddressParts, types: &Types) -> i64 { + if let Some((size, _)) = address.layout { + return i64::from(size); + } + match address.pointee.and_then(|ty| types.get(ty)) { + Some(Type::Pointer(_)) => 8, + Some(Type::Slice(_) | Type::Str) => 16, + // Resolve the root's layout in the runtime helper. + // Keep the original pointer operation's exception handler. + _ => -1, + } +} + +pub fn lower_address_observers(body: &mut Body, types: &mut Types, debug: Option<&mut DebugInfo>) { + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let long = types.scalar(ScalarType::I64); + let mut constants = FxHashMap::default(); + let mut prologue = Vec::new(); + let mut constant = |body: &mut Body, bits: i64| { + *constants.entry(bits).or_insert_with(|| { + let id = ConstId::new(body.constants.len()); + body.constants.push(Constant::Scalar( + Scalar::integer(ScalarType::I64, bits as u128).unwrap(), + )); + let id_inst = InstId::new(body.instructions.len()); + let result = ValueId::new(body.values.len()); + body.values.push(Value { + ty: long, + def: ValueDef::Inst(id_inst), + }); + body.instructions.push(Inst { + op: Op::Constant(id), + result: Some(result), + }); + prologue.push(id_inst); + result + }) + }; + for index in 0..body.instructions.len() { + let Op::Call { method, kind, args } = body.instructions[index].op else { + continue; + }; + let method = &body.methods[method.index()]; + if method.owner != OWNER + || !matches!(kind, CallKind::JvmStatic | CallKind::Virtual) + || !matches!( + types.get(method.returns), + Some(Type::Scalar(ScalarType::I64 | ScalarType::U64)) + ) + { + continue; + } + let name = match method.name.as_str() { + "addr" => "locationAddr", + "offset_from" | "offsetFrom" => "offsetLocations", + "offset_from_unsigned" => "offsetLocationsUnsigned", + "byte_offset_from" => "byteOffsetLocations", + "byte_offset_from_unsigned" => "byteOffsetLocationsUnsigned", + "align_offset" => "alignLocation", + _ => continue, + }; + let count = if name == "locationAddr" { 1 } else { 2 }; + if args.len as usize != count + || method.params.len() + usize::from(kind == CallKind::Virtual) != count + { + continue; + } + let returns = method.returns; + let args = body.args[args.range()].to_vec(); + let first = address_parts(body, types, args[0]); + let distance = name.contains("Offset") || name.starts_with("offset"); + let second = distance + .then(|| address_parts(body, types, args[1])) + .flatten(); + if first.is_none() || (distance && second.is_none()) { + continue; + } + if name == "alignLocation" && types.get(body.value_type(args[1])).unwrap().carrier() != 2 { + continue; + } + let step = first.as_ref().map_or(-1, |address| stride(address, types)); + let parts = |address: Option| -> [ValueId; 2] { + body.args[address.unwrap().parts.range()] + .try_into() + .unwrap() + }; + let mut values = parts(first).to_vec(); + let mut params = vec![object, long]; + if distance { + values.extend(parts(second)); + params.extend([object, long]); + } + if name.starts_with("offset") || name == "alignLocation" { + values.push(constant(body, step)); + params.push(long); + } + if name == "alignLocation" { + values.push(args[1]); + params.push(body.value_type(args[1])); + } + let method = MethodId::new(body.methods.len()); + body.methods.push(MethodRef { + owner: OWNER.into(), + name: name.into(), + params, + returns, + interface: false, + }); + body.instructions[index].op = Op::Call { + method, + kind: CallKind::JvmStatic, + args: List::append(&mut body.args, values), + }; + } + let count = prologue.len() as u32; + prologue.append(&mut body.blocks[body.entry.index()].instructions); + body.blocks[body.entry.index()].instructions = prologue; + if let Some(debug) = debug { + for event in &mut debug.events { + if event.block == body.entry { + event.position += count; + } + } + } +} diff --git a/compiler-core/src/opt/address_observers_tests.rs b/compiler-core/src/opt/address_observers_tests.rs new file mode 100644 index 00000000..80eac92e --- /dev/null +++ b/compiler-core/src/opt/address_observers_tests.rs @@ -0,0 +1,164 @@ +use super::*; +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; + +#[test] +fn address_observers_keep_components_and_exception_handlers() { + for name in [ + "addr", + "offset_from", + "offset_from_unsigned", + "byte_offset_from", + "byte_offset_from_unsigned", + "align_offset", + ] { + for typed in [false, true] { + for protected in [false, true] { + let mut types = Types::default(); + let long = types.scalar(ScalarType::I64); + let int = types.scalar(ScalarType::I32); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let pointer = types.intern(Type::Pointer(int)); + let codec = types.symbol("test/Codec#pair#Ltest/Pair;"); + let mut b = Builder::new(&types, long); + let root = b.parameter(b.current(), object); + let offset = b.parameter(b.current(), long); + let parts = b.args([root, offset]); + let packed = b + .emit( + if typed { + Op::TypedAddressPack { + parts, + size: 16, + codec: Some(codec), + } + } else { + Op::AddressPack(parts) + }, + Some(pointer), + ) + .unwrap(); + let alignment = b.constant(long, Scalar::integer(ScalarType::I64, 8).unwrap()); + let (params, args) = match name { + "addr" => (vec![], vec![packed]), + "align_offset" => (vec![long], vec![packed, alignment]), + _ => (vec![pointer], vec![packed, packed]), + }; + let method = b.method(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: name.into(), + params, + returns: long, + interface: false, + }); + let args = b.args(args); + let call = Op::Call { + method, + kind: CallKind::Virtual, + args, + }; + let (result, handler) = if protected { + let handler = b.create_block(); + (b.invoke(call, Some(long), handler).unwrap(), Some(handler)) + } else { + (b.emit(call, Some(long)).unwrap(), None) + }; + b.terminate(Terminator::Return(Some(result))); + if let Some(handler) = handler { + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + } + let mut body = b.finish().unwrap(); + lower_address_observers(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + assert!(!live(&body, &types).values[packed.index()], "{name}"); + let (id, inst) = body.instructions.iter().enumerate().find(|(_, inst)| { + matches!(inst.op, Op::Call { method, .. } if body.methods[method.index()].name != name) + }).unwrap(); + let Op::Call { method, kind, args } = inst.op else { + unreachable!() + }; + assert_eq!(kind, CallKind::JvmStatic); + let selected = &body.methods[method.index()].name; + assert!( + selected.ends_with("Locations") + || selected.ends_with("LocationsUnsigned") + || selected == "alignLocation" + || selected == "locationAddr" + ); + if name.starts_with("offset") || name == "align_offset" { + let size = + body.args[args.start as usize + if name == "align_offset" { 2 } else { 4 }]; + assert_eq!( + body.scalar_value(size).unwrap().bits(), + if typed { 16 } else { 4 } + ); + } + if protected { + assert!(body.blocks.iter().any(|block| matches!(block.terminator, + Some(Terminator::Invoke { inst, .. }) if inst.index() == id))); + } + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } + } + } +} + +#[test] +fn dynamic_stride_comes_from_the_root() { + let mut types = Types::default(); + let long = types.scalar(ScalarType::I64); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let record = types.symbol("test/Record"); + let record = types.intern(Type::Class(record)); + let pointer = types.intern(Type::Pointer(record)); + let mut b = Builder::new(&types, long); + let root = b.parameter(b.current(), object); + let offset = b.parameter(b.current(), long); + let origin = b.parameter(b.current(), object); + let parts = b.args([root, offset]); + let packed = b.emit(Op::AddressPack(parts), Some(pointer)).unwrap(); + let origin_parts = b.args([origin, offset]); + let origin = b + .emit(Op::AddressPack(origin_parts), Some(pointer)) + .unwrap(); + let method = b.method(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: "offsetFrom".into(), + params: vec![pointer, pointer], + returns: long, + interface: false, + }); + let args = b.args([packed, origin]); + let result = b + .emit( + Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }, + Some(long), + ) + .unwrap(); + b.terminate(Terminator::Return(Some(result))); + let mut body = b.finish().unwrap(); + lower_address_observers(&mut body, &mut types, None); + verify(&body, &types).unwrap(); + assert!(!live(&body, &types).values[packed.index()]); + let Op::Call { args, .. } = body + .instructions + .iter() + .find(|inst| matches!(inst.op, Op::Call { .. })) + .unwrap() + .op + else { + unreachable!() + }; + let args = &body.args[args.range()]; + assert!(!live(&body, &types).values[origin.index()]); + assert_eq!(args[3], offset); + assert_eq!(body.scalar_value(args[4]).unwrap().bits() as i64, -1); + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); +} diff --git a/compiler-core/src/opt/memory_copies.rs b/compiler-core/src/opt/memory_copies.rs new file mode 100644 index 00000000..bd78ec3f --- /dev/null +++ b/compiler-core/src/opt/memory_copies.rs @@ -0,0 +1,53 @@ +//! Pass address components to byte-copy intrinsics. +//! The runtime handles overlap, alias tracking and encoded references. +use super::address_parts::address_parts; +use crate::ir::*; + +pub fn lower_memory_copies(body: &mut Body, types: &mut Types) { + for index in 0..body.instructions.len() { + let inst = body.instructions[index]; + let Op::Call { + method, + kind: CallKind::JvmStatic, + args, + } = inst.op + else { + continue; + }; + let method = &body.methods[method.index()]; + if method.owner != "org/rustlang/runtime/Pointer" + || !matches!(method.name.as_str(), "copy" | "copyNonOverlapping") + || inst.result.is_some() + || args.len != 3 + { + continue; + } + let nonoverlapping = method.name == "copyNonOverlapping"; + let args = &body.args[args.range()]; + let count = args[2]; + let (Some(source), Some(destination)) = ( + address_parts(body, types, args[0]), + address_parts(body, types, args[1]), + ) else { + continue; + }; + let (Some((a, ac)), Some((b, bc))) = (source.layout, destination.layout) else { + continue; + }; + let layouts = [(args[0], a, ac), (args[1], b, bc)].map(|(value, size, codec)| { + types.layout(AddressLayout { + value: types.pointee(body.value_type(value)).unwrap(), + size, + codec, + }) + }); + let mut parts = body.args[source.parts.range()].to_vec(); + parts.extend_from_slice(&body.args[destination.parts.range()]); + parts.push(count); + body.instructions[index].op = Op::CopyStorage { + parts: List::append(&mut body.args, parts), + layouts, + nonoverlapping, + }; + } +} diff --git a/compiler-core/src/opt/memory_copies_tests.rs b/compiler-core/src/opt/memory_copies_tests.rs new file mode 100644 index 00000000..f8d34839 --- /dev/null +++ b/compiler-core/src/opt/memory_copies_tests.rs @@ -0,0 +1,100 @@ +use super::lower_memory_copies; +use crate::ir::*; +use crate::scalar::{Scalar, ScalarType}; + +#[test] +fn exact_copy_components_keep_handlers_and_do_not_materialize_pointers() { + assert!( + std::mem::size_of::() <= 20, + "copy metadata must not inflate ordinary IR instructions" + ); + for typed in [false, true] { + for protected in [false, true] { + let mut types = Types::default(); + let unit = types.intern(Type::Unit); + let long = types.scalar(ScalarType::I64); + let int = types.scalar(ScalarType::I32); + let object = types.symbol("java/lang/Object"); + let object = types.intern(Type::Class(object)); + let pointer = types.intern(Type::Pointer(int)); + let codec = types.symbol("test/Codec#value#[I"); + let mut b = Builder::new(&types, unit); + let root = b.parameter(b.current(), object); + let offset = b.constant(long, Scalar::integer(ScalarType::I64, 4).unwrap()); + let count = b.constant(long, Scalar::integer(ScalarType::I64, 16).unwrap()); + let parts = b.args([root, offset]); + let op = if typed { + Op::TypedAddressPack { + parts, + size: 8, + codec: Some(codec), + } + } else { + Op::AddressPack(parts) + }; + let source = b.emit(op, Some(pointer)).unwrap(); + let method = b.method(MethodRef { + owner: "org/rustlang/runtime/Pointer".into(), + name: "copyNonOverlapping".into(), + params: vec![pointer, pointer, long], + returns: unit, + interface: false, + }); + let args = b.args([source, source, count]); + let call = Op::Call { + method, + kind: CallKind::JvmStatic, + args, + }; + if protected { + let handler = b.create_block(); + b.invoke(call, None, handler); + b.terminate(Terminator::Return(None)); + b.switch_to(handler); + b.terminate(Terminator::Rethrow); + } else { + b.emit(call, None); + b.terminate(Terminator::Return(None)); + } + let mut body = b.finish().unwrap(); + lower_memory_copies(&mut body, &mut types); + verify(&body, &types).unwrap(); + let (index, op) = body + .instructions + .iter() + .enumerate() + .find(|(_, i)| matches!(i.op, Op::CopyStorage { .. })) + .unwrap(); + let Op::CopyStorage { + layouts, + nonoverlapping, + .. + } = op.op + else { + unreachable!() + }; + assert!(nonoverlapping); + let layouts = layouts.map(|ty| { + let Type::Layout(id) = types.get(ty).unwrap() else { + panic!() + }; + let layout = types.get_layout(id); + (layout.size, layout.codec) + }); + assert_eq!( + layouts, + if typed { + [(8, Some(codec)); 2] + } else { + [(4, None); 2] + } + ); + assert!(!super::live(&body, &types).values[source.index()]); + if protected { + assert!(body.blocks.iter().any(|b| matches!(b.terminator, + Some(Terminator::Invoke { inst, .. }) if inst.index() == index))); + } + crate::jvm::select::compile(&body, &types, &mut Default::default()).unwrap(); + } + } +} diff --git a/compiler-core/src/opt/mod.rs b/compiler-core/src/opt/mod.rs index 1146ba85..a61cc504 100644 --- a/compiler-core/src/opt/mod.rs +++ b/compiler-core/src/opt/mod.rs @@ -64,8 +64,22 @@ pub use typed_loads::lower_typed_loads; #[cfg(test)] mod typed_loads_tests; +mod address_observers; +pub use address_observers::lower_address_observers; +#[cfg(test)] +mod address_observers_tests; + +mod address_intrinsics; +pub use address_intrinsics::lower_address_intrinsics; +#[cfg(test)] +mod address_intrinsics_tests; + +mod memory_copies; +pub use memory_copies::lower_memory_copies; #[cfg(test)] mod array_view_tests; +#[cfg(test)] +mod memory_copies_tests; mod simplify; mod unreachable; From 16e140a8f9e88fb6a82a1c35e0f1b3a69b563310 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 14:29:28 +1000 Subject: [PATCH 23/61] make Java exports explicit --- Readme.md | 22 +++- src/java_exports.rs | 74 +++++++++--- src/lower2/jvm_gen/bridges.rs | 91 ++------------- tests/cargo_jvm/workflows/src/lib.rs | 4 + tests/integration/cdylib/src/lib.rs | 4 + tests/integration/inner_classes/src/lib.rs | 2 + tests/integration/java_exports/Cargo.lock | 14 +++ tests/integration/java_exports/Cargo.toml | 7 ++ tests/integration/java_exports/Main.java | 28 +++++ .../java_exports/provider/Cargo.toml | 4 + .../java_exports/provider/src/lib.rs | 24 ++++ tests/integration/java_exports/src/lib.rs | 108 ++++++++++++++++++ tests/integration/jvm_hello/src/lib.rs | 4 + tests/integration/jvm_link_names/src/lib.rs | 4 + tests/integration/jvm_macros/src/lib.rs | 11 +- tests/integration/struct_methods/Main.java | 29 +++++ tests/integration/struct_methods/src/lib.rs | 37 ++++++ .../integration/trait_implementors/src/lib.rs | 4 + tests/kotlin/async_interop/src/lib.rs | 4 + tests/kotlin/readme_interop/src/lib.rs | 4 + 20 files changed, 380 insertions(+), 99 deletions(-) create mode 100644 tests/integration/java_exports/Cargo.lock create mode 100644 tests/integration/java_exports/Cargo.toml create mode 100644 tests/integration/java_exports/Main.java create mode 100644 tests/integration/java_exports/provider/Cargo.toml create mode 100644 tests/integration/java_exports/provider/src/lib.rs create mode 100644 tests/integration/java_exports/src/lib.rs diff --git a/Readme.md b/Readme.md index 7feef542..04518dd7 100644 --- a/Readme.md +++ b/Readme.md @@ -88,6 +88,21 @@ classes and interfaces rather than opaque native handles (see without JNI glue ([test and demo](tests/integration/jvm_link_names/src/lib.rs)). +Java exposure is explicit: `pub` controls Rust visibility, while +`#[jvm_codegen::export]` on a public function, type, or module preserves its +Java-facing API. Marking a module includes its public descendants; marking a type +includes its public inherent methods. Register the tool with +`#![feature(register_tool)]` and `#![register_tool(jvm_codegen)]`. +For a library intended entirely for Java, add `#![feature(custom_inner_attributes)]` +and `#![jvm_codegen::export]` at the crate root, as in the examples below. +Mark the types Java constructs or inspects as well as the functions it calls. +Use concrete, nongeneric exports as Java entry points for generic Rust APIs. +Ordinary Rust dependencies remain internal: the compiler can flatten their +representations, merge their classes, and remove unused methods. Rebuild Rust +dependencies with the same backend; generated internal classes are not a stable +Java API. Importing existing Java classes with `jvm::class` or `jvm::interface` +does not require an export marker. + For example, one Rust API can expose an enum and accept both a JVM implementation of a Rust trait and a standard JVM lambda. Its result can then cross the Rust/Kotlin async bridge. The complete example is kept executable in @@ -96,6 +111,10 @@ cross the Rust/Kotlin async bridge. The complete example is kept executable in **Rust** ```rust +#![feature(register_tool, custom_inner_attributes)] +#![register_tool(jvm_codegen)] +#![jvm_codegen::export] + pub trait BatchObserver { fn accept(&mut self, processed: u32) -> bool; } @@ -247,8 +266,9 @@ Add `jvm = { package = "rcj", git = "https://github.com/IntegralPilot/rustc_code to `[dependencies]`. ```rust -#![feature(extern_types, register_tool)] +#![feature(extern_types, register_tool, custom_inner_attributes)] #![register_tool(jvm_codegen)] +#![jvm_codegen::export] #[jvm::class("java.time.LocalDate", rename_all = "camelCase")] impl JavaLocalDate { diff --git a/src/java_exports.rs b/src/java_exports.rs index 4d37e363..b54272bf 100644 --- a/src/java_exports.rs +++ b/src/java_exports.rs @@ -1,12 +1,66 @@ //! Discover and materialize the Rust library surface exposed to Java. use super::*; +/// Export attributes select the Java API. Rust visibility alone does not select it. +pub(crate) fn is_exported(tcx: TyCtxt<'_>, def_id: DefId) -> bool { + fn scope(tcx: TyCtxt<'_>, mut current: Option) -> bool { + while let Some(def_id) = current { + #[allow(deprecated)] + if tcx.get_all_attrs(def_id).iter().any(|attribute| { + let path = attribute.path(); + path.len() == 2 && path[0].as_str() == "jvm_codegen" && path[1].as_str() == "export" + }) { + return true; + } + current = tcx.opt_parent(def_id); + } + false + } + if scope(tcx, Some(def_id)) { + return true; + } + // An exported type's inherent methods are part of its Java receiver API. + if let Some(implementation) = tcx + .opt_associated_item(def_id) + .and_then(|i| i.impl_container(tcx)) + { + let ty = tcx + .type_of(implementation) + .instantiate_identity() + .skip_norm_wip(); + if let TyKind::Adt(adt, _) = ty.kind() { + return scope(tcx, Some(adt.did())); + } + } + false +} + +/// KotlinFutureInterop requires these named schemas. Definition paths stay constant through +/// reexports. +pub(crate) fn is_runtime_type(tcx: TyCtxt<'_>, def_id: DefId) -> bool { + tcx.crate_name(def_id.krate) == rustc_span::sym::core + && matches!( + lower1::jvm_names::class_for_def_id(tcx, def_id).as_str(), + "org/rustlang/core/task/wake/RawWaker" + | "org/rustlang/core/task/wake/RawWakerVTable" + | "org/rustlang/core/task/wake/Waker" + | "org/rustlang/core/task/wake/Context" + ) +} + pub(super) fn ensure_trait_interface<'tcx>( tcx: TyCtxt<'tcx>, trait_def_id: DefId, data_types: &mut Definitions<'tcx>, ) { let interface_name = lower1::jvm_names::class_for_def_id(tcx, trait_def_id); + if !trait_def_id.is_local() && data_types.has_upstream_type(&interface_name) { + data_types + .foreign_interfaces + .borrow_mut() + .insert(interface_name); + return; + } let methods = trait_interface_methods(tcx, trait_def_id, &interface_name, data_types); match data_types.get_mut(&interface_name) { @@ -63,9 +117,8 @@ pub(super) fn trait_interface_methods<'tcx>( continue; } - let mir_sig = tcx.instantiate_bound_regions_with_erased( - tcx.type_of(def_id).skip_binder().fn_sig(tcx), - ); + let mir_sig = tcx + .instantiate_bound_regions_with_erased(tcx.type_of(def_id).skip_binder().fn_sig(tcx)); let explicit_inputs = mir_sig.inputs(); let output = mir_sig.output(); let instance = Instance::new_raw( @@ -122,12 +175,6 @@ pub(super) fn trait_interface_methods<'tcx>( methods } -pub(super) fn crate_emits_library_artifact(tcx: TyCtxt<'_>) -> bool { - tcx.crate_types() - .iter() - .any(|crate_type| !matches!(crate_type, CrateType::Executable)) -} - pub(super) fn is_lowerable_java_public_function(tcx: TyCtxt<'_>, def_id: DefId) -> bool { if !matches!(tcx.def_kind(def_id), DefKind::Fn | DefKind::AssocFn) { return false; @@ -166,7 +213,8 @@ pub(super) fn java_public_surface_def_ids( JavaPublicSurface::Exported => effective_visibilities.is_exported(local_def_id), JavaPublicSurface::Reachable => effective_visibilities.is_reachable(local_def_id), }; - is_public_enough.then_some(local_def_id.to_def_id()) + (is_public_enough && is_exported(tcx, local_def_id.to_def_id())) + .then_some(local_def_id.to_def_id()) }) .collect(); @@ -199,15 +247,11 @@ pub(super) fn materialize_java_public_data_type<'tcx>( } } -pub(super) fn lower_public_library_exports<'tcx>( +pub(super) fn lower_java_exports<'tcx>( tcx: TyCtxt<'tcx>, oomir_module: &mut lower1::context::Module<'tcx>, lowered_instances: &Lock>>, ) { - if !crate_emits_library_artifact(tcx) { - return; - } - let function_defs = java_public_surface_def_ids(tcx, JavaPublicSurface::Exported) .into_iter() .filter(|def_id| is_lowerable_java_public_function(tcx, *def_id)) diff --git a/src/lower2/jvm_gen/bridges.rs b/src/lower2/jvm_gen/bridges.rs index 1b1fd501..b4ae25ff 100644 --- a/src/lower2/jvm_gen/bridges.rs +++ b/src/lower2/jvm_gen/bridges.rs @@ -1,60 +1,6 @@ //! Native JVM bridges emission. use super::*; -pub(in crate::lower2) fn create_relative_pointer_bridge( - cp: &mut InternedConstantPool, - class_name: &str, - method_name: &str, - signature: &oomir::Signature, - access_flags: MethodAccessFlags, - owner_is_interface: bool, -) -> jvm::Result { - debug_assert!(signature.is_static); - let relative_signature = signature.relative_pointer_abi_signature(); - let relative_name = format!("{method_name}{}", oomir::RELATIVE_POINTER_METHOD_SUFFIX); - let class_index = cp.add_class(class_name)?; - let target = if owner_is_interface { - cp.add_interface_method_ref(class_index, &relative_name, &relative_signature.to_string())? - } else { - cp.add_method_ref(class_index, &relative_name, &relative_signature.to_string())? - }; - - let mut instructions = Vec::new(); - let mut local = 0u16; - - for (_, ty) in &signature.params { - if !ty.has_jvm_value() { - continue; - } - instructions.push(get_load_instruction(ty, local)?); - let size = get_type_size(ty); - local += size; - - if matches!(ty, Type::Pointer(_)) { - instructions.push(Instruction::Lconst_0); - instructions.push(Instruction::Lconst_0); - } - } - instructions.push(Instruction::Invokestatic(target)); - instructions.push(return_instruction_for_type(&signature.ret)); - - let descriptor = signature.to_string(); - Ok(jvm::Method { - access_flags, - name_index: cp.add_utf8(method_name)?, - descriptor_index: cp.add_utf8(&descriptor)?, - attributes: vec![code_attribute_for_descriptor( - cp, - local, - instructions, - &descriptor, - true, - Some(class_name), - method_name, - )?], - }) -} - pub(super) fn create_static_instance_bridge( cp: &mut InternedConstantPool, class_name_jvm: &str, @@ -118,37 +64,18 @@ pub(super) fn create_static_instance_bridge( } instructions.push(return_instruction_for_type(bridge_signature.ret.as_ref())); - let mut parameters = Vec::new(); - for (param_name, param_ty) in &bridge_signature.params { - if !param_ty.has_jvm_value() { - continue; - } - let name_index = cp.add_utf8(param_name)?; - parameters.push(jvm::attributes::MethodParameter { - name_index, - access_flags: MethodAccessFlags::empty(), - }); - } - let method_parameters_attribute_name_index = cp.add_utf8("MethodParameters")?; - Ok(jvm::Method { access_flags: MethodAccessFlags::PUBLIC | MethodAccessFlags::STATIC, name_index, descriptor_index, - attributes: vec![ - code_attribute_for_descriptor( - cp, - next_local, - instructions, - &bridge_descriptor, - true, - None, - method_name, - )?, - Attribute::MethodParameters { - name_index: method_parameters_attribute_name_index, - parameters, - }, - ], + attributes: vec![code_attribute_for_descriptor( + cp, + next_local, + instructions, + &bridge_descriptor, + true, + None, + method_name, + )?], }) } diff --git a/tests/cargo_jvm/workflows/src/lib.rs b/tests/cargo_jvm/workflows/src/lib.rs index 0bee4e5a..4d11b6aa 100644 --- a/tests/cargo_jvm/workflows/src/lib.rs +++ b/tests/cargo_jvm/workflows/src/lib.rs @@ -1,3 +1,7 @@ +#![feature(register_tool)] +#![register_tool(jvm_codegen)] + +#[jvm_codegen::export] pub fn triple(value: u32) -> u32 { value * 3 } diff --git a/tests/integration/cdylib/src/lib.rs b/tests/integration/cdylib/src/lib.rs index 700b9bca..d01101ff 100644 --- a/tests/integration/cdylib/src/lib.rs +++ b/tests/integration/cdylib/src/lib.rs @@ -1,3 +1,7 @@ +#![feature(custom_inner_attributes)] +#![feature(register_tool)] +#![register_tool(jvm_codegen)] +#![jvm_codegen::export] pub fn multiply(left: i32, right: i32) -> i32 { left * right } diff --git a/tests/integration/inner_classes/src/lib.rs b/tests/integration/inner_classes/src/lib.rs index c8e716f9..f90f3394 100644 --- a/tests/integration/inner_classes/src/lib.rs +++ b/tests/integration/inner_classes/src/lib.rs @@ -1,3 +1,5 @@ +#![feature(custom_inner_attributes)] +#![jvm_codegen::export] #![feature(register_tool)] #![register_tool(jvm_codegen)] diff --git a/tests/integration/java_exports/Cargo.lock b/tests/integration/java_exports/Cargo.lock new file mode 100644 index 00000000..fdb189e7 --- /dev/null +++ b/tests/integration/java_exports/Cargo.lock @@ -0,0 +1,14 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 4 + +[[package]] +name = "export-provider" +version = "0.1.0" + +[[package]] +name = "java-exports" +version = "0.1.0" +dependencies = [ + "export-provider", +] diff --git a/tests/integration/java_exports/Cargo.toml b/tests/integration/java_exports/Cargo.toml new file mode 100644 index 00000000..aa86fd3e --- /dev/null +++ b/tests/integration/java_exports/Cargo.toml @@ -0,0 +1,7 @@ +[package] +name = "java-exports" +version = "0.1.0" +edition = "2024" + +[dependencies] +export-provider = { path = "provider" } diff --git a/tests/integration/java_exports/Main.java b/tests/integration/java_exports/Main.java new file mode 100644 index 00000000..d7f67c66 --- /dev/null +++ b/tests/integration/java_exports/Main.java @@ -0,0 +1,28 @@ +public class Main { + public static void main(String[] args) { + if (java_exports.java_exports.answer() != 42) throw new AssertionError(); + java_exports.Counter counter = new java_exports.Counter(40); + if (counter.add(2) != 42 || counter.value != 42) throw new AssertionError(); + if (!(java_exports.api.calls.choose(true) instanceof java_exports.api.Choice.First)) { + throw new AssertionError(); + } + export_provider.Shared shared = java_exports.java_exports.upstream(new export_provider.Shared(38)); + if (shared.value != 40) throw new AssertionError(); + if (shared.add(2) != 42 || shared.value != 42) throw new AssertionError(); + org.rustlang.runtime.Pointer internal = java_exports.java_exports.internal_drop_value(); + if (internal.getObject() instanceof org.rustlang.runtime.RustDrop) { + throw new AssertionError("unexported Rust type should not acquire a Java drop callback"); + } + java_exports.java_exports.destroy_internal(internal); + if (java_exports.java_exports.drop_count() != 2) { + throw new AssertionError("typed and trait-object destruction must still run"); + } + new java_exports.ExportedDrop(7.0).rustDrop(); + if (java_exports.java_exports.drop_count() != 3) { + throw new AssertionError("exported Java destruction callback must remain available"); + } + for (var method : java_exports.java_exports.class.getDeclaredMethods()) { + if (method.getName().equals("internal_function")) throw new AssertionError(); + } + } +} diff --git a/tests/integration/java_exports/provider/Cargo.toml b/tests/integration/java_exports/provider/Cargo.toml new file mode 100644 index 00000000..30bdc483 --- /dev/null +++ b/tests/integration/java_exports/provider/Cargo.toml @@ -0,0 +1,4 @@ +[package] +name = "export-provider" +version = "0.1.0" +edition = "2024" diff --git a/tests/integration/java_exports/provider/src/lib.rs b/tests/integration/java_exports/provider/src/lib.rs new file mode 100644 index 00000000..d314d33e --- /dev/null +++ b/tests/integration/java_exports/provider/src/lib.rs @@ -0,0 +1,24 @@ +#![feature(register_tool)] +#![register_tool(jvm_codegen)] + +#[jvm_codegen::export] +pub struct Shared { + pub value: i32, +} + +impl Shared { + #[inline(never)] + pub fn add(&mut self, amount: i32) -> i32 { + self.value += amount; + self.value + } +} + +pub struct Internal { + pub value: i32, +} + +#[inline(never)] +pub fn internal_value(value: Internal) -> i32 { + value.value +} diff --git a/tests/integration/java_exports/src/lib.rs b/tests/integration/java_exports/src/lib.rs new file mode 100644 index 00000000..86ceaaf3 --- /dev/null +++ b/tests/integration/java_exports/src/lib.rs @@ -0,0 +1,108 @@ +#![feature(register_tool)] +#![register_tool(jvm_codegen)] + +// Rust-public implementation details do not acquire a Java method surface. +pub struct Internal { + pub value: i32, +} + +impl Internal { + #[inline(never)] + pub fn read(&self) -> i32 { + self.value + } +} + +pub fn internal_function(value: i32) -> i32 { + Internal { value }.read() +} + +#[jvm_codegen::export] +pub fn answer() -> i32 { + internal_function(42) +} + +#[jvm_codegen::export] +pub struct Counter { + pub value: i32, +} + +impl Counter { + pub fn add(&mut self, amount: i32) -> i32 { + self.value += amount; + self.value + } +} + +#[jvm_codegen::export] +pub mod api { + pub enum Choice { + First, + Second, + } + + pub mod calls { + pub fn choose(value: bool) -> super::Choice { + if value { + super::Choice::First + } else { + super::Choice::Second + } + } + } +} + +// Rust calls and Java receiver methods must share the exported upstream schema. +#[jvm_codegen::export] +pub fn upstream(mut value: export_provider::Shared) -> export_provider::Shared { + value.add(export_provider::internal_value(export_provider::Internal { + value: 2, + })); + value +} + +static DROP_COUNT: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0); + +// Rust visibility alone does not require an erased Java destruction callback. +pub struct InternalDrop { + value: f64, +} + +impl Drop for InternalDrop { + fn drop(&mut self) { + assert_eq!(self.value, 42.0); + DROP_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + } +} + +#[jvm_codegen::export] +pub fn internal_drop_value() -> *mut InternalDrop { + Box::into_raw(Box::new(InternalDrop { value: 42.0 })) +} + +#[jvm_codegen::export] +pub unsafe fn destroy_internal(value: *mut InternalDrop) { + unsafe { + drop(Box::from_raw(value)); + } + // Trait erasure still needs a concrete destruction adapter. + let erased: Box = Box::new(InternalDrop { value: 42.0 }); + drop(erased); +} + +#[jvm_codegen::export] +pub struct ExportedDrop { + pub value: f64, +} + +impl Drop for ExportedDrop { + fn drop(&mut self) { + assert_eq!(self.value, 7.0); + DROP_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + } +} + +#[jvm_codegen::export] +pub fn drop_count() -> usize { + DROP_COUNT.load(std::sync::atomic::Ordering::SeqCst) +} diff --git a/tests/integration/jvm_hello/src/lib.rs b/tests/integration/jvm_hello/src/lib.rs index d2d9c6b2..0693d8a2 100644 --- a/tests/integration/jvm_hello/src/lib.rs +++ b/tests/integration/jvm_hello/src/lib.rs @@ -1,3 +1,7 @@ +#![feature(custom_inner_attributes)] +#![feature(register_tool)] +#![register_tool(jvm_codegen)] +#![jvm_codegen::export] pub struct Ciallo { pub count: i32, pub desc: &'static str, diff --git a/tests/integration/jvm_link_names/src/lib.rs b/tests/integration/jvm_link_names/src/lib.rs index 58ed5f2c..ea44d067 100644 --- a/tests/integration/jvm_link_names/src/lib.rs +++ b/tests/integration/jvm_link_names/src/lib.rs @@ -1,3 +1,7 @@ +#![feature(custom_inner_attributes)] +#![feature(register_tool)] +#![register_tool(jvm_codegen)] +#![jvm_codegen::export] #![feature(extern_types)] unsafe extern "C" { diff --git a/tests/integration/jvm_macros/src/lib.rs b/tests/integration/jvm_macros/src/lib.rs index 4987a74b..c1306042 100644 --- a/tests/integration/jvm_macros/src/lib.rs +++ b/tests/integration/jvm_macros/src/lib.rs @@ -1,3 +1,5 @@ +#![feature(custom_inner_attributes)] +#![jvm_codegen::export] #![feature(extern_types, register_tool)] #![register_tool(jvm_codegen)] @@ -67,6 +69,13 @@ pub enum MacroRoot { Other(i32), } +#[inline(never)] +fn update_through_reborrow(state: &mut JavaState, value: i32) { + let borrowed = &mut *state; + borrowed.set_value(value); + assert_eq!(state.get_value(), value); +} + pub fn exercise() -> i64 { interfaces::exercise(); unsafe { @@ -77,7 +86,7 @@ pub fn exercise() -> i64 { assert_eq!((&*first).get_value(), 7); assert_eq!((&*first).get_wide(), 4_000_000_000); - (&mut *first).set_value(13); + update_through_reborrow(&mut *first, 13); assert_eq!((&*first).get_value(), 13); let second = JavaState::new(11, 9); diff --git a/tests/integration/struct_methods/Main.java b/tests/integration/struct_methods/Main.java index 7dc46c07..79166bec 100644 --- a/tests/integration/struct_methods/Main.java +++ b/tests/integration/struct_methods/Main.java @@ -11,10 +11,21 @@ public static void main(String[] args) throws Exception { assertRustFinalizeIsNotJavaFinalizer(); assertFieldlessNoArgsConstructorRemains(); assertConstantStructUsesDeclarationOrder(); + assertPrivateCarrierHasNoReceiverBridges(); + assertOptionalViewBridges(); System.out.println("Struct method mapping test passed!"); } + private static void assertOptionalViewBridges() { + struct_methods.OptionalViewBridge bridge = new struct_methods.OptionalViewBridge(); + org.rustlang.runtime.Utf8View empty = org.rustlang.runtime.Utf8View.fromJavaString(""); + if (bridge.has_value(null) || bridge.has_value(bridge.round_trip(null)) + || !bridge.has_value(empty) || !bridge.has_value(bridge.round_trip(empty))) { + throw new AssertionError("optional view bridges must distinguish None from empty borrows"); + } + } + private static void assertNoFieldedNoArgsConstructor() { try { struct_methods.NamedCounter.class.getConstructor(); @@ -119,4 +130,22 @@ private static void assertConstantStructUsesDeclarationOrder() { throw new AssertionError("constant structs should use declaration-order constructor arguments"); } } + + private static void assertPrivateCarrierHasNoReceiverBridges() throws Exception { + if (struct_methods.struct_methods.private_counter() != 12) { + throw new AssertionError("private Rust calls and destruction must still run"); + } + Class carrier; + try { + carrier = Class.forName("struct_methods.PrivateCounter"); + } catch (ClassNotFoundException eliminatedCarrier) { + return; + } + for (Method method : carrier.getDeclaredMethods()) { + if (method.getName().equals("advance") || method.getName().equals("drop") + || method.getName().equals("rustDrop")) { + throw new AssertionError("private carrier should not have an unused receiver bridge: " + method); + } + } + } } diff --git a/tests/integration/struct_methods/src/lib.rs b/tests/integration/struct_methods/src/lib.rs index 356e7dad..a20fae0c 100644 --- a/tests/integration/struct_methods/src/lib.rs +++ b/tests/integration/struct_methods/src/lib.rs @@ -1,3 +1,7 @@ +#![feature(custom_inner_attributes)] +#![feature(register_tool)] +#![register_tool(jvm_codegen)] +#![jvm_codegen::export] pub struct NamedCounter { pub name: &'static str, pub count: u32, @@ -89,3 +93,36 @@ pub fn finish_number() -> u32 { finish_trait(&mut number); number } + +struct PrivateCounter(u32); + +impl PrivateCounter { + #[inline(never)] + fn advance(&mut self) { + self.0 += 5; + } +} + +impl Drop for PrivateCounter { + fn drop(&mut self) { + assert_eq!(self.0, 12); + } +} + +pub fn private_counter() -> u32 { + let mut counter = std::hint::black_box(PrivateCounter(7)); + counter.advance(); + counter.0 +} + +pub struct OptionalViewBridge; + +impl OptionalViewBridge { + pub fn has_value(&self, value: Option<&str>) -> bool { + value.is_some() + } + + pub fn round_trip(&self, value: Option<&'static str>) -> Option<&'static str> { + std::hint::black_box(value) + } +} diff --git a/tests/integration/trait_implementors/src/lib.rs b/tests/integration/trait_implementors/src/lib.rs index 4857bcce..3fdaa147 100644 --- a/tests/integration/trait_implementors/src/lib.rs +++ b/tests/integration/trait_implementors/src/lib.rs @@ -1,3 +1,7 @@ +#![feature(custom_inner_attributes)] +#![feature(register_tool)] +#![register_tool(jvm_codegen)] +#![jvm_codegen::export] pub mod bound_lifetimes; /// A Java class can implement this generated interface and Rust will dispatch diff --git a/tests/kotlin/async_interop/src/lib.rs b/tests/kotlin/async_interop/src/lib.rs index d2ee0d92..76c09d5d 100644 --- a/tests/kotlin/async_interop/src/lib.rs +++ b/tests/kotlin/async_interop/src/lib.rs @@ -1,3 +1,7 @@ +#![feature(custom_inner_attributes)] +#![feature(register_tool)] +#![register_tool(jvm_codegen)] +#![jvm_codegen::export] use core::future::Future; use core::pin::Pin; use core::task::{Context, Poll}; diff --git a/tests/kotlin/readme_interop/src/lib.rs b/tests/kotlin/readme_interop/src/lib.rs index 586508bb..1d1f8ce2 100644 --- a/tests/kotlin/readme_interop/src/lib.rs +++ b/tests/kotlin/readme_interop/src/lib.rs @@ -1,3 +1,7 @@ +#![feature(custom_inner_attributes)] +#![feature(register_tool)] +#![register_tool(jvm_codegen)] +#![jvm_codegen::export] pub trait BatchObserver { fn accept(&mut self, processed: u32) -> bool; } From 7d7d014a35f0505fe7a5c2135bb47fbf799068e3 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 15:54:49 +1000 Subject: [PATCH 24/61] model typed storage in OOMIR --- src/oomir.rs | 39 +++++- src/oomir/construct/types.rs | 77 ++++++++++-- src/oomir/emission.rs | 79 ++++++++++++ src/oomir/types.rs | 236 ++++++++++++++++++++++------------- src/oomir/visit.rs | 13 +- 5 files changed, 337 insertions(+), 107 deletions(-) diff --git a/src/oomir.rs b/src/oomir.rs index f298cf7d..3d0711c4 100644 --- a/src/oomir.rs +++ b/src/oomir.rs @@ -3,6 +3,7 @@ mod forward; pub use forward::{MethodForwarder, ReceiverPointer}; mod body; pub(crate) mod construct; +pub mod fields; pub(crate) mod outline; pub use body::SsaBody; pub type SsaFunction = Function>; @@ -23,12 +24,13 @@ pub mod scalar; mod visit; pub use jvm_compiler_core::jvm::abi::{ - POINTER_CLASS, RELATIVE_POINTER_METHOD_SUFFIX, SLICE_VIEW_CLASS, UTF8_VIEW_CLASS, - relative_pointer_byte_offset_field, relative_pointer_element_offset_field, + POINTER_CLASS, SLICE_VIEW_CLASS, TAGGED_LONG_CLASS, UTF8_VIEW_CLASS, }; pub const JAVA_STRING_CLASS: &str = "java/lang/String"; pub const CALLER_LOCATION_PARAM_NAME: &str = "__caller_location"; +pub const ENUM_TAG_METHOD: &str = "$rust$tag"; pub use jvm_compiler_core::debug::SourceLocation; +pub use jvm_compiler_core::ir::HeapOp; /// A source-level Rust variable that can be represented by a JVM local slot. #[derive(Debug, Clone, PartialEq, Eq, Hash)] @@ -53,9 +55,6 @@ pub struct Module> { pub shared_data_types: Option>>, /// Canonical emission shards use the same immutable representation tables. pub shared_context: Option>>, - /// Static methods removed into the canonical data-type contribution table - /// that still use the component-carrying internal pointer ABI. - pub relative_static_methods: Arc>, /// JVM interfaces referenced by this shard but defined in another crate. pub external_interfaces: HashSet, pub statics: HashMap, @@ -73,7 +72,6 @@ impl Module { suppressed_data_types: self.suppressed_data_types, shared_data_types: self.shared_data_types, shared_context: self.shared_context, - relative_static_methods: self.relative_static_methods, external_interfaces: self.external_interfaces, statics: self.statics, } @@ -94,6 +92,17 @@ impl Module { } impl Module { + pub fn component_method(&self, owner: &str, name: &str) -> bool { + component_method(owner, name) + || (name == "call" + && self.data_type(owner).is_some_and(|data| { + let (DataType::Class { interfaces, .. } + | DataType::Interface { interfaces, .. }) = data; + interfaces + .iter() + .any(|p| p.starts_with("org/rustlang/runtime/FnPtr_")) + })) + } pub fn data_type(&self, name: &str) -> Option<&DataType> { self.data_types.get(name).or_else(|| { self.shared_data_types @@ -167,6 +176,7 @@ pub struct Static { pub allocation_alignment: usize, pub allocation_codec_class_name: Option, pub is_thread_local: bool, + pub is_private: bool, } impl Static { @@ -215,6 +225,8 @@ pub enum AdtHelperKind { enum_class: String, variants: Vec, values: Vec, + /// Private variants answer directly without naming every sibling class. + dispatch: Option, }, EnumIsVariant { enum_class: String, @@ -233,9 +245,21 @@ pub enum AdtHelperKind { }, } +/// Static method namespaces require no Rust storage, constructors, or copy helpers. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum ClassKind { + Value, + /// A named public Java value keeps its declared field contract. + JavaValue, + /// DST fields describe decoded views. They do not supply authoritative inline storage. + MemoryView, + Static, +} + #[derive(Debug, Clone, PartialEq, Eq)] pub enum DataType { Class { + kind: ClassKind, is_abstract: bool, super_class: Option, fields: Vec<(String, Type)>, @@ -254,12 +278,14 @@ impl Hash for DataType { std::mem::discriminant(self).hash(state); match self { Self::Class { + kind, is_abstract, super_class, fields, methods, interfaces, } => { + kind.hash(state); is_abstract.hash(state); super_class.hash(state); fields.hash(state); @@ -288,6 +314,7 @@ impl DataType { pub fn clean_duplicates(&mut self) { match self { DataType::Class { + kind: _, is_abstract: _, super_class: _, fields, diff --git a/src/oomir/construct/types.rs b/src/oomir/construct/types.rs index 425899ee..647a77ad 100644 --- a/src/oomir/construct/types.rs +++ b/src/oomir/construct/types.rs @@ -105,6 +105,10 @@ impl Vocabulary { use ir::Type as I; use oomir::Type as T; let repr = match ty { + T::TaggedI64 => { + self.add(&T::I64); + I::TaggedI64 + } T::Void | T::Unit => I::Unit, T::Boolean => I::Scalar(ScalarType::Bool), T::Char | T::U16 => I::Scalar(ScalarType::U16), @@ -120,9 +124,25 @@ impl Vocabulary { T::F64 => I::Scalar(ScalarType::F64), T::Class(name) => I::Class(self.types.symbol(name)), T::Interface(name) => I::Interface(self.types.symbol(name)), - T::Pointer(inner) => I::Pointer(self.add(inner)), + T::Pointer(inner) => { + let mut value = self.add(inner); + if let Some(layout) = &inner.layout { + let codec = layout.codec.as_ref().map(|name| self.types.symbol(name)); + let layout_type = ir::AddressLayout { + value, + size: layout.size, + codec, + }; + let id = self.types.layout(layout_type); + if id.index() == COMMON_DESCRIPTORS.len() + self.descriptors.len() { + self.descriptors.push(String::new()); + } + value = id; + } + I::Pointer(value) + } T::Slice(inner) => I::Slice(self.add(inner)), - T::Array(inner) | T::MutableReference(inner) => { + T::Array(inner) => { let inner = if inner.has_jvm_value() { self.add(inner) } else { @@ -130,11 +150,6 @@ impl Vocabulary { }; I::Array(inner) } - T::Reference(inner) => { - let id = self.add(inner); - self.ids.insert(ty.clone(), id); - return id; - } T::Str => I::Str, }; let id = self.types.intern(repr); @@ -227,6 +242,38 @@ impl Vocabulary { self.add(&oomir::Type::Class(class_name.clone())); } match instruction { + Heap { .. } => { + self.add(&oomir::Type::pointer(oomir::Type::U8)); + } + TaggedPack { .. } => { + self.add(&oomir::Type::TaggedI64); + self.add(&oomir::Type::I64); + } + TaggedPart { .. } => { + self.add(&oomir::Type::I64); + } + AddressRetype { layout, .. } | ViewAddress { layout, .. } => { + self.add(&layout.pointer_type); + if let oomir::Operand::Constant(oomir::Constant::String(codec)) = &layout.codec { + self.types.symbol(codec); + } + } + AddressOffset { ty, .. } => { + self.add(ty); + } + MemoryProject { projection, .. } => { + self.add(&oomir::Type::Class(projection.owner.clone())); + self.add(&oomir::Type::pointer(oomir::Type::Class( + projection.owner.clone(), + ))); + self.add(&projection.pointee); + self.add(&oomir::Type::pointer(projection.pointee.clone())); + self.add(&oomir::Type::java_string()); + } + MemoryLoad { pointee, .. } | MemoryStore { pointee, .. } => { + self.add(pointee); + self.add(&oomir::Type::pointer(pointee.clone())); + } InvokeStatic { method_ty, .. } | InvokeRustStatic { method_ty, .. } | InvokeVirtual { method_ty, .. } @@ -329,10 +376,19 @@ pub(crate) fn source_type(types: &ir::Types, id: ir::TypeId) -> oomir::Type { }, I::Class(id) => T::Class(types.symbol_name(id).unwrap().into()), I::Interface(id) => T::Interface(types.symbol_name(id).unwrap().into()), - I::Pointer(inner) => T::Pointer(Box::new(source_type(types, inner))), + I::Pointer(inner) => match types.get(inner).unwrap() { + I::Layout(id) => { + let ir::AddressLayout { value, size, codec } = types.get_layout(id); + T::pointer(source_type(types, value)) + .with_address_layout(size, codec.map(|s| types.symbol_name(s).unwrap().into())) + } + _ => T::pointer(source_type(types, inner)), + }, + I::Layout(_) => panic!("a pointee layout is not a JVM value"), I::Array(inner) => T::Array(Box::new(source_type(types, inner))), I::Slice(inner) => T::Slice(Box::new(source_type(types, inner))), I::Str => T::Str, + I::TaggedI64 => T::TaggedI64, } } @@ -364,11 +420,10 @@ mod tests { Class("java/lang/Throwable".into()), Class("A".into()), Interface("I".into()), - Pointer(Box::new(I32)), - Pointer(Box::new(I64)), + oomir::Type::pointer(I32), + oomir::Type::pointer(I64), Array(Box::new(F16)), Array(Box::new(I16)), - Reference(Box::new(U8)), ]; let mut vocabulary = Vocabulary::default(); for ty in &types { diff --git a/src/oomir/emission.rs b/src/oomir/emission.rs index 5568ee8b..084962e2 100644 --- a/src/oomir/emission.rs +++ b/src/oomir/emission.rs @@ -24,8 +24,37 @@ pub struct BasicBlock { pub use jvm_compiler_core::scalar::BinaryOp; +/// Field locations keep the Rust layout independently of the generated carrier. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct MemoryProjection { + pub owner: String, + pub field: String, + pub pointee: Type, + pub offset: u64, + pub size: u64, + pub codec: Option, +} + +/// Address layouts stay explicit until instruction selection. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct AddressLayout { + pub pointer_type: Type, + pub size: Operand, + pub codec: Operand, +} + #[derive(Debug, Clone, PartialEq, Eq, Hash)] pub enum Instruction { + TaggedPack { + dest: String, + value: Operand, + tag: Operand, + }, + TaggedPart { + dest: String, + value: Operand, + index: u8, + }, SourceLocation(SourceLocation), // metadata. does not emit JVM bytecode. LocalVariableScope(Vec), // same UnwindStart { @@ -33,6 +62,7 @@ pub enum Instruction { }, // metadata for a protected JVM region. UnwindEnd, // ends the current protected region. Rethrow, // resumes the current Rust unwind. + Unreachable, // reaching this terminator is undefined behavior. Binary { op: BinaryOp, dest: String, @@ -88,6 +118,55 @@ pub enum Instruction { dest: String, src: Operand, // Source operand (could be Variable or Constant, though in this context, it's likely Variable) }, + /// Owned reads make independent copies. Live reads can require MemoryCommit after field + /// mutation. + MemoryLoad { + dest: String, + pointer: Operand, + pointee: Type, + owned: bool, + }, + MemoryStore { + pointer: Operand, + pointee: Type, + value: Operand, + }, + ValueCopy { + dest: String, + source: Operand, + }, + MemoryCommit { + pointer: Operand, + }, + Heap { + operation: HeapOp, + args: Vec, + dest: Option, + }, + MemoryProject { + dest: String, + base: Operand, + projection: Box, + }, + AddressRetype { + dest: Option, + source: Operand, + layout: Box, + }, + ViewAddress { + dest: Option, + source: Operand, + layout: Box, + }, + AddressOffset { + dest: Option, + source: Operand, + count: Operand, + ty: Type, + bytes: bool, + wrapping: bool, + subtract: bool, + }, ThrowNewWithMessage { exception_class: String, // e.g., "java/lang/RuntimeException" message: String, // The message from the panic/assert diff --git a/src/oomir/types.rs b/src/oomir/types.rs index c05a4a50..ffbb1742 100644 --- a/src/oomir/types.rs +++ b/src/oomir/types.rs @@ -2,7 +2,6 @@ use super::*; #[derive(Debug, Clone, PartialEq, Eq, Hash)] -#[allow(dead_code)] /* Reference variant currently unused */ pub enum Type { Void, /// Rust's inhabited, zero-sized unit value. It has no JVM stack value or local slot. @@ -20,14 +19,42 @@ pub enum Type { F16, F32, F64, - Pointer(Box), // A sized Rust reference or raw pointer. - MutableReference(Box), - Reference(Box), // Representing references, not currently constructed but might be useful in future for more complex things. - Array(Box), // Representing arrays - Slice(Box), // A view over an array with an offset and length. - Str, // A borrowed UTF-8 byte view. - Class(String), // For structs, enums, and potentially Objects - Interface(String), // dyn TraitName + Pointer(Pointee), // A sized Rust reference or raw pointer. + Array(Box), // Representing arrays + Slice(Box), // A view over an array with an offset and length. + TaggedI64, + Str, // A borrowed UTF-8 byte view. + Class(String), // For structs, enums, and potentially Objects + Interface(String), // dyn TraitName +} + +/// Source-language layout attached to an address, independent of its JVM carrier. +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct AddressSchema { + pub size: u32, + pub codec: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Hash)] +pub struct Pointee { + pub value: Box, + pub layout: Option>, +} +impl AsRef for Pointee { + fn as_ref(&self) -> &Type { + &self.value + } +} +impl std::ops::Deref for Pointee { + type Target = Type; + fn deref(&self) -> &Type { + &self.value + } +} +impl std::ops::DerefMut for Pointee { + fn deref_mut(&mut self) -> &mut Type { + &mut self.value + } } pub fn is_non_null_class_name(class_name: &str) -> bool { @@ -38,52 +65,93 @@ pub fn is_non_null_class_name(class_name: &str) -> bool { } impl Type { - /// A readable, descriptor-stable token for generated JVM ABI helper names. - pub fn jvm_abi_name_token(&self) -> String { - fn identifier(raw: &str) -> String { - let mut result = String::with_capacity(raw.len()); - let mut separator = false; - for ch in raw.chars() { - if ch.is_ascii_alphanumeric() || matches!(ch, '_' | '$') { - result.push(ch); - separator = false; - } else if !separator && !result.is_empty() { - result.push('_'); - separator = true; - } - } - while result.ends_with('_') { - result.pop(); - } - if result.is_empty() { - "Object".to_string() - } else { - result - } + pub fn pointer(value: Type) -> Self { + Self::Pointer(Pointee { + value: Box::new(value), + layout: None, + }) + } + pub fn with_address_layout(self, size: u32, codec: Option) -> Self { + let Self::Pointer(mut pointee) = self else { + return self; + }; + pointee.layout = Some(std::sync::Arc::new(AddressSchema { size, codec })); + Self::Pointer(pointee) + } + + pub fn materialize_address( + &self, + cp: &mut jvm_compiler_core::classfile::constant_pool::InternedConstantPool, + code: &mut Vec, + ) -> jvm_compiler_core::classfile::Result<()> { + if let Self::Pointer(pointee) = self + && let Some(layout) = &pointee.layout + { + jvm_compiler_core::jvm::abi::materialize_typed_address( + cp, + code, + layout.size, + layout.codec.as_deref(), + ) + } else { + jvm_compiler_core::jvm::abi::materialize_address(cp, code, self.address_plan()) } + } + pub fn scalar_address_size(&self) -> Option { + let Self::Pointer(inner) = self else { + return None; + }; + Some(match inner.as_ref() { + Self::Boolean | Self::I8 | Self::U8 => 1, + Self::I16 | Self::U16 | Self::F16 => 2, + Self::I32 | Self::U32 | Self::F32 => 4, + Self::I64 | Self::U64 | Self::F64 => 8, + _ => return None, + }) + } + pub fn address_plan(&self) -> u32 { + use jvm_compiler_core::jvm::abi; match self { - Type::Void | Type::Unit => "void".to_string(), - Type::Boolean => "boolean".to_string(), - Type::I8 | Type::U8 => "byte".to_string(), - Type::I16 => "short".to_string(), - Type::Char | Type::U16 => "char".to_string(), - Type::I32 | Type::U32 => "int".to_string(), - Type::I64 | Type::U64 => "long".to_string(), - Type::F16 => "binary16".to_string(), - Type::F32 => "float".to_string(), - Type::F64 => "double".to_string(), - Type::Pointer(_) => "Pointer".to_string(), - Type::MutableReference(inner) | Type::Array(inner) => { - format!("Array_{}", inner.jvm_abi_name_token()) - } - Type::Reference(inner) => inner.jvm_abi_name_token(), - Type::Slice(_) => "SliceView".to_string(), - Type::Str => "Utf8View".to_string(), - Type::Class(name) | Type::Interface(name) => identifier(name), + Self::Pointer(inner) => match inner.as_ref() { + Self::Slice(_) | Self::Str => abi::STORED_VIEW, + Self::Pointer(_) => abi::STORED_ADDRESS, + _ => self.scalar_address_size().unwrap_or(0), + }, + _ => 0, } } + pub fn component_shape(&self) -> Option { + use jvm_compiler_core::ir::ComponentShape; + if matches!(self, Self::TaggedI64) { + Some(ComponentShape::TaggedI64) + } else if matches!(self, Self::Slice(_) | Self::Str) { + Some(ComponentShape::View) + } else if self.scalar_address_size().is_some() { + Some(ComponentShape::Address) + } else if matches!(self, Self::Pointer(_)) { + Some(ComponentShape::StorageAddress) + } else { + None + } + } + + pub fn components(&self) -> Option + use<>> { + use jvm_compiler_core::ir::ComponentShape; + self.component_shape().map(|shape| { + let object = Type::Class("java/lang/Object".into()); + let (parts, count) = match shape { + ComponentShape::TaggedI64 => ([Self::I64, Self::I64, Self::Unit], 2), + ComponentShape::View => ([object, Self::I32, Self::U64], 3), + ComponentShape::Address | ComponentShape::StorageAddress => { + ([object, Self::I64, Self::Unit], 2) + } + }; + parts.into_iter().take(count) + }) + } + /// The JVM's own immutable string class, used only for JVM ABI values. pub fn java_string() -> Self { Self::Class(JAVA_STRING_CLASS.to_string()) @@ -115,8 +183,7 @@ impl Type { let mut arrays = 0; let (primitive, name) = loop { match ty { - Type::Reference(inner) => ty = inner, - Type::Array(inner) | Type::MutableReference(inner) if inner.has_jvm_value() => { + Type::Array(inner) if inner.has_jvm_value() => { arrays += 1; ty = inner; } @@ -124,8 +191,8 @@ impl Type { arrays += 1; break ('L', "java/lang/Object"); } - Type::MutableReference(_) => break ('L', "java/lang/Object"), Type::Pointer(_) => break ('L', POINTER_CLASS), + Type::TaggedI64 => break ('L', TAGGED_LONG_CLASS), Type::Str => break ('L', UTF8_VIEW_CLASS), Type::Slice(_) => break ('L', SLICE_VIEW_CLASS), Type::Class(name) | Type::Interface(name) => break ('L', name.as_str()), @@ -180,13 +247,13 @@ impl Type { /// Returns None for primitive types. pub fn to_jvm_internal_name(&self) -> Option { match self { + Type::TaggedI64 => Some(TAGGED_LONG_CLASS.to_string()), Type::Str => Some(UTF8_VIEW_CLASS.to_string()), Type::Class(name) | Type::Interface(name) => Some(name.replace('.', "/")), Type::Pointer(_) => Some(POINTER_CLASS.to_string()), - Type::Reference(inner) => inner.to_jvm_internal_name(), // delegate to inner type // For array-valued types, the descriptor is the component class name // expected by `anewarray`. Mutable references use one-element arrays. - Type::Array(_) | Type::MutableReference(_) => Some(self.to_jvm_descriptor()), + Type::Array(_) => Some(self.to_jvm_descriptor()), Type::Slice(_) => Some(SLICE_VIEW_CLASS.to_string()), // Primitives don't have an internal name for anewarray. _ => None, @@ -223,13 +290,12 @@ impl Type { Type::F32 => Some(JVMInstruction::Fastore), Type::F64 => Some(JVMInstruction::Dastore), // Reference types: - Type::Str + Type::TaggedI64 + | Type::Str | Type::Class(_) | Type::Interface(_) | Type::Array(_) - | Type::Slice(_) - | Type::Reference(_) - | Type::MutableReference(_) => Some(JVMInstruction::Aastore), + | Type::Slice(_) => Some(JVMInstruction::Aastore), Type::Pointer(_) => Some(JVMInstruction::Aastore), Type::Void => None, Type::Unit => None, @@ -241,15 +307,16 @@ impl Type { match constant { Constant::Unit => Type::Unit, Constant::StaticRef { ty, .. } => ty.clone(), - Constant::FunctionPointer { interface_name, .. } => { + Constant::FunctionPointer { interface_name, .. } + | Constant::FunctionHandle { interface_name, .. } => { Type::Interface(interface_name.clone()) } Constant::FactoryCall { ty, .. } => ty.clone(), Constant::StaticCall { ty, .. } => ty.clone(), - Constant::PointerAddress { pointee, .. } => Type::Pointer(pointee.clone()), - Constant::RepeatedBytePointer { pointee, .. } => Type::Pointer(pointee.clone()), - Constant::ByteArrayPointer { pointee, .. } => Type::Pointer(pointee.clone()), - Constant::InternedPointer { pointee, .. } => Type::Pointer(pointee.clone()), + Constant::PointerAddress { pointee, .. } => Type::pointer(*pointee.clone()), + Constant::RepeatedBytePointer { pointee, .. } => Type::pointer(*pointee.clone()), + Constant::ByteArrayPointer { pointee, .. } => Type::pointer(*pointee.clone()), + Constant::InternedPointer { pointee, .. } => Type::pointer(*pointee.clone()), Constant::Null(ty) => ty.clone(), Constant::I8(_) => Type::I8, Constant::U8(_) => Type::U8, @@ -268,11 +335,11 @@ impl Type { Constant::Boolean(_) => Type::Boolean, Constant::Char(_) => Type::Char, Constant::Str(_) => Type::Str, - Constant::String(_) => Type::java_string(), + Constant::String(_) | Constant::LiteralString(_) => Type::java_string(), Constant::Instance { class_name, params, .. } if class_name == POINTER_CLASS => { - Type::Pointer(Box::new(if params.len() == 3 { + Type::pointer(if params.len() == 3 { params .first() .map(Type::from_constant) @@ -281,7 +348,10 @@ impl Type { // Address-only constructors carry no JVM pointee value from // which to infer a more specific OOMIR type. Type::Unit - })) + }) + } + Constant::Instance { class_name, .. } if class_name == TAGGED_LONG_CLASS => { + Type::TaggedI64 } Constant::Instance { class_name, .. } => Type::Class(class_name.to_string()), } @@ -311,11 +381,10 @@ impl Type { pub fn is_jvm_reference_type(&self) -> bool { matches!( self, - Type::Reference(_) - | Type::Pointer(_) - | Type::MutableReference(_) + Type::Pointer(_) | Type::Array(_) | Type::Slice(_) + | Type::TaggedI64 | Type::Str | Type::Class(_) | Type::Interface(_) @@ -350,11 +419,9 @@ impl Type { Type::Pointer(_) => Some(POINTER_CLASS.to_string()), Type::Array(_) => Some(self.to_jvm_descriptor()), // Array descriptor works for checkcast/anewarray Type::Slice(_) => Some(SLICE_VIEW_CLASS.to_string()), + Type::TaggedI64 => Some(TAGGED_LONG_CLASS.to_string()), Type::Str => Some(UTF8_VIEW_CLASS.to_string()), - Type::Reference(inner) => inner.to_jvm_descriptor_or_internal_name(), - Type::MutableReference(inner) => { - Type::Array(inner.clone()).to_jvm_descriptor_or_internal_name() - } // MutableReference is treated as an array + // MutableReference is treated as an array _ => None, } } @@ -370,11 +437,8 @@ impl Type { false } // Handle nested types recursively - Type::MutableReference(inner) - | Type::Pointer(inner) - | Type::Reference(inner) - | Type::Array(inner) - | Type::Slice(inner) => inner.replace_class(old_name, new_name), + Type::Array(inner) | Type::Slice(inner) => inner.replace_class(old_name, new_name), + Type::Pointer(inner) => inner.replace_class(old_name, new_name), // Primitive types and Void are unaffected. Type::Void | Type::Unit @@ -391,6 +455,7 @@ impl Type { | Type::F16 | Type::F32 | Type::F64 + | Type::TaggedI64 | Type::Str => { // No class names to replace here false @@ -407,10 +472,9 @@ impl Type { ); match self { Type::Class(name) | Type::Interface(name) => Some(name), + Type::TaggedI64 => Some(TAGGED_LONG_CLASS), Type::Str => Some(UTF8_VIEW_CLASS), - Type::Array(inner) | Type::MutableReference(inner) | Type::Reference(inner) => { - inner.get_class_name() - } + Type::Array(inner) => inner.get_class_name(), // Method dispatch through a Rust reference targets the pointee; // pointer-native methods are redirected explicitly during lowering. Type::Pointer(inner) => inner.get_class_name(), @@ -437,15 +501,9 @@ mod tests { (U16, "C"), (Class("java.lang.É".into()), "Ljava/lang/É;"), (Interface("java/lang/É".into()), "Ljava/lang/É;"), - (Pointer(Box::new(Unit)), "Lorg/rustlang/runtime/Pointer;"), + (Type::pointer(Unit), "Lorg/rustlang/runtime/Pointer;"), (Array(Box::new(Unit)), "[Ljava/lang/Object;"), - (MutableReference(Box::new(Unit)), "Ljava/lang/Object;"), - ( - Array(Box::new(MutableReference(Box::new(Unit)))), - "[Ljava/lang/Object;", - ), - (Reference(Box::new(Array(Box::new(U8)))), "[B"), - (MutableReference(Box::new(Array(Box::new(I8)))), "[[B"), + (Array(Box::new(Array(Box::new(I8)))), "[[B"), ]; for (ty, expected) in &cases { assert_eq!(ty.to_jvm_descriptor(), *expected); diff --git a/src/oomir/visit.rs b/src/oomir/visit.rs index 5598c07e..8575cfe1 100644 --- a/src/oomir/visit.rs +++ b/src/oomir/visit.rs @@ -6,8 +6,19 @@ macro_rules! operand_visitor { pub fn $name<'a>(&'a $($mutable)? self, mut visit: impl FnMut(&'a $($mutable)? Operand)) { use Instruction::*; match self { + TaggedPack { value, tag, .. } => { visit(value); visit(tag); } + TaggedPart { value, .. } => visit(value), Binary { op1, op2, .. } => { visit(op1); visit(op2); } Not { src, .. } | Neg { src, .. } | Move { src, .. } => visit(src), + ValueCopy { source, .. } => visit(source), + MemoryLoad { pointer, .. } | MemoryCommit { pointer } => visit(pointer), + Heap { args, .. } => args.iter().for_each(&mut visit), + MemoryProject { base, .. } => visit(base), + AddressRetype { source, layout, .. } | ViewAddress { source, layout, .. } => { + visit(source); visit(&layout.size); visit(&layout.codec); + } + AddressOffset { source, count, .. } => { visit(source); visit(count); } + MemoryStore { pointer, value, .. } => { visit(pointer); visit(value); } Branch { condition, .. } => visit(condition), Return { operand } => { if let Some(operand) = operand { visit(operand); } } InvokeStatic { args, .. } | InvokeRustStatic { args, .. } => { @@ -31,7 +42,7 @@ macro_rules! operand_visitor { SetJvmField { object, value, .. } => { visit(object); visit(value); } GetField { object, .. } | GetJvmField { object, .. } | Cast { op: object, .. } => visit(object), SourceLocation(_) | LocalVariableScope(_) | UnwindStart { .. } | UnwindEnd - | Rethrow | CreateFunctionPointer { .. } | GetStaticField { .. } | Jump { .. } + | Rethrow | Unreachable | CreateFunctionPointer { .. } | GetStaticField { .. } | Jump { .. } | ThrowNewWithMessage { .. } | Label { .. } => {} } } From 7b8fab11b77912b1f8881308e4525cdc0eb986a9 Mon Sep 17 00:00:00 2001 From: Michael Reeves Date: Thu, 1 Oct 2026 18:28:35 +1000 Subject: [PATCH 25/61] lower memory operations through components --- src/allocator_shims.rs | 9 +- src/lower1/control_flow/calls/imports.rs | 25 ++ src/oomir/construct/addresses.rs | 163 +++++++++++++ src/oomir/construct/arrays.rs | 31 ++- src/oomir/construct/heap.rs | 44 ++++ src/oomir/construct/memory.rs | 144 +++++++++++ src/oomir/construct/mod.rs | 144 ++++++++++- src/oomir/construct/operations.rs | 123 +++++++++- src/oomir/construct/pointer_tests.rs | 281 ++++++++++++++++++++++ src/oomir/construct/pointers.rs | 221 ++++++++++++----- src/oomir/construct/tests.rs | 291 ++++++++++++++++++++++- src/oomir/construct/wrappers.rs | 4 - 12 files changed, 1374 insertions(+), 106 deletions(-) create mode 100644 src/oomir/construct/addresses.rs create mode 100644 src/oomir/construct/heap.rs create mode 100644 src/oomir/construct/memory.rs create mode 100644 src/oomir/construct/pointer_tests.rs diff --git a/src/allocator_shims.rs b/src/allocator_shims.rs index 722fd914..537aa07e 100644 --- a/src/allocator_shims.rs +++ b/src/allocator_shims.rs @@ -15,7 +15,7 @@ pub(super) fn allocator_shim_target_signature( } AllocatorTy::Ptr => params.push(( input.name.to_string(), - oomir::Type::Pointer(Box::new(oomir::Type::U8)), + oomir::Type::pointer(oomir::Type::U8), )), AllocatorTy::Usize => { params.push((input.name.to_string(), oomir::Type::U64)); @@ -27,7 +27,7 @@ pub(super) fn allocator_shim_target_signature( } let ret = match method.output { - AllocatorTy::ResultPtr => oomir::Type::Pointer(Box::new(oomir::Type::U8)), + AllocatorTy::ResultPtr => oomir::Type::pointer(oomir::Type::U8), AllocatorTy::Never | AllocatorTy::Unit => oomir::Type::Void, AllocatorTy::Layout | AllocatorTy::Ptr | AllocatorTy::Usize => { panic!("invalid allocator shim output type") @@ -215,10 +215,7 @@ pub(super) fn emit_allocator_shims<'tcx>( result.clone(), ); if matches!(method.output, AllocatorTy::Never) { - instructions.push(oomir::Instruction::ThrowNewWithMessage { - exception_class: "java/lang/AssertionError".to_string(), - message: "Diverging allocator call returned unexpectedly".to_string(), - }); + instructions.push(oomir::Instruction::Unreachable); } else { instructions.push(oomir::Instruction::Return { operand: result.map(|name| oomir::Operand::Variable { diff --git a/src/lower1/control_flow/calls/imports.rs b/src/lower1/control_flow/calls/imports.rs index 0aa489d9..ad8c49b0 100644 --- a/src/lower1/control_flow/calls/imports.rs +++ b/src/lower1/control_flow/calls/imports.rs @@ -33,6 +33,31 @@ pub(super) fn emit<'tcx>( ), ); } + let heap = if !jvm_import.interface && jvm_import.class_name == oomir::POINTER_CLASS { + match (jvm_import.method_name.as_str(), rust_descriptor.as_str()) { + ("allocateBytes", "(JJ)Lorg/rustlang/runtime/Pointer;") => { + Some(oomir::HeapOp::Allocate) + } + ( + "reallocateBytes", + "(Lorg/rustlang/runtime/Pointer;JJJ)Lorg/rustlang/runtime/Pointer;", + ) => Some(oomir::HeapOp::Reallocate), + ("deallocateBytes", "(Lorg/rustlang/runtime/Pointer;)V") => { + Some(oomir::HeapOp::Deallocate) + } + _ => None, + } + } else { + None + }; + if let Some(operation) = heap { + instructions.push(oomir::Instruction::Heap { + operation, + args: oomir_operands, + dest: effective_dest, + }); + return; + } instructions.push(oomir::Instruction::InvokeStatic { class_name: jvm_import.class_name, method_name: jvm_import.method_name, diff --git a/src/oomir/construct/addresses.rs b/src/oomir/construct/addresses.rs new file mode 100644 index 00000000..8c3e20a5 --- /dev/null +++ b/src/oomir/construct/addresses.rs @@ -0,0 +1,163 @@ +//! Rust address operations carry layout and arithmetic independently of Java. +use super::*; + +fn static_size(value: &oomir::Operand) -> Option { + use oomir::{Constant as C, Operand::Constant}; + let value = match value { + Constant(C::I32(n)) => u32::try_from(*n).ok()?, + Constant(C::U32(n)) => *n, + Constant(C::I64(n)) => u32::try_from(*n).ok()?, + Constant(C::U64(n)) => u32::try_from(*n).ok()?, + _ => return None, + }; + (value <= i32::MAX as u32).then_some(value) +} + +impl Emission<'_> { + pub(super) fn address(&mut self, instruction: oomir::Instruction) -> Result<()> { + use oomir::Instruction::*; + let (dest, value) = match instruction { + AddressOffset { + dest, + source, + count, + ty, + bytes, + wrapping, + subtract, + } => { + let pointer = self.operand(source)?; + let count = self.operand(count)?; + let mut offset = self.adapt(count, self.ty(&oomir::Type::I64))?; + if subtract { + offset = self + .emit(ir::Op::Neg(offset), Some(self.ty(&oomir::Type::I64))) + .unwrap(); + } + let value = self + .emit( + ir::Op::Offset { + pointer, + offset, + bytes, + wrapping, + }, + Some(self.ty(&ty)), + ) + .unwrap(); + (dest, value) + } + AddressRetype { + dest, + source, + layout, + } => (dest, self.address_layout(source, *layout, false)?), + ViewAddress { + dest, + source, + layout, + } => (dest, self.address_layout(source, *layout, true)?), + _ => unreachable!("non-address instruction"), + }; + if let Some(dest) = dest { + self.write(&dest, value)?; + } + Ok(()) + } + + pub(super) fn address_layout( + &mut self, + source: oomir::Operand, + layout: oomir::AddressLayout, + view: bool, + ) -> Result { + let oomir::AddressLayout { + pointer_type, + size, + codec, + } = layout; + let source_type = source.get_type().expect("typed address source"); + let codec_id = match &codec { + oomir::Operand::Constant(oomir::Constant::Null(_)) => Some(None), + oomir::Operand::Constant(oomir::Constant::String(name)) => Some(Some( + self.vocabulary + .types + .find_symbol(name) + .expect("registered address codec"), + )), + _ => None, + }; + let input = self.operand(source)?; + let returns = self.ty(&pointer_type); + // DST casts keep the slice length for later tail construction. Sized casts can discard it. + let thin_view = match &pointer_type { + oomir::Type::Pointer(inner) => match inner.as_ref() { + oomir::Type::Class(name) => self + .context + .fields + .get(name) + .is_some_and(|layout| layout.direct), + oomir::Type::Interface(_) => false, + _ => true, + }, + _ => false, + }; + let typed = matches!(pointer_type, oomir::Type::Pointer(_)) + && if view { + thin_view && matches!(source_type, oomir::Type::Slice(_) | oomir::Type::Str) + } else { + matches!(source_type, oomir::Type::Pointer(_)) + }; + if typed && let (Some(size), Some(codec)) = (static_size(&size), codec_id) { + let op = if view { + let mut values = Vec::with_capacity(3); + for (index, ty) in [types::object(), oomir::Type::I32, oomir::Type::U64] + .iter() + .enumerate() + { + values.push( + self.emit( + ir::Op::ViewPart { + view: input, + index: index as u8, + }, + Some(self.ty(ty)), + ) + .unwrap(), + ); + } + ir::Op::ViewAddress { + parts: self.builder.args(values), + size, + codec, + } + } else if pointer_type.scalar_address_size() == Some(size) && codec.is_none() { + ir::Op::Cast(input) + } else { + ir::Op::RetypeAddress { + pointer: input, + size, + codec, + } + }; + return Ok(self.emit(op, Some(returns)).unwrap()); + } + let source_type = if view { types::object() } else { source_type }; + let size = self.operand(size)?; + let codec = self.operand(codec)?; + Ok(self + .call( + oomir::POINTER_CLASS.into(), + if view { "fromSlice" } else { "retype" }.into(), + vec![ + self.ty(&source_type), + self.ty(&oomir::Type::U64), + self.ty(&oomir::Type::java_string()), + ], + returns, + ir::CallKind::JvmStatic, + vec![input, size, codec], + )? + .unwrap()) + } +} diff --git a/src/oomir/construct/arrays.rs b/src/oomir/construct/arrays.rs index 4b6f2fe1..15fe33f0 100644 --- a/src/oomir/construct/arrays.rs +++ b/src/oomir/construct/arrays.rs @@ -28,7 +28,14 @@ impl Emission<'_> { return self.write(&dest, value); } let value = self - .emit(Op::ArrayGet { array, index }, Some(element)) + .emit( + Op::ArrayGet { + native: false, + array, + index, + }, + Some(element), + ) .unwrap(); self.write(&dest, value) } @@ -48,20 +55,12 @@ impl Emission<'_> { return Ok(()); } if copy && self.vocabulary.types.get(element).unwrap().carrier() == 5 { - value = self - .call( - oomir::POINTER_CLASS.into(), - "copyManagedValue".into(), - vec![self.ty(&types::object())], - self.ty(&types::object()), - CallKind::JvmStatic, - vec![value], - )? - .unwrap(); + value = self.copy_value(value); } let value = self.adapt(value, element)?; self.emit( Op::ArraySet { + native: false, array, index, value, @@ -78,6 +77,16 @@ impl Emission<'_> { ) -> Result<()> { let array = self.operand(array)?; let value = self.operand(value)?; + if let Some(ir::Type::Array(element)) = self + .vocabulary + .types + .get(self.builder.body.value_type(array)) + && ir::StorageSlot::scalar(element, &self.vocabulary.types).is_some() + { + let value = self.adapt(value, element)?; + self.emit(Op::ArrayFill { array, value }, None); + return Ok(()); + } let copy = self.constant(oomir::Constant::Boolean(copy))?; self.call( oomir::POINTER_CLASS.into(), diff --git a/src/oomir/construct/heap.rs b/src/oomir/construct/heap.rs new file mode 100644 index 00000000..2b58e7d5 --- /dev/null +++ b/src/oomir/construct/heap.rs @@ -0,0 +1,44 @@ +//! Platform allocation enters SSA as storage operations, before JVM selection. +use super::*; + +impl Emission<'_> { + pub(super) fn heap( + &mut self, + operation: ir::HeapOp, + operands: Vec, + dest: Option, + ) -> Result<()> { + let pointer = self.ty(&oomir::Type::pointer(oomir::Type::U8)); + let object = self.ty(&oomir::Type::Class("java/lang/Object".into())); + let long = self.ty(&oomir::Type::I64); + let mut values = operands + .into_iter() + .map(|op| self.operand(op)) + .collect::>>()?; + let mut args = Vec::new(); + if operation != ir::HeapOp::Allocate { + let address = self.adapt(values.remove(0), pointer)?; + for (index, ty) in [(0, object), (1, long)] { + args.push( + self.emit(ir::Op::AddressPart { address, index }, Some(ty)) + .unwrap(), + ); + } + } + for value in values { + args.push(self.adapt(value, long)?); + } + let returns = operation != ir::HeapOp::Deallocate; + let args = self.builder.args(args); + let root = self.emit(ir::Op::Heap { operation, args }, returns.then_some(object)); + if let (Some(root), Some(dest)) = (root, dest) { + let zero = self.constant(oomir::Constant::I64(0))?; + let parts = self.builder.args([root, zero]); + let address = self + .emit(ir::Op::AddressPack(parts), Some(pointer)) + .unwrap(); + self.write(&dest, address)?; + } + Ok(()) + } +} diff --git a/src/oomir/construct/memory.rs b/src/oomir/construct/memory.rs new file mode 100644 index 00000000..fed2eb2c --- /dev/null +++ b/src/oomir/construct/memory.rs @@ -0,0 +1,144 @@ +//! Semantic places and memory accesses become SSA without helper recognition. +use super::*; +use ir::{CallKind, Op}; + +impl Emission<'_> { + pub(super) fn copy_value(&mut self, source: ir::ValueId) -> ir::ValueId { + let ty = self.builder.body.value_type(source); + if matches!( + self.vocabulary.types.get(ty), + Some( + ir::Type::TaggedI64 + | ir::Type::Scalar(_) + | ir::Type::Pointer(_) + | ir::Type::Slice(_) + | ir::Type::Str + ) + ) { + source + } else { + self.emit(Op::CopyValue(source), Some(ty)).unwrap() + } + } + + pub(super) fn memory(&mut self, instruction: oomir::Instruction) -> Result<()> { + use oomir::Instruction::*; + match instruction { + ValueCopy { dest, source } => { + let source = self.operand(source)?; + let value = self.copy_value(source); + self.write(&dest, value)?; + } + MemoryLoad { + dest, + pointer, + pointee, + owned, + } => { + let pointer = self.operand(pointer)?; + let pointer = + self.adapt(pointer, self.ty(&oomir::Type::pointer(pointee.clone())))?; + let value = self + .emit( + if owned && pointee.is_jvm_reference_type() { + Op::LoadCopy(pointer) + } else { + Op::Load(pointer) + }, + Some(self.ty(&pointee)), + ) + .unwrap(); + self.write(&dest, value)?; + } + MemoryStore { + pointer, + pointee, + value, + } => { + let pointer = self.operand(pointer)?; + let pointer = + self.adapt(pointer, self.ty(&oomir::Type::pointer(pointee.clone())))?; + let value = self.operand(value)?; + let value = self.adapt(value, self.ty(&pointee))?; + self.emit(Op::Store { pointer, value }, None); + } + MemoryCommit { pointer } => { + let pointer = self.operand(pointer)?; + self.emit(Op::Commit(pointer), None); + } + MemoryProject { + dest, + base, + projection, + } => { + let oomir::MemoryProjection { + owner, + field, + pointee, + offset, + size, + codec, + } = *projection; + let base_type = oomir::Type::Class(owner.clone()); + let pointer_type = oomir::Type::pointer(base_type.clone()); + let result_type = oomir::Type::pointer(pointee.clone()); + let direct = self.context.fields.get(&owner).is_some_and(|layout| { + layout.direct + && layout + .members + .iter() + .any(|(name, ty)| name == &field && ty == &pointee) + }) && base.get_type().as_ref() == Some(&pointer_type); + let base = self.operand(base)?; + let value = if direct { + let base = self.adapt(base, self.ty(&pointer_type))?; + let field = self.builder.field(ir::FieldRef { + owner: self.ty(&base_type), + name: field, + ty: self.ty(&pointee), + is_static: false, + }); + let projection = self.builder.projection(ir::PointerProjection { + field, + offset, + size, + codec, + }); + self.emit( + Op::Project { base, projection }, + Some(self.ty(&result_type)), + ) + .unwrap() + } else { + let mut args = vec![base]; + for value in [ + oomir::Constant::String(owner), + oomir::Constant::LiteralString(field), + oomir::Constant::U64(offset), + oomir::Constant::U64(size), + codec.map_or_else( + || oomir::Constant::Null(oomir::Type::java_string()), + oomir::Constant::String, + ), + ] { + args.push(self.constant(value)?); + } + let string = self.ty(&oomir::Type::java_string()); + let long = self.ty(&oomir::Type::U64); + self.call( + oomir::POINTER_CLASS.into(), + "projectStructField".into(), + vec![string, string, long, long, string], + self.ty(&result_type), + CallKind::Virtual, + args, + )? + .unwrap() + }; + self.write(&dest, value)?; + } + _ => unreachable!("non-memory instruction"), + } + Ok(()) + } +} diff --git a/src/oomir/construct/mod.rs b/src/oomir/construct/mod.rs index b610cc20..0764ad30 100644 --- a/src/oomir/construct/mod.rs +++ b/src/oomir/construct/mod.rs @@ -9,12 +9,20 @@ use jvm_compiler_core::{ }; use rustc_hash::{FxHashMap as HashMap, FxHashSet as HashSet}; use std::sync::Arc; +mod addresses; mod arithmetic; mod arrays; mod context; +mod heap; +mod memory; mod operations; +#[cfg(test)] +mod pointer_tests; mod pointers; mod types; +#[cfg(test)] +mod view_tests; +mod views; mod wrappers; pub(crate) use context::Context; use types::Vocabulary; @@ -255,7 +263,7 @@ pub(crate) fn seal(function: oomir::Function, context: &Context) -> Result Result