module Runtime.NativeFiniteWords import Compiler.MachineX86Native import Compiler.MachineX86NativeAssembly import Std.Foundation import Std.Natural family NativeFiniteWordFormat : Type 0 constructor NativeFiniteF32 constructor NativeFiniteBF16 end-family def nativeFiniteFormatTag = (lambda unrestricted format : (family NativeFiniteWordFormat) . (eliminate NativeFiniteWordFormat (lambda unrestricted current : (family NativeFiniteWordFormat) . Nat) format (branch NativeFiniteF32 . 0) (branch NativeFiniteBF16 . 1))) -- The checkpoint validator examines representations, not floating-point -- arithmetic: exponent 255 means infinity or NaN in either IEEE binary32 or -- bfloat16. This is the same predicate used by the native host routine below. def finiteF32Bits = (lambda unrestricted bits : Nat . (naturalSelect (naturalEqual (naturalModuloUnchecked (naturalDivideUnchecked bits 8388608) 256) 255) 0 1)) def finiteBF16Bits = (lambda unrestricted bits : Nat . (naturalSelect (naturalEqual (naturalModuloUnchecked (naturalDivideUnchecked bits 128) 256) 255) 0 1)) def finiteBF16Pair = (lambda unrestricted packed : Nat . (naturalMultiply (finiteBF16Bits (naturalModuloUnchecked packed 65536)) (finiteBF16Bits (naturalDivideUnchecked packed 65536)))) def finiteRAX = (constructor X86NativeRegister64 X86NativeRAX) def finiteRCX = (constructor X86NativeRegister64 X86NativeRCX) def finiteRDI = (constructor X86NativeRegister64 X86NativeRDI) def finiteRSI = (constructor X86NativeRegister64 X86NativeRSI) def finiteR8 = (constructor X86NativeRegister64 X86NativeR8) def finiteEmit = (lambda unrestricted instruction : (family X86NativeInstruction) . (lambda unrestricted tail : (family X86NativeAssembly) . (constructor X86NativeAssembly X86NativeAssemblyEmit instruction tail))) def finiteLabel = (lambda unrestricted name : Bytes . (lambda unrestricted tail : (family X86NativeAssembly) . (constructor X86NativeAssembly X86NativeAssemblyLabel name tail))) def finiteJump = (lambda unrestricted name : Bytes . (lambda unrestricted tail : (family X86NativeAssembly) . (constructor X86NativeAssembly X86NativeAssemblyJump name tail))) def finiteBranch = (lambda unrestricted condition : (family X86NativeCondition) . (lambda unrestricted name : Bytes . (lambda unrestricted tail : (family X86NativeAssembly) . (constructor X86NativeAssembly X86NativeAssemblyJumpCondition condition name tail)))) def finiteMove = (lambda unrestricted destination : (family X86NativeRegister64) . (lambda unrestricted source : (family X86NativeRegister64) . (finiteEmit (constructor X86NativeInstruction X86NativeMoveRegister64 source destination)))) def finiteAnd = (lambda unrestricted register : (family X86NativeRegister64) . (lambda unrestricted mask : Nat . (finiteEmit (constructor X86NativeInstruction X86NativeAndImmediate64 register (x86NativeImmediate32FromNatural mask))))) def finiteCompare = (lambda unrestricted register : (family X86NativeRegister64) . (lambda unrestricted value : Nat . (finiteEmit (constructor X86NativeInstruction X86NativeCompareImmediate64 register (x86NativeImmediate32FromNatural value))))) def finiteWordCheck = (lambda unrestricted half : Nat . (lambda unrestricted tail : (family X86NativeAssembly) . (nat-eliminate (lambda unrestricted mode : Nat . (family X86NativeAssembly)) (finiteAnd finiteRAX 2139095040 (finiteCompare finiteRAX 2139095040 (finiteBranch (constructor X86NativeCondition X86NativeConditionZero) b"reject" tail))) (lambda unrestricted predecessor : Nat . (lambda unrestricted unused : (family X86NativeAssembly) . (finiteMove finiteRCX finiteRAX (finiteAnd finiteRAX 32640 (finiteCompare finiteRAX 32640 (finiteBranch (constructor X86NativeCondition X86NativeConditionZero) b"reject" (finiteEmit (constructor X86NativeInstruction X86NativeShiftRightImmediate64 finiteRCX (constructor X86NativeImmediate8 X86NativeImmediate8Value (byte 16))) (finiteAnd finiteRCX 32640 (finiteCompare finiteRCX 32640 (finiteBranch (constructor X86NativeCondition X86NativeConditionZero) b"reject" tail)))))))))) half))) -- System V: RDI points to `extent` mapped bytes; RSI is that extent. The -- caller chooses the format by embedding one of these two routine images. -- Every four-byte word is checked, including both halves of a BF16 pair. -- Zero or unaligned extents and a wrapping end address fail closed. The -- caller must prove the mapped span belongs to its typed staging region. def nativeFiniteWordsAssembly = (lambda unrestricted half : Nat . (finiteEmit (constructor X86NativeInstruction X86NativeTestRegister64 finiteRSI finiteRSI) (finiteBranch (constructor X86NativeCondition X86NativeConditionZero) b"reject" (finiteMove finiteRAX finiteRSI (finiteAnd finiteRAX 3 (finiteBranch (constructor X86NativeCondition X86NativeConditionNotZero) b"reject" (finiteMove finiteR8 finiteRDI (finiteEmit (constructor X86NativeInstruction X86NativeAddRegister64 finiteRSI finiteR8) (finiteEmit (constructor X86NativeInstruction X86NativeCompareRegister64 finiteRDI finiteR8) (finiteBranch (constructor X86NativeCondition X86NativeConditionBelow) b"reject" (finiteLabel b"scan" (finiteEmit (constructor X86NativeInstruction X86NativeLoadMemory32ZeroExtend64 finiteRAX finiteRDI (x86NativeDisplacement32FromNatural 0)) (finiteWordCheck half (finiteEmit (constructor X86NativeInstruction X86NativeAddImmediate64 finiteRDI (x86NativeImmediate32FromNatural 4)) (finiteEmit (constructor X86NativeInstruction X86NativeCompareRegister64 finiteR8 finiteRDI) (finiteBranch (constructor X86NativeCondition X86NativeConditionZero) b"accept" (finiteJump b"scan" (finiteLabel b"accept" (finiteEmit (constructor X86NativeInstruction X86NativeClear32 finiteRAX) (finiteEmit (constructor X86NativeInstruction X86NativeReturn) (finiteLabel b"reject" (finiteEmit (constructor X86NativeInstruction X86NativeMoveImmediate32 finiteRAX (x86NativeImmediate32FromNatural 1)) (finiteEmit (constructor X86NativeInstruction X86NativeReturn) (constructor X86NativeAssembly X86NativeAssemblyEnd)))))))))))))))))))))))) def nativeFiniteWordsCode = (lambda unrestricted half : Nat . (eliminate X86NativeAssemblyResult (lambda unrestricted current : (family X86NativeAssemblyResult) . Bytes) (x86NativeAssemble (nativeFiniteWordsAssembly half)) (branch X86NativeAssemblyEncoded code . code) (branch X86NativeAssemblyEncodeDuplicateLabel name . b"") (branch X86NativeAssemblyEncodeOffsetOverflow . b"") (branch X86NativeAssemblyMissingLabel name . b"") (branch X86NativeAssemblyDisplacementOutOfRange name . b""))) def nativeFiniteF32Routine : Bytes = (nativeFiniteWordsCode 0) def nativeFiniteBF16Routine : Bytes = (nativeFiniteWordsCode 1) def nativeFiniteRoutineFor = (lambda unrestricted format : (family NativeFiniteWordFormat) . (eliminate NativeFiniteWordFormat (lambda unrestricted current : (family NativeFiniteWordFormat) . Bytes) format (branch NativeFiniteF32 . nativeFiniteF32Routine) (branch NativeFiniteBF16 . nativeFiniteBF16Routine)))