diff --git a/UnitTest/AttrParser.lean b/UnitTest/AttrParser.lean index 0d63f8f282..aa63bf3d0e 100644 --- a/UnitTest/AttrParser.lean +++ b/UnitTest/AttrParser.lean @@ -546,6 +546,7 @@ macro "#assert " e:term : command => /-! ## LLVM Byte type -/ #assert expectSuccessType "!llvm.byte<64>" (LLVM.ByteType.mk 64) +#assert expectSuccessType "!llvm.array<2 x byte<8>>" (LLVM.ArrayType.mk 2 $ LLVM.ByteType.mk 8) /-! ## LLVM Struct type (parsed opaquely; see `parseOptionalLLVMStructType`) diff --git a/Veir/Parser/AttrParser.lean b/Veir/Parser/AttrParser.lean index e7c7e0aa99..538ce3e6b3 100644 --- a/Veir/Parser/AttrParser.lean +++ b/Veir/Parser/AttrParser.lean @@ -133,15 +133,18 @@ def parseOptionalFloatType : AttrParserM (Option FloatType) := do /-- Parse an optional byte type. - A byte type is represented as `!llvm.byte` where bitwidth is a positive integer. + A byte type is represented as `!llvm.byte`, or `byte` when `short`. -/ -def parseOptionalByteType : AttrParserM (Option LLVM.ByteType) := do - let token ← peekToken - let .exclamationIdent := token.kind | return none - let input := (← getThe ParserState).input - let typeName := { token.slice with start := token.slice.start + 1 }.of input - if typeName ≠ "llvm.byte".toByteArray then return none - let _ ← consumeToken +def parseOptionalByteType (short := false) : AttrParserM (Option LLVM.ByteType) := do + if short then + let .true ← parseOptionalKeyword "byte".toByteArray | return none + else + let token ← peekToken + let .exclamationIdent := token.kind | return none + let input := (← getThe ParserState).input + let typeName := { token.slice with start := token.slice.start + 1 }.of input + if typeName ≠ "llvm.byte".toByteArray then return none + let _ ← consumeToken parsePunctuation "<" let bitwidth ← parseInteger false false parsePunctuation ">" @@ -1109,8 +1112,8 @@ partial def parseOptionalLLVMStructType (short := false) : AttrParserM (Option T /-- Parse a type within an LLVM-dialect type body, accepting the LLVM "pretty-print" sugar keywords `void`, `ptr`, `x86_amx`, `ppc_fp128`, `label`, `metadata`, `token`, and - the bare nested forms `array<...>`, `struct<...>`, `func<...>`, and `target<...>` in - addition to the regular MLIR type forms. + the bare nested forms `byte<...>`, `array<...>`, `struct<...>`, `func<...>`, and + `target<...>` in addition to the regular MLIR type forms. The bare nested form exists because the LLVM dialect has a custom directive `PrettyLLVMType`, which allows types from the LLVM dialect to be written @@ -1134,6 +1137,8 @@ partial def parseLLVMType (errorMsg : String := "type expected") : AttrParserM T return LLVM.VoidType.mk if ← parseOptionalKeyword "ptr".toByteArray then return (LLVM.PointerType.mk : TypeAttr) + if let some type ← parseOptionalByteType true then + return type if let some type ← parseOptionalLLVMArrayType true then return type if let some type ← parseOptionalLLVMStructType true then