diff --git a/IR.g4 b/IR.g4 new file mode 100644 index 00000000..ff003212 --- /dev/null +++ b/IR.g4 @@ -0,0 +1,82 @@ +grammar IR; + +// --- Parser rules --- + +program: line* EOF; + +line: labelMarker | instruction; + +labelMarker: LABEL; + +instruction + : dest '=' instrCall # instrWithDest + | instrCall # instrNoDest + | dest '=' const_ # constAssign + | dest '=' left=reg op=(PLUS | MINUS) right=reg # infixInstr + | left=reg op=(PLUS_EQ | MINUS_EQ) right=reg # compoundAssignInstr + ; + +dest + : reg # destReg + | '_' # destDiscard + | '[' regList ']' # destList + ; + +regList: reg (',' reg)*; + +instrCall: instrName '(' args ')'; + +instrName: IDENTIFIER ('<' typeName '>')?; + +typeName: TYPE_KEYWORD; + +args: (arg (',' arg)*)?; + +arg + : value # positionalArg + | IDENTIFIER ':' value # labeledArg + ; + +value + : reg # valReg + | LABEL # valLabel + | INT # valInt + | '[' regList ']' # valRegList + ; + +const_ + : STRING # constString + | MINUS? INT # constInt + | BOOL # constBool + ; + +reg: REG; + +// --- Lexer rules --- + +WS: [ \t]+ -> skip; +NEWLINE: [\r\n]+ -> skip; + +// Must come before IDENTIFIER so keywords are not swallowed +TYPE_KEYWORD: 'int' | 'str' | 'portion' | 'monetary'; +BOOL: 'true' | 'false'; + +REG: '$' [a-zA-Z_] [a-zA-Z0-9_]*; +LABEL: '#' [a-zA-Z_] [a-zA-Z0-9_]*; +INT: [0-9]+; +STRING: '"' ('\\"' | ~[\r\n"])* '"'; +IDENTIFIER: [a-z] [a-z0-9_]*; + +LPAREN: '('; +RPAREN: ')'; +LBRACKET: '['; +RBRACKET: ']'; +COMMA: ','; +EQ: '='; +PLUS: '+'; +MINUS: '-'; +PLUS_EQ: '+='; +MINUS_EQ: '-='; +LT: '<'; +GT: '>'; +UNDERSCORE: '_'; diff --git a/Justfile b/Justfile index ae23cdee..b586574f 100644 --- a/Justfile +++ b/Justfile @@ -15,6 +15,7 @@ tidy: generate: @antlr4 -Dlanguage=Go Lexer.g4 Numscript.g4 -o internal/parser/antlrParser -package antlrParser @mv internal/parser/antlrParser/_lexer.go internal/parser/antlrParser/lexer.go + @antlr4 -Dlanguage=Go IR.g4 -o internal/ir/internal/syntax/antlrParser -package antlrParser tests: @go test -race -covermode=atomic \ diff --git a/bytecode-verifier.md b/bytecode-verifier.md new file mode 100644 index 00000000..facc9760 --- /dev/null +++ b/bytecode-verifier.md @@ -0,0 +1,119 @@ +# Bytecode Verifier + +`vm.Verify` is a static pass over a `vm.Program`. A nil result means the +execution loop cannot read out of bounds or crash on that program. + +It is **opt-in**. `Exec` does not call it, and neither does `compiler.Compile`: +the VM assumes its own compiler's output is well formed, and paying for a full +static pass on every compile would make that assumption cost something. The +assumption is earned by tests instead — see [What keeps it +honest](#what-keeps-it-honest). + +Run it on any program that did not come out of `ir.Assemble` in this process: +anything read from a file, a wire, or a cache. + +```go +program, err := vm.DecodeProgram(bytes) +if err != nil { ... } +if err := vm.VerifyWithVars(program, vars); err != nil { ... } +result, execErr := vm.Exec(ctx, vm.NewVm(program), vars, store) +``` + +`Verify(p)` checks the program alone. `VerifyWithVars(p, vars)` also checks that +`vars` carries every variable the program loads — a program that reads a +variable is only safe against the vars it will actually be given, so a caller +that passes vars should use the second form. + +## What it checks + +| Check | What it prevents | +|---|---| +| Instruction stream decodes end to end | A truncated multi-word instruction, whose ext word `Exec` reads past the end of the stream | +| Every opcode is known | `Exec` falling through to its `default` arm | +| Const-pool indices are in range | `Op_LoadInt`/`Op_LoadStr` indexing past the pool | +| Jumps land on an instruction boundary | Landing on an ext word, whose ignored opcode byte would then be decoded as an instruction | +| Register indices are below the bank's declared `MaxReg` | Indexing past a register bank — the banks are sized from those counts, so this is what makes the sizing safe | +| `nilReg` only in optional operands | Reading register 255 where `Exec` dereferences unconditionally | +| Flag operands are 0 or 1 | `Op_MarkEnd` tests `A == 1`, so a 2 silently commits a region that meant to rewind | +| Definite assignment | Reading a register not written on every path reaching the instruction | +| The current asset is set before it is used | A `send` or `pull` before any `set_current_asset` | +| Mark discipline (see below) | An `Op_MarkEnd` with no open mark, a send/save/asset change inside a region, a run ending with a region open | +| (`VerifyWithVars`) Var-pool indices are in range | `Op_LoadVar*` indexing past the pool, or dereferencing a nil `*Vars` | + +Two properties come for free rather than as their own pass: + +- **Type confusion.** A register's bank is part of its identity, so a slot + written as an int and later read as a string is a read of a register that was + never written. `ir.Typecheck` covers the same ground one level up, on the IR, + where a register still has a name. +- **Termination.** Jump deltas are unsigned and relative to the following + instruction, so a backward jump cannot be encoded. This is also what makes the + definite-assignment dataflow a single ordered pass rather than a worklist: + every predecessor of a step is earlier in the stream. + +### Mark discipline + +`Op_MarkPush`/`Op_MarkEnd` take no operand, so mark depth is a function of +position in the instruction stream. The pass carries a depth along the same +forward dataflow as definite assignment, and rejects a program where: + +- predecessors disagree on the depth at a join, +- an `Op_MarkEnd` runs at depth 0, +- an `Op_SendToAccount`, `Op_SetCurrentAsset` or `Op_Save` runs at depth > 0, +- the run can end (falling off the last instruction, or jumping past it) at depth > 0. + +`ir/instr.go` asks emitters to keep this decidable ("never emit a mark op on +only one side of a branch"). + +The VM still enforces the same rules at execution time, via +`runstate.HasOpenMark()` and the `errSendWhileMarkOpen` / +`errSetAssetWhileMarkOpen` / `errSaveWhileMarkOpen` sentinels in `vm.go`: +`Exec` does not require a verified program. + +**`internal/funds` must not be relaxed along with it.** The tree-walking +interpreter shares `RunState.MarkEnd` (see `interpreter.go`), and it is not +verified, so `ErrNoOpenMark` and the INVARIANT documented on `MarkEnd` stay. + +## What it does not check + +### Unbounded sources + +`Op_PullAccount` with both cap and overdraft nil is statically decidable, but +it is a legitimate user-facing error (`InvalidUncappedSource`, "unbounded source +is not allowed here"), not a malformed program. Deliberately left to run time. + +### Anything semantic + +A verified program can still produce nonsense: allotment portions that don't sum +to 1, assets that don't line up across a pull and a send, postings that make no +business sense. The verifier is about the VM's own memory safety, and the fuzz +tests tolerate garbage output on purpose (`_, _ = vm.Exec(...)`). + +### Totality of `Exec` + +`Exec` is total only for callers that verified first. Making it unconditionally +total would mean verifying inside `Exec`, and paying for it on a path where the +bytecode is nearly always the compiler's own. + +### Cost + +Definite assignment is O(steps × registers), with a map allocated per step. +Fine at current program sizes; if programs grow, the assigned-sets want to be +bitsets over a dense register numbering. + +## What keeps it honest + +The verifier is a second, independent model of what each opcode reads and +writes — `decodeInstr` mirrors, operand for operand, the matching arm of `Exec`. +Two models drift. Four things push back: + +| Test | Property | +|---|---| +| `compiler.TestCompiledCorpusPassesVerify` | Every script in the corpus compiles to bytecode the verifier accepts | +| `vm.assembleIR` (all of `ir_test.go`) | Every IR-driven test program verifies, including sequences the compiler never emits | +| `vm.FuzzExec` | Arbitrary bytes: whatever verifies, `Exec` runs without panicking | +| `compiler.FuzzMutatedBytecode` | Corrupted real bytecode: same property, but close enough to valid to reach the deep checks | + +An opcode added to `instruction.go` but not to `decodeInstr` is rejected as +unknown, so the corpus and IR tests fail loudly rather than silently skipping +it. That is the intended failure mode. diff --git a/bytecode_version_test.go b/bytecode_version_test.go new file mode 100644 index 00000000..8a6eb7eb --- /dev/null +++ b/bytecode_version_test.go @@ -0,0 +1,62 @@ +package numscript_test + +import ( + "bytes" + "encoding/binary" + "testing" + + "github.com/formancehq/numscript" + "github.com/stretchr/testify/require" +) + +// The public surface a host needs to keep stored bytecode and the executing +// build in step: the current version, the version stamped on what Compile and +// Encode produce, a peek that reads it off raw bytes, and the typed rejection. +func TestBytecodeVersionPublicAPI(t *testing.T) { + varsEncoder, program, err := numscript.Compile(`vars { + monetary $amt +} + +send $amt ( + source = @src + destination = @dst +)`) + require.NoError(t, err) + require.Equal(t, numscript.CurrentBytecodeVersion, program.Version) + + programBytes := program.Encode() + peeked, err := numscript.PeekCompiledProgramVersion(programBytes) + require.NoError(t, err) + require.Equal(t, numscript.CurrentBytecodeVersion, peeked) + + decoded, err := numscript.DecodeCompiledProgram(programBytes) + require.NoError(t, err) + require.Equal(t, numscript.CurrentBytecodeVersion, decoded.Version) + + vars, err := varsEncoder.Encode(map[string]string{"amt": "USD/2 100"}) + require.NoError(t, err) + require.Equal(t, numscript.CurrentBytecodeVersion, vars.Version) + + varsBytes := vars.Encode() + peeked, err = numscript.PeekVarsVersion(varsBytes) + require.NoError(t, err) + require.Equal(t, numscript.CurrentBytecodeVersion, peeked) + + // A blob from the next major is reported, not misread: the peek still + // tells which version it is, the decoders reject it with the typed error. + next := numscript.BytecodeVersion{Major: numscript.CurrentBytecodeVersion.Major + 1} + foreign := bytes.Clone(programBytes) + binary.LittleEndian.PutUint16(foreign[4:], next.Major) + binary.LittleEndian.PutUint16(foreign[6:], next.Minor) + + peeked, err = numscript.PeekCompiledProgramVersion(foreign) + require.NoError(t, err) + require.Equal(t, next, peeked) + require.False(t, numscript.CurrentBytecodeVersion.CanRead(peeked)) + + _, err = numscript.DecodeCompiledProgram(foreign) + var unsupported numscript.UnsupportedBytecodeVersionError + require.ErrorAs(t, err, &unsupported) + require.Equal(t, next, unsupported.Encoded) + require.Equal(t, numscript.CurrentBytecodeVersion, unsupported.Supported) +} diff --git a/compiler-architecture.md b/compiler-architecture.md new file mode 100644 index 00000000..22619346 --- /dev/null +++ b/compiler-architecture.md @@ -0,0 +1,543 @@ +# Numscript vm architecture + +## Design goals + +Because of new requirements in the ledger v3, we are redesigning how numscript is architected. +Previous architecture (still implemented in this repo) consisted in a simple data flow: +We parsed the ast, and then walked the Ast to interpret it and emit postings: + +```mermaid +stateDiagram + direction TB + Still --> s8:parse + s8 --> s7:run (with vars) + Still:source code + s8:Ast + s7:postings +``` + +while this is the simplest architecture we could implement, and its performance was still good enough for our use cases, a few things changed with the ledger v3 design: + +1. Higher ledger TPS: numscript is more likely to become a bottleneck. So it's now justified to pay with more complexity for better perfs. Tree walker interpreter is usually suboptimal, has a lot of pointer chasing, and makes it hard to optimise things +2. It needs a way to send the programs around the nodes in a compact and efficient way. This could be solved in the previous architecture by using the syntax itself as a serializations format, or with some rpc encoding, but would still require complex operations on the nodes +3. We now have a specific section which is sequential and needs maximum perf. So we now prefer an architecture that allows to have pre-computed optimisations in the parallel path, so that the sequential one is highly optimised + +that lead to researching a new implementation that would fit those design goals better + +## The overall architecture + +We now compile the parsed Ast (using `compiler.Compile`) and get a `vm.Program` struct and a `compiler.VarsEncoder` struct (or a compilation error). The compiler can optionally run an optimisation pass. + +The `vm.Program` is used to create a `vm.Vm` instance (via `vm.NewVm`). Creating a `vm.Vm` instance will allocate the relevant registers and data, but it's designed so that we can reuse the same `vm.Vm` instance across execution of the same program. + +The `VarsEncoder` struct knows how to encode a json payload (a `map[string]string`) into `vm.Vars`. + +Finally, we can obtain our postings and meta output by running the `vm.Exec` function, by passing the `vm.Vm` instance, a `Store` implementation (used by the vm to fetch balances and meta), and the `vm.Vars`. + +An important property: Both `vm.Program` and `vm.Vars` can be encoded and decoded as bytes. This way, we can orchestrate the previously mentioned flow: + +- The leader node can parse and pre-compile numscript into a `vm.Program` and keep the `compiler.VarsEncoder`. The `vm.Program` is serialised into bytes and sent to nodes, which'll deserialise it into `vm.Program` again, and used to create the instance of the `vm.Vm`. +- On each tx, the leader takes the json payload and turns it into `vm.Vars` with the `vm.Encoder`. `vm.Vars` are serialised, sent and deserialised back. Now node can finally run the highly optimised warm `vm.Vm` instance. + +Diagram is roughly like this: + +```mermaid +stateDiagram + direction TB + s2 --> s3:encode + s3 --> s2:decode + s2 --> s4 + s4 --> s7:exec + Still --> s8:parse + s8 --> s2:compilation + s8 --> s6:compilation + s6 --> s9:encode vars payload + s9 --> s10:encode + s10 --> s9:decode + s9 --> s7 + s2:compiled program + s3:program bytecode repr + s4:vm instance + s7:postings + Still:source code + s8:Ast + s6:vars encoder + s9:vm.Vars + s10:vars bytecode repr +``` + +In the next section we'll see how each of those building blocks works exactly + +## Vm architecture + +The vm's instance is composed by: + +- the `vm.Program` (the compilation artifact) +- registers +- the runtime's state, which keeps track of allocated funds and the accounts' balances. + +The bytecode's instructions have a fixed 32bit size: + +```go +package vm + +type Instruction struct { + Opcode byte + A byte + B byte + C byte +} +``` + +Because of the fixed-size, the vm can keep the hydrated buffer of `[]Instruction` stream instead of having to parse things on the fly at runtime (or without having to model it via heap-allocated structs) while still being compact enough to benefit good CPU locality. + +Instructions come in 2 formats: either `ABC` (3 arguments of 1 byte each) or `ABB` (2 arguments, with 1 having 1 byte size and the other a little endian repr of 2bytes). +If an instruction doesn't fit the 4 bytes limit, we simply extend it with the `Instruction` after that. + +Instructions are fetched and evaluated one at the time until they are finished (no HALT instructions to stop, so that bytecode always terminates by design). + +Instructions can move data by manipulating registers. Registers banks are separated by type (so that we don't have to have a single heap-allocated value, nor unsafe pointers or manually handled unsafe memory). With "type" here we mean the internal representation of data, which isn't the same as numscript types (there isn't necessarily a 1-1 relationship). For example, both strings, assets and accounts are represented via the golang `string` type. The `bool` bank is the other direction: it has no numscript counterpart at all, and exists so that a branch condition can't be a monetary quantity. + +> Note: scopes turned out not to need a representation change. `scoped(account, "scope")` +> compiles to a plain second `string` register carried alongside the account's own, using +> the same nilable-optional-operand idiom `Color`/`Overdraft` already use on `PullAccount` — +> exactly the same non-bank treatment `Monetary` gets as a `(str, int)` pair. See +> `compiler.compileAccountExpr`/`accountValue` and `ir.AssertValidScope`. + +A simple example of an instruction is: + +``` +ADD_INT 0x00 0x01 0x02 +``` + +which behaves like this: + +``` +int_registers[0x00] = int_registers[0x01] + int_registers[0x02] +``` + +Note that this model plays very well with golang's `big.Int` mutable API. + +The instruction set has + +- a few pure, binary or unary arithmetic/logic operations (string add, int add, portion mul and sub, int/portion conversions, etc). There is no min instruction: a min is a comparison plus a copy +- comparisons that write a bool register (`LT_INT`, `EQ_INT`, `STR_EQ`, `IS_ZERO`, ...) and `NOT` +- a few domain instructions which call the `funds.RunState`'s API (such as `PULL_ACCOUNT`, `SEND_TO_ACCOUNT`, `SAVE`). Those domain primitives can allocate funds, pull them to allocate postings, etc. This runtime logic is shared with the interpreter implementation. +- `MARK_PUSH` / `MARK_END`, which open and close a backtracking region for `oneof` +- jumps on a bool register (`JMP_IF_TRUE`, `JMP_IF_FALSE`) and an unconditional `JMP`, which can only jump forward (so that the vm always halts by design) +- constant pool loading instructions: `LOAD_STR(dest:u8, idx:u16)`, which performs `str_regs[dest] = program.str_pool[idx]`, and `LOAD_INT`. +- `LOAD_VAR_STR(dest:u8, idx:u16)`, which performs `str_regs[dest] = vars.str_pool[idx]`, and `LOAD_VAR_INT` instructions, to load vars encoded in the `vm.Vars` struct + +There is no allotment instruction either: an allotment share is computed with `INT_TO_PORTION`, `MUL_PORTION` and `PORTION_TO_INT`, plus a fixup that hands the rounding leftover to the earliest shares. [instruction-encoding.md](instruction-encoding.md) has the full opcode table. + +The VM implementation itself is trivial, and most of the complexity is moved to the compiler + +### Bytecode encoding + +The program bytecode encoding is designed so that the hydration can be fast, and so that it stays stable across versions. + +After a magic word (so that we reject right away random bytes that didn't come from the compiler), we have a small header with the bytecode version and the number of sections. The version lets a decoder reject a payload it cannot read instead of silently misreading it. + +``` +| "NUMB" 4 B | magic ++-------------------------------+ +| major : u16 2 B | header +| minor : u16 2 B | +| count : u16 2 B | ++-------------------------------+ +| section 0 | sections +| section 1 | +| ... | +| section `count`-1 | ++-------------------------------+ + +section ++-----------+-----------+---------------+ +| tag : u16 | len : u32 | content ... | len = content byte length ++-----------+-----------+---------------+ + 2 B 4 B +``` + +Each section id identified by its `tag` identifier. `len` is the number of bytes it takes. + +Current sections are + +- Instructions, hydrated into a `[]vm.Instruction` slice +- Str pool, hydrated into a `[]string` slice +- Int pool, hydrated into a `[]big.Int` slice + +Missing sections are valid and considered as empty. Unkown sections are allowed and skipped, unless the 15-th bit(`0x8000`) is set; in that case the program is rejected. + +> Note: the compiler is free to arrange the sections in any order (e.g. it may pad or align them in future optimizations) + +> Note: this design would make it possible to have very fast hydration by re-intepreting the instruction slice via unsafe casting, or by using mmap. In our case this is more dangerous than useful, but it's a nice property to have + +#### Bytecode version + +The version is `major.minor` (`vm.BytecodeVersion`, exported as `numscript.BytecodeVersion`), two `u16` header fields with major first. It versions the bytecode format only: the compiler and the library are released independently of it, and a new compiler version does not imply a new bytecode version. Sixteen bits each is deliberately more than a byte: the header width is the one part of the format that cannot be widened later without also changing the magic word, and minor is the field that moves on every additive change. + +The split encodes the compatibility rule. A reader accepts a payload of its own major whose minor is no newer than its own, and nothing else: + +- a **minor** bump is additive — new opcodes, new sections, new optional operands — and leaves the meaning of everything an older writer could produce untouched, so a 1.1 VM runs 1.0 bytecode; +- a 1.0 VM does **not** run 1.1 bytecode: it may happen to know every opcode a given payload uses, but that is not assumed; +- a **major** bump changes the meaning of existing encodings, so a 2.0 VM runs no 1.x bytecode at all. + +`Encode` always stamps `vm.CurrentBytecodeVersion`; the decoders return `vm.UnsupportedBytecodeVersionError` (carrying the encoded and the supported version) for a payload they cannot read. `PeekProgramVersion` / `PeekVarsVersion` read the version off the raw header without decoding the rest and without applying the rule, so a host can report which version an unreadable payload was written with. A host that stores bytecode compiled by one build and executes it with another should compare the stored payload's version with `CurrentBytecodeVersion` (or `CanRead`, if it accepts older minors) before trusting it to run. + +### Constant pools + +Both the str and int pool start with the count of elems and then have contiguos sequence of strings/ints. + +``` +str pool section int pool section ++-------------+--------------+ +-------------+-------------+ +| count : u32 | records ... | | count : u32 | records ... | ++-------------+--------------+ +-------------+-------------+ + +str record int record ++------------+-------------+ +------+-------------+----------------+ +| len : u32 | raw bytes | | sign | magSz : u32 | magnitude ... | ++------------+-------------+ +------+-------------+----------------+ + 1 B big-endian +``` + +Int follows the same encoding as its `.Bytes()` and `.SetBytes()` methods. + +### Vars encoding + +`vm.Vars` uses the exact same encoding, just without the instructions section: the magic word is `"NVAR"`, and it only carries the str pool and int pool sections. + +``` +| "NVAR" 4 B | magic ++-------------------------------+ +| major : u16 2 B | header +| minor : u16 2 B | +| count : u16 2 B | ++-------------------------------+ +| str pool section | +| int pool section | ++-------------------------------+ +``` + +The version/section framing and the pool encoding are the same code as the program's, just parametrized by the magic word and the set of sections. + +An important property is that the `vm.Vars` don't have a 1-1 correspondence with the vars. The `vm.Vars` only encodes ints and strings. Composite objects, such as monetaries, are split into 2 different vars. This keeps data encoding minimal, and makes optimizations surface simpler to implement (see optimisations section below). + +The same principle now applies inside the VM, not just at the vars boundary: a monetary **is** a (str asset, int amount) register pair everywhere. There is no monetary register bank and no instruction that builds or projects one. + +The compiler is free to choose any encoding it wants for the vars (e.g. the first value in the str pool doesn't have to be the first string variable). Behaviour can change across versions. + +### Soundness verification + +`vm.Verify(program)` and `vm.VerifyWithVars(program, vars)` (in `internal/vm/verify.go`) analyse the bytecode statically, so a bug in the compiler or a corrupted payload is caught before it can make the VM crash. See [bytecode-verifier.md](bytecode-verifier.md) for the full list of checks and what each one prevents. In short: + +- No undefined opcodes, and no truncated multi-word instruction +- Const indices stay inside the const pool; with `VerifyWithVars`, var indices stay inside the vars pool +- Register indices stay below the counts the program declares +- Jumps land on an instruction boundary. They can only go forward (the delta is unsigned), so every program halts +- No read before write, on every path. This is what makes reusing a VM instance safe +- Mark regions (`oneof`) open and close in matching pairs on every path, and no send, `save` or asset change runs inside one + +The check is opt-in: `Exec` and `compiler.Compile` don't call it, since the compiler's output is valid by construction (the test corpus runs every compiled script through it). Call it on any program that didn't just come out of the compiler, for example after a Raft node decodes the bytes payload. + +## Compiler + +Instead of emitting a `vm.Instruction{}` stream directly, the compiler emits a `[]ir.Instr` slice. That's an intermediate representation of the instruction which isn't strictly necessary, but allows us to dump, manipulate or analyse instruction without having to run a fully-fledged disassembler every time. After the compilation, the `[]ir.Instr` are assembled into `[]vm.Instruction`. +The instruction set is mostly similar, but there are a few differences. +The most crucial one is that instead of many separate pools of 256 registers, there's a single infinite stream of registers. +We'll materialise those "logical" registers into actual physical registers during assembly, and perfom register allocation policies so that we'll be able to fit scripts within the 256 registers constraint. +We are able to fully typecheck the `[]ir.Instr` program, so that we know that we aren't passing logical registers that were created with a different type. + +Other differences in the instruction set include: + +- instead of `LOAD_INT` or `LOAD_STRING` referencing constant pool index, we have a `loadInt{ dest reg; value big.Int }` and `loadString{ dest reg; value string}` which handle populating and deduping constant pool when assembling, or using specialised instructions like `LOAD_INT_IMMEDIATE` instructions which contain the number in the payload itself. +- we have a `labelMarker struct{ label string }` pseudo-instruction. This way the jump can reference an instruction that hasn't been emitted yet without complex hacks at compile time + +This split allows us to implement peephole optimisations (bytecode rewriting) - see the "optimisation" section. + +We'll use the `irInstr` notation in the following sections: + +``` +// instructions can have many args, which may have labels, +// and may write the result into another register +$my_reg = some_instr($arg_reg, label: $another_arg) + +// consts use literals directly: +$some_int = 42 +$some_str = "USD/2" + +// special syntax for int math: +$tot = $x + $y +// auto-increment syntax +$tot += $x +``` + +This notation is a real format: it has a grammar ([IR.g4](IR.g4)), a parser (`ir.Parse`) and a dumper (`ir.Dump`), and instructions round-trip through it. It's fully specified in [ir-textual-format.md](ir-textual-format.md) — including the instruction reference, the argument conventions and the known round-trip caveats. + +In the sections below, a meta-notation is used for parametrized exprs/sources/dests + +#### Bounded send statement + +```num +send ( + source = + destination = +) +``` + +``` +$asset, $amount = // two regs, no instruction of its own +set_current_asset($asset) // needed for pull_account and send_to_account + +// a source always compiles by putting the pulled amt into a reg +$pulled = + +// we check if we managed to pull enough funds, or fail due to missing funds +check_enough_funds($pulled, $amount) + + +``` + +#### Plain account source/dest (bounded) + +Let's compile the `@src` source account, bounded by the value in the `$amount` reg. +It'll write pulled amount into the `$pulled` reg: + +``` +$src = "src" +$overdraft = 0 +$eq = str_eq($src, $world) +jmp_if_false($eq, #not_world) + $pulled = pull_account(account: $src, cap: $amount) + jmp(#pull_end) +#not_world + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) +#pull_end +``` + +The branch is `@world`. A pull with no `overdraft` operand is *unbounded*: it makes +the whole cap available without ever reading a balance, which is exactly what +`@world` means. The VM knows nothing about the name — the account is a register, so +it could equally come from a var, an interpolation or a `meta()` read, and the +comparison has to happen at run time. `$world` is a single `load_str "world"` in +the program prologue, so it dominates every such branch no matter what jumps the +branch sits between. + +This is emitted for *every* source account, literal `@world` or not. One code path +is easier to trust than a compile-time-folded one, and collapsing it back down is a +peephole's job: const-fold `str_eq` over two known `load_str`s, then drop the dead +arm. Until those land, the diamond is the cost of the VM not knowing about `@world`. + +When the source is `allowing unbounded overdraft` there is no `overdraft` operand +to drop, so both arms would be identical and the branch is skipped entirely. + +the plain `@dest` destination account will look like: + +``` +$dest = "dest" +send_to_account(account: $dest) +``` + +Here's a full example of a send statement: + +``` +send [USD/2 10] ( + source = @src + destination = @dest +) +``` + +output: + +``` +$world = "world" +$asset = "USD/2" +$amount = 10 +set_current_asset($asset) +$src = "src" +$overdraft = 0 +$eq = str_eq($src, $world) +jmp_if_false($eq, #not_world) + $pulled = pull_account(account: $src, cap: $amount) + jmp(#pull_end) +#not_world + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) +#pull_end +check_enough_funds($pulled, $amount) +$dest = "dest" +send_to_account(account: $dest) +``` + +#### Inorder sources (bounded) + +Let's compile the inorder source `, .., `, by storing the pulled amt in the `$pulled` register, bounded by the `$amount` cap. + +``` +$pulled = 0 +$remaining = int_copy($amount) + +// first source +$pulled_s1 = +$pulled += $pulled_s1 +$remaining -= $pulled_s1 +$exhausted = is_zero($remaining) +jmp_if_true($exhausted, #inorder_end) + +// second source +$pulled_s2 = +$pulled += $pulled_s2 +$remaining -= $pulled_s2 +$exhausted = is_zero($remaining) +jmp_if_true($exhausted, #inorder_end) + +.. + +$pulled_sn = +$pulled += $pulled_sn +// last one doesn't need jump + +#inorder_end +``` + +#### Max source (bounded) + +Let's compile `max from `, bounded by the `$amount` cap. + +``` +$max_asset, $max_amount = +assert_same_asset($max_asset, $asset) // $asset is the current asset (set via set_current_asset) + +// $cap = min($max_amount, $amount): there is no min opcode, so it is a comparison +// and a copy. Copying the left operand first saves the else arm's `jmp`. +$cap = int_copy($max_amount) +$lt = lt_int($max_amount, $amount) +jmp_if_true($lt, #min_end) +$cap = int_copy($amount) +#min_end + +$pulled = +``` + +The speculative copy is only sound because `$cap` is freshly allocated: an aliased dest would clobber `$amount` before the else arm reads it. If there's no outer cap (an unbounded context), the whole min is skipped and the inner source is capped by `$max_amount` directly. + +#### Allotment source (bounded) + +Let's compile the allotment source ` from , .., from `, bounded by the `$amount` cap. +Note that an allotment source is always bounded. + +``` +// only for a that could be negative, i.e. a division with a non-literal operand +assert_non_negative_portion() + +// portions must cover exactly 1 (no `remaining` clause here) +$leftover = 1 - - .. - +assert_leftover_exact($leftover) // plain "assert_leftover" if there's a remaining clause + +// split $amount across the portions. There is no allotment instruction: each +// share is a floored product, and the leftover from flooring is handed to the +// earliest shares one unit at a time. n is statically known, so the fixup is +// unrolled -- and only n-1 blocks are needed, since the shortfall is < n. +$amount_p = int_to_portion($amount) +$share_1 = portion_to_int(mul_portion(, $amount_p)) +.. +$share_n = portion_to_int(mul_portion(, $amount_p)) +// then, for i in 1..n-1: if $total < $amount { $share_i += 1; $total += 1 } + +$pulled_s1 = +check_enough_funds($pulled_s1, $share_1) + +.. + +$pulled_sn = +check_enough_funds($pulled_sn, $share_n) +``` + +### Optimisations + +> [!NOTE] +> Peephole optimisations aren't implemented yet. There is a POC in another branch, to measure how much perf could be impacted, but it's too soon to consider + +You may have noticed that the previous compilation examples emit _a lot_ of garbage. +That's done by design: the compiler must be simple and declarative. We don't want dozens of special cases in the compilation logic, which must express a general, albeit redundant, template which focuses on correctness. + +One whole class of that garbage is gone for good, though, and not via a peephole: monetaries used to be boxed with `mk_monetary` and immediately unboxed with `get_asset`/`get_amount`, so 3 of the 12 instructions for a simple `send` were pure round-tripping. Representing a monetary as a register pair removes them at codegen time, which is why there is no `monetaryFold` peephole to write. + +That change also *enables* a peephole that was previously out of reach. `assert_same_asset` used to compare two `get_asset(mk_monetary(..))` results, whose provenance is invisible without folding first; now both operands are plain `load_str`s, so an asset comparison between two literals is statically decidable and the assert can be dropped outright. + +However, the `irInstr` layer allows us to easily rewrite the instructions so that we remove garbage instructions, precomputing more aggressively, rewrite them into more efficient code (this is called [peephole optimisation](https://en.wikipedia.org/wiki/Peephole_optimization)) + +Each peephole is independent and is expressed as a `func(instr []irInstr) []irInstr`, which returns the new instructions set, or nil if it didn't change. +Each peephole is independently testable and reviewable. + +We apply each peephole optimisation sequentially, and repeat until we reach a fixed point for each peephole (the program `p` such that `f(p) == p`, where `f` is the peephole function). + +Note that proving that a peephole function _does_ have a fixed point is usually simple, whereas proving that the function composition of all the peepholes isn't. Pratically speaking, we can avoid non-terminating optimisation passes by imposing a max amount of optimisation passes. A clever order of peepholes should make convergence quite fast anyway. + +> TODO list some peepholes + +The `@world` diamond (see [Plain account source/dest](#plain-account-sourcedest-bounded)) is +the clearest case to date, and it needs two passes that compose: + +1. **Const-fold `str_eq`** — when both operands trace back to `load_str` constants, the + comparison is statically decidable: replace it with `true` or `false`. +2. **Dead-branch elimination** — a conditional jump on a register holding a known bool + is either a no-op or an unconditional `jmp`; then everything between a `jmp` and the + next reachable label is unreachable. `is_zero` over a `load_int` folds the same way, + which is what makes the quantity branches reachable for this pass too. + +Together they collapse a literal `@world` source back to a single unbounded +`pull_account`, and any other literal account back to a single bounded one — i.e. to +exactly the code the compiler emitted when the VM still special-cased the name. The +prologue's `load_str "world"` also becomes dead once no branch refers to it. + +### Registers allocator + +> [!NOTE] +> Currently implemented allocator is a bump allocator: allocate a fresh register for each distinct logical register. A linear-scan allocator is prototyped in another branch. + +After optimisation pass is (optionally) run, we can materialise logical registers into physical registers of each type bank during assembly phase. + +A good register allocation algorithm can reduce the number of needed registers. +For example, consider the `($x + $y) * $z` expression: + +``` +$x = load_var(idx: 0) +$y = load_var(idx: 1) +$z = load_var(idx: 2) +$w = $x + $y +$res = $w * $z +``` + +A naive allocation (bump allocation: materialise each distinct logical register into a fresh physical register) would assemble this into: + +``` +// need 5 registers in total +LOAD_VAR_INT(dest: 0, idx: 0) +LOAD_VAR_INT(dest: 1, idx: 1) +LOAD_VAR_INT(dest: 2, idx: 2) +ADD_INT(dest: 3, left: 0, right: 1) +MUL_INT(dest: 4, left: 3, right: 2) +``` + +Whereas an optimal allocation would produce something like this: + +``` +// need 2 registers in total +LOAD_VAR_INT(dest: 0, idx: 0) +LOAD_VAR_INT(dest: 1, idx: 1) +ADD_INT(dest: 0, left: 0, right: 1) +LOAD_VAR_INT(dest: 1, idx: 2) +MUL_INT(dest: 0, left: 0, right: 1) +``` + +What a better register allocation buys us is: + +1. better CPU locality, thus higher runtime speed (probably irrelevant gain in our case) +2. less memory used: the initial vm load will have to load less registers (although max number of registers per bank is 256 anyway) +3. avoid having to forbid scripts that overflow the 256 registers limit, or having to implement register spilling behaviour (the most important improvement) + +Registers allocation is a widely studied topic, so we don't really have to discover anything new. +There are more aggressive and expensive algorithms that are able to produce the most optimal registers allocation (e.g. by having to compute graph coloring, a provably expensive problem), which we don't need in our case: we still need decent perf at compile time as well, and a simpler allocation will most likely be "good enough". +Specifically, a [linear scan allocation](https://web.cs.ucla.edu/~palsberg/course/cs132/linearscan.pdf) will get us very close to the optimal allocation with `O(n)` cost. + +> Note: Claude argues that, for our instruction set, linear scan would produce _exactly_ the same result as the optimal allocation algorithms. I haven't yet put effort in understanding whether that's the case and why that is diff --git a/instruction-encoding.md b/instruction-encoding.md new file mode 100644 index 00000000..3763a212 --- /dev/null +++ b/instruction-encoding.md @@ -0,0 +1,521 @@ +# VM Bytecode Specification + +Instructions are **4 bytes** wide: `[Opcode: 8] [A: 8] [B: 8] [C: 8]`. + +- Registers are split into **per-type banks** (`int_regs`, `str_regs`, `por_regs`, `bool_regs`); an operand indexes the bank implied by the opcode. There is no monetary bank: a monetary is a (`str_regs` asset, `int_regs` amount) pair, so the instructions that deal in monetaries take or return the two halves separately. +- `0xFF` in a register slot means **nil** (absent optional operand). +- **`Bx`** = a `u16` formed by slots `B`,`C` (little-endian); used for pool indices and jump targets. **`sBx`** is its signed form. +- Most instructions are one word. A few extend into **continuation words** (shown as `↳ cont.`); an instruction's length is fixed by its opcode. +- There is **no `HALT`**: programs terminate by design (jumps are forward-only). +- Opcodes are grouped by category with gaps, so new instructions slot into a category without renumbering. Unused values are reserved (users can't emit them, so we stay free to define them later). + +> Opcode values match the constants in `internal/vm/instruction.go`; operand layouts match `decodeInstr` in `internal/vm/verify.go`. + +--- + +## 1. State & Assertions + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
OpcodeHexNameABCDescription
00x00SET_CURRENT_ASSETasset--Sets the current asset (used by PULL_ACCOUNT / SEND_TO_ACCOUNT) from str_regs[A]
10x01ASSERT_SAME_ASSETxy-Traps unless str_regs[A] and str_regs[B] are the same asset
20x02ASSERT_VALID_ACCOUNTacc--Traps if the account name in str_regs[A] is malformed
30x03ASSERT_NON_NEGATIVE_BALANCEamtacc-Traps if int_regs[A] is negative; B = account (for the error)
40x04ASSERT_LEFTOVERporexact-Traps if por_regs[A] is negative; when B == 1 (no remaining) also traps if non-zero
50x05CHECK_ENOUGH_FUNDSpulledtarget-Traps (missing funds) unless int_regs[A] == int_regs[B]: the pulled amount must match the target exactly
60x06ASSERT_VALID_COLORcolor--Traps if the color in str_regs[A] is malformed (only uppercase letters; the empty string is valid)
70x07ASSERT_NON_NEGATIVE_AMOUNTamt--Traps if int_regs[A] (a sent/saved amount) is negative
80x08ASSERT_VALID_SCOPEscope--Traps if the scope in str_regs[A] is malformed
90x09ASSERT_NON_NEGATIVE_PORTIONpor--Traps if por_regs[A] (an allotment clause portion) is negative
100x0AASSERT_UNSCOPEDscopeacc-Traps (cannot cast a scoped account to string) if str_regs[A] is not empty; B = account (for the error). Emitted before a scoped account is interpolated into an account name
0x0B..0x0F reserved
+ +## 2. Constants & Variables + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
OpcodeHexNameABCDescription
160x10LOAD_INTdestBx (const idx)int_regs[A] = int_pool[Bx]
170x11LOAD_STRdestBx (const idx)str_regs[A] = str_pool[Bx]
180x12LOAD_VAR_INTdestBx (var idx)int_regs[A] = vars.int_pool[Bx]
190x13LOAD_VAR_STRdestBx (var idx)str_regs[A] = vars.str_pool[Bx]
200x14LOAD_INT_IMMEDIATEdestsBx (i16 value)int_regs[A] = (big.Int)sBx — small literals inline, no pool entry. Reserved; not implemented
210x15CONST_TRUEdest--bool_regs[A] = true — the value is in the opcode, so there is nothing to decode and no pool entry
220x16CONST_FALSEdest--bool_regs[A] = false
0x17..0x1F reserved
+ +## 3. Metadata + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
OpcodeHexNameABCDescription
320x20SET_TX_METAkeyval-Sets transaction metadata str_regs[A] = str_regs[B]
330x21SET_ACCOUNT_METAacckeyvalSets account metadata: account A, key B, value C. 2 words:
​​↳ cont.scope--Scope reg (0xFF = unscoped)
340x22META_STRdestacckeystr_regs[A] = meta(account B, key C). 2 words:
​​↳ cont.scope--Scope reg (0xFF = unscoped)
350x23META_INTdestacckeyas META_STR, typed int, same continuation word
360x24META_PORTIONdestacckeyas META_STR, typed portion, same continuation word
370x25META_MONETARYdest assetacckeyParses the value as a monetary. One store read yields both halves, so this is the only two-destination read: str_regs[A] = asset
​​↳ cont.dest amtscope-int_regs[A] = amount; scope reg in B (0xFF = unscoped)
0x26..0x2F reserved
+ +## 4. Arithmetic & Constructors (binary) + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
OpcodeHexNameABCDescription
480x30ADD_INTdestleftrightint_regs[A] = int_regs[B] + int_regs[C]
490x31SUB_INTdestleftrightint_regs[A] = int_regs[B] - int_regs[C]
0x32 reserved (was MIN_INT: a min is a comparison and a copy, so it is LT_INT plus a branch)
510x33SUB_PORTIONdestleftrightpor_regs[A] = por_regs[B] - por_regs[C]. Its counterpart ADD_PORTION is at 0x38, not adjacent, because 0x32 is burned and 0x34..0x37 were taken
520x34MK_PORTIONdestnumdenpor_regs[A] = int_regs[B] / int_regs[C]
0x35 reserved (was MK_MONETARY: a monetary is a register pair, nothing to construct)
540x36ADD_STRINGdestleftrightstr_regs[A] = str_regs[B] + str_regs[C]
0x37 reserved (was STR_EQ: moved to the comparison group, §7, now 0x62)
560x38ADD_PORTIONdestleftrightpor_regs[A] = por_regs[B] + por_regs[C]. A rational sum, so unequal denominators combine correctly and the result is normalised; it may exceed 1
570x39MUL_PORTIONdestleftrightpor_regs[A] = por_regs[B] * por_regs[C]. With PORTION_TO_INT, this is how an allotment share is computed
0x3A..0x3F reserved
+ +## 5. Unary & Conversions + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
OpcodeHexNameABCDescription
0x40..0x41 reserved (were GET_AMOUNT / GET_ASSET: projecting a monetary is naming one of its two registers, so it costs no instruction)
660x42INT_COPYdestsrc-int_regs[A] = int_regs[B] (fresh copy). One copy per bank, none crossing banks; the family is split across 0x42..0x43 and 0x4A..0x4B because 0x44..0x49 were already spoken for. No monetary copy: a monetary is a (str, int) pair, so copy the halves
670x43PORTION_COPYdestsrc-por_regs[A] = por_regs[B] (fresh copy)
680x44NEG_INTdestsrc-int_regs[A] = -int_regs[B]
690x45INT_TO_STRINGdestsrc-str_regs[A] = str(int_regs[B])
700x46PORTION_TO_STRINGdestsrc-str_regs[A] = str(por_regs[B])
710x47MONETARY_TO_STRINGdestassetamtstr_regs[A] = str_regs[B] + " " + str(int_regs[C]) — takes both halves, so it is ternary despite living in this section
0x48 reserved (was IS_ZERO: moved to the comparison group, §7, now 0x63)
0x49 reserved (was NOT: moved to the bool-ops group, §8, now 0x70)
740x4ASTR_COPYdestsrc-str_regs[A] = str_regs[B]
750x4BBOOL_COPYdestsrc-bool_regs[A] = bool_regs[B]
760x4CINT_TO_PORTIONdestsrc-por_regs[A] = int_regs[B], exact
770x4DPORTION_TO_INTdestsrc-int_regs[A] = floor(por_regs[B])
0x4E..0x4F reserved
+ +## 6. Funds & Postings + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
OpcodeHexNameABCDescription
800x50PULL_ACCOUNTdestacccapPulls the current asset from account str_regs[B] into the source queue, capped by int_regs[C] (0xFF = uncapped); pulled amount → int_regs[A]. 2 words:
​​↳ cont.overdraftcolorscopeOverdraft bound int_regs[A] (0xFF = unbounded), color str_regs[B] and scope str_regs[C] (0xFF = none). Cap and overdraft both 0xFF traps (unbounded source)
810x51SEND_TO_ACCOUNTacccapscopeDrains the source queue in the current asset to account str_regs[A], up to int_regs[B] (0xFF = everything queued), scope str_regs[C] (0xFF = unscoped). A = 0xFF is kept: the funds go back to their sources and no posting is emitted. Traps while a mark is open
820x52SAVEaccassetamountReduces the balance of account str_regs[A] for asset str_regs[B] by int_regs[C] (C = 0xFF ⇒ save all), floored at 0. Traps while a mark is open. 2 words:
​​↳ cont.scope--Scope reg (0xFF = unscoped)
0x53 reserved (was MK_ALLOTMENT: an allotment share is built from MUL_PORTION, PORTION_TO_INT and the leftover fixup, so there is no variadic instruction)
840x54BALANCEdest amtaccassetint_regs[A] = balance(account B, asset C) from the run-state. Only the amount: the resulting monetary's asset is operand C, which the caller already holds. 2 words:
​​↳ cont.scope--Scope reg (0xFF = unscoped)
850x55MARK_PUSH---Opens a region at the current source-queue depth and posting count, for oneof backtracking. The run-state keeps the stack of open regions; no register is involved
860x56MARK_ENDrewind--Closes the innermost region. A = 1 rewinds it: repays everything pulled and reverses everything posted since the matching MARK_PUSH. A = 0 commits it. Any other value is rejected by the verifier; traps if no region is open. Textual IR: mark_rewind() / mark_commit()
0x57..0x5F reserved (e.g. PULL_ACCOUNT specializations). This block used to run to 0x8F; §7 and §8 took 0x60..0x7F out of it, leaving nine slots for the four specializations sketched in instruction.go
+ +## 7. Comparisons + +Every bool producer lives here. `A` = dest (a `bool_regs` index) for all of them; the operand banks are what the opcode implies. `IS_ZERO` is unary and the rest are binary — they are one group because they are one *category*, not one arity. + +Only `<` and `==` exist, per type. The other four surface operators are **normalised by the front end**: + +| surface | lowering | +|---|---| +| `a < b` | `Lt(a, b)` | +| `a > b` | `Lt(b, a)` — operands swapped | +| `a <= b` | `Not(Lt(b, a))` | +| `a >= b` | `Not(Lt(a, b))` | +| `a == b` | `Eq(a, b)` | +| `a != b` | `Not(Eq(a, b))` | + +12 surface operators, 5 opcodes. Every extra predicate is another case in the SMT encoder and in any formal model of the VM, so the cost would be paid three times over. LLVM does the same, canonicalising `sgt` to `slt` with swapped operands in InstCombine so downstream passes only ever see one form. + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
OpcodeHexNameABCDescription
960x60LT_INTdestleftrightbool_regs[A] = int_regs[B] < int_regs[C]. Strict
970x61EQ_INTdestleftrightbool_regs[A] = int_regs[B] == int_regs[C]
980x62STR_EQdestleftrightbool_regs[A] = str_regs[B] == str_regs[C]. The only string comparison that yields a value rather than trapping (cf. ASSERT_SAME_ASSET). Was 0x37
990x63IS_ZEROdestsrc-bool_regs[A] = int_regs[B].Sign() == 0 — the projection from a quantity to a condition, since the jumps take a bool. Tests the sign, so a negative amount is not zero. Kept alongside EQ_INT because it needs no materialised zero and it is on every quantity branch. Was 0x48
1000x64LT_PORTIONdestleftrightbool_regs[A] = por_regs[B] < por_regs[C]. Strict, and by value — see EQ_PORTION
1010x65EQ_PORTIONdestleftrightbool_regs[A] = por_regs[B] == por_regs[C]. Value equality: 1/2 == 2/4 is true. big.Rat normalises on construction, so the rationals are compared — comparing numerator/denominator pairs separately would give the wrong answer
0x66..0x6F reserved for < and == on types that don't exist yet. Str gets equality only, never ordering. Bool equality, and structural comparison of tuples/arrays, are front-end expansions rather than opcodes. No named-but-unimplemented constants live here on purpose: a live opcode with no emitter invites a second lowering path that no test exercises
+ +## 8. Bool ops + + + + + + + + + + + + + + + + + + + + + + +
OpcodeHexNameABCDescription
1120x70NOTdestsrc-bool_regs[A] = !bool_regs[B] — the only operation whose operand and result are both bools, and what the four derived operators above are built from. Was 0x49
0x71..0x7F reserved for and/or, if they ever pay for themselves — both are expressible as branches, so neither is needed for completeness
0x80..0x8F reserved
+ +## 9. Control Flow + + + + + + + + + + + + + + + + + + + + + + + + + + + + + +
OpcodeHexNameABCDescription
1440x90JMP_IF_FALSEcondBx (forward delta)If bool_regs[A] is false, skip Bx instructions: pc += Bx, where pc already points at the next instruction. Being an unsigned delta, the jump is forward-only (guarantees termination). A quantity is not a condition — project it with IS_ZERO first
1450x91JMP—Bx (forward delta)Unconditional: pc += Bx. Forward-only, as above
1460x92JMP_IF_TRUEcondBx (forward delta)The dual of JMP_IF_FALSE, so either edge of a condition is one instruction and no negation opcode is needed
0x93..0xFF reserved
diff --git a/internal/analysis/check.go b/internal/analysis/check.go index d2ec4243..428111e7 100644 --- a/internal/analysis/check.go +++ b/internal/analysis/check.go @@ -5,6 +5,7 @@ import ( "slices" "strings" + "github.com/formancehq/numscript/internal/builtins" "github.com/formancehq/numscript/internal/flags" "github.com/formancehq/numscript/internal/parser" "github.com/formancehq/numscript/internal/utils" @@ -55,17 +56,17 @@ func (r VarOriginFnCallResolution) GetParams() []string { return r.Params } func (r StatementFnCallResolution) GetParams() []string { return r.Params } const ( - // Statemetn fns - FnSetTxMeta = "set_tx_meta" - FnSetAccountMeta = "set_account_meta" + // Statement fns + FnSetTxMeta = builtins.SetTxMeta + FnSetAccountMeta = builtins.SetAccountMeta // Expr fns - FnVarOriginMeta = "meta" - FnVarOriginBalance = "balance" - FnVarOriginOverdraft = "overdraft" - FnVarOriginGetAsset = "get_asset" - FnVarOriginGetAmount = "get_amount" - FnVarOriginScoped = "scoped" + FnVarOriginMeta = builtins.Meta + FnVarOriginBalance = builtins.Balance + FnVarOriginOverdraft = builtins.Overdraft + FnVarOriginGetAsset = builtins.GetAsset + FnVarOriginGetAmount = builtins.GetAmount + FnVarOriginScoped = builtins.Scoped ) var Builtins = map[string]FnCallResolution{ diff --git a/internal/builtins/builtins.go b/internal/builtins/builtins.go new file mode 100644 index 00000000..d9bec8eb --- /dev/null +++ b/internal/builtins/builtins.go @@ -0,0 +1,12 @@ +package builtins + +const ( + SetTxMeta = "set_tx_meta" + SetAccountMeta = "set_account_meta" + Meta = "meta" + Balance = "balance" + Overdraft = "overdraft" + GetAsset = "get_asset" + GetAmount = "get_amount" + Scoped = "scoped" +) diff --git a/internal/cmd/assemble.go b/internal/cmd/assemble.go new file mode 100644 index 00000000..2832f1e6 --- /dev/null +++ b/internal/cmd/assemble.go @@ -0,0 +1,116 @@ +package cmd + +import ( + "fmt" + "io" + "os" + "strings" + + "github.com/formancehq/numscript/internal/ir" + + "github.com/spf13/cobra" +) + +type AssembleArgs struct { + OutputPath string +} + +// stdioPath is the conventional stand-in for stdin/stdout, accepted both as the +// input path and as --output. +const stdioPath = "-" + +// defaultBytecodePath derives the output path from the IR path: "x.ir" becomes +// "x.numb", anything else just gains the suffix. Reading from stdin has no path +// to derive from, so it writes to stdout. +func defaultBytecodePath(irPath string) string { + if irPath == stdioPath { + return stdioPath + } + return strings.TrimSuffix(irPath, ".ir") + ".numb" +} + +func readIRSource(irPath string) ([]byte, error) { + if irPath == stdioPath { + return io.ReadAll(os.Stdin) + } + return os.ReadFile(irPath) +} + +func assemble(irPath string, opts AssembleArgs) error { + content, err := readIRSource(irPath) + if err != nil { + return err + } + src := string(content) + + instrs, irErrs := ir.Parse(src) + if len(irErrs) != 0 { + for _, irErr := range irErrs { + fmt.Fprintln(os.Stderr, irErr.Error()) + fmt.Fprint(os.Stderr, irErr.Range.ShowOnSource(src)) + } + return fmt.Errorf("assembling failed") + } + + if err := ir.Typecheck(instrs); err != nil { + fmt.Fprintln(os.Stderr, err.Error()) + return fmt.Errorf("assembling failed") + } + + program, err := ir.Assemble(instrs) + if err != nil { + fmt.Fprintln(os.Stderr, err.Error()) + return fmt.Errorf("assembling failed") + } + + bytecode := program.Encode() + + outputPath := opts.OutputPath + if outputPath == "" { + outputPath = defaultBytecodePath(irPath) + } + if outputPath == stdioPath { + _, err := os.Stdout.Write(bytecode) + return err + } + + return os.WriteFile(outputPath, bytecode, 0o644) +} + +func getAssembleCmd() *cobra.Command { + opts := AssembleArgs{} + + cmd := cobra.Command{ + Use: "assemble ", + Short: "Assemble a textual IR file into bytecode", + Long: `Assemble a textual IR file into the binary bytecode the vm executes. + +The output goes to the input path with a ".numb" extension, for example: +assemble folder/my-script.ir +will write 'folder/my-script.numb'. + +Use --output to write elsewhere, or --output - to write the bytecode to stdout. + +Pass - as the path to read the IR from stdin, in which case the bytecode goes to +stdout unless --output says otherwise: +cat folder/my-script.ir | numscript assemble - > folder/my-script.numb + +The IR format tracks an unstable instruction set and is not a public interface. +`, + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + err := assemble(args[0], opts) + if err != nil { + cmd.SilenceErrors = true + cmd.SilenceUsage = true + return err + } + + return nil + }, + } + + cmd.Flags().StringVarP(&opts.OutputPath, "output", "o", "", "Path where to write the bytecode ('-' for stdout)") + + return &cmd +} diff --git a/internal/cmd/bytecode_run.go b/internal/cmd/bytecode_run.go new file mode 100644 index 00000000..56a6d1bc --- /dev/null +++ b/internal/cmd/bytecode_run.go @@ -0,0 +1,277 @@ +package cmd + +import ( + "context" + "encoding/json" + "fmt" + "math/big" + "os" + "sort" + "strings" + + "github.com/formancehq/numscript/internal/funds" + "github.com/formancehq/numscript/internal/interpreter" + "github.com/formancehq/numscript/internal/vm" + + "github.com/spf13/cobra" +) + +// VarsPoolFile is the raw form of vm.Vars: the compiler's VarsEncoder maps +// declared variable names onto pool slots, but it lives in the source, not in +// the bytecode, so a bytecode-only run has to name the slots by index. Ints are +// strings so that values beyond float64/int64 survive the JSON round-trip. +type VarsPoolFile struct { + Strings []string `json:"strings"` + Ints []string `json:"ints"` +} + +// BytecodeInputsFile is the `run` inputs file plus the vars pool. It is a +// separate type so that the vm's index-addressed vars stay out of the +// interpreter's inputs shape; the shared fields keep the same json names, so one +// .inputs.json works for both commands. +type BytecodeInputsFile struct { + Meta interpreter.AccountsMetadata `json:"metadata"` + Balances interpreter.Balances `json:"balances"` + VarsPool *VarsPoolFile `json:"varsPool"` +} + +type BytecodeRunArgs struct { + InputsPath string + VarsPath string + OutFormatOpt string + SkipVerify bool +} + +// vmMetaKey identifies one metadata slot: account, scope and key. +type vmMetaKey struct { + account string + scope string + key string +} + +// vmStore is a vm.Store over the rows of an inputs file. +type vmStore struct { + balances map[funds.PairKey]*big.Int + meta map[vmMetaKey]string +} + +func (s vmStore) GetBalance(_ context.Context, account, scope, asset, color string) (*big.Int, error) { + // the caller owns what it gets: the run state mutates balances in place + if v, ok := s.balances[funds.PairKey{Account: account, Scope: scope, Asset: asset, Color: color}]; ok { + return new(big.Int).Set(v), nil + } + return new(big.Int), nil +} + +func (s vmStore) GetMetadata(_ context.Context, account, scope, key string) (string, bool, error) { + v, ok := s.meta[vmMetaKey{account: account, scope: scope, key: key}] + return v, ok, nil +} + +func newVmStore(inputsPath string, inputs BytecodeInputsFile) (vmStore, error) { + store := vmStore{ + balances: make(map[funds.PairKey]*big.Int, len(inputs.Balances)), + meta: make(map[vmMetaKey]string, len(inputs.Meta)), + } + + for _, row := range inputs.Balances { + amount := row.Amount + if amount == nil { + amount = new(big.Int) + } + store.balances[funds.PairKey{Account: row.Account, Scope: row.Scope, Asset: row.Asset, Color: row.Color}] = amount + } + + for _, row := range inputs.Meta { + store.meta[vmMetaKey{account: row.Account, scope: row.Scope, key: row.Key}] = row.Value + } + + return store, nil +} + +// loadVars resolves the vars pool from either the inputs file or an encoded +// .nvar blob. Both absent is legal: vm.Exec accepts a nil *Vars. +func loadVars(inputsPath string, inputs BytecodeInputsFile, varsPath string) (*vm.Vars, error) { + if inputs.VarsPool != nil && varsPath != "" { + return nil, fmt.Errorf("cannot use --vars together with the 'varsPool' key of '%s'", inputsPath) + } + + if varsPath != "" { + content, err := os.ReadFile(varsPath) + if err != nil { + return nil, err + } + vars, err := vm.DecodeVars(content) + if err != nil { + return nil, fmt.Errorf("failed to decode vars file '%s': %w", varsPath, err) + } + return &vars, nil + } + + if inputs.VarsPool == nil { + return nil, nil + } + + ints := make([]big.Int, len(inputs.VarsPool.Ints)) + for i, raw := range inputs.VarsPool.Ints { + if _, ok := ints[i].SetString(raw, 10); !ok { + return nil, fmt.Errorf("invalid inputs file '%s': varsPool.ints[%d] is not an integer: %q", inputsPath, i, raw) + } + } + + return &vm.Vars{ + StringsPool: inputs.VarsPool.Strings, + IntsPool: ints, + Version: vm.CurrentBytecodeVersion, + }, nil +} + +func bytecodeRun(bytecodePath string, opts BytecodeRunArgs) error { + bytecode, err := os.ReadFile(bytecodePath) + if err != nil { + return err + } + + program, err := vm.DecodeProgram(bytecode) + if err != nil { + return fmt.Errorf("failed to decode bytecode file '%s': %w", bytecodePath, err) + } + + inputsPath := opts.InputsPath + if inputsPath == "" { + inputsPath = bytecodePath + ".inputs.json" + } + + inputsContent, err := os.ReadFile(inputsPath) + if err != nil { + return err + } + + var inputs BytecodeInputsFile + err = json.Unmarshal(inputsContent, &inputs) + if err != nil { + return fmt.Errorf("failed to parse inputs file '%s' as JSON: %w", inputsPath, err) + } + + if err := validateInputRows(inputsPath, inputs.Balances, inputs.Meta); err != nil { + return err + } + + store, err := newVmStore(inputsPath, inputs) + if err != nil { + return err + } + + vars, err := loadVars(inputsPath, inputs, opts.VarsPath) + if err != nil { + return err + } + + // this command is the one place bytecode arrives from outside the process, + // so it is the one place that has to assume nothing about it: Exec would + // crash rather than error on a malformed program + if !opts.SkipVerify { + if _, err := vm.VerifyWithVars(program, vars); err != nil { + return fmt.Errorf("bytecode file '%s' is malformed: %w", bytecodePath, err) + } + } + + result, execErr := vm.Exec(context.Background(), vm.NewVm(program), vars, store) + if execErr != nil { + fmt.Fprintln(os.Stderr, execErr.Error()) + return fmt.Errorf("execution failed") + } + + switch opts.OutFormatOpt { + case OutputFormatJson: + return showBytecodeJson(result) + case OutputFormatPretty: + return showBytecodePretty(result) + default: + return fmt.Errorf("invalid output format: %s", opts.OutFormatOpt) + } +} + +func showBytecodeJson(result funds.ExecutionResult) error { + out, err := json.Marshal(result) + if err != nil { + return fmt.Errorf("error marshaling result to JSON: %w", err) + } + + _, err = os.Stdout.Write(out) + return err +} + +func showBytecodePretty(result funds.ExecutionResult) error { + fmt.Println("Postings:") + fmt.Println(interpreter.PrettyPrintPostings(result.Postings)) + + // interpreter.PrettyPrintMeta takes map[string]Value; the vm's metadata is + // already stringified + if len(result.Metadata) != 0 { + fmt.Println("Meta:") + fmt.Print(prettyPrintStringMeta(result.Metadata)) + } + + if len(result.AccountsMetadata) != 0 { + fmt.Println("Accounts meta:") + fmt.Print(result.AccountsMetadata.PrettyPrint()) + } + + return nil +} + +func prettyPrintStringMeta(meta map[string]string) string { + keys := make([]string, 0, len(meta)) + for key := range meta { + keys = append(keys, key) + } + sort.Strings(keys) + + var sb strings.Builder + for _, key := range keys { + fmt.Fprintf(&sb, " %s: %s\n", key, meta[key]) + } + return sb.String() +} + +func getBytecodeRunCmd() *cobra.Command { + opts := BytecodeRunArgs{} + + cmd := cobra.Command{ + Use: "bytecode-run", + Short: "Execute a bytecode file", + Long: `Execute a bytecode file, taking as inputs a json file containing balances, metadata and the vars pool. + +The inputs file has to have the same name as the bytecode file plus a ".inputs.json" suffix, for example: +bytecode-run folder/my-script.numb +will expect a 'folder/my-script.numb.inputs.json' file where to read inputs from. + +You can explicitly specify where the inputs file should be using the optional --inputs argument. + +Unlike 'run', variables are not passed by name: the bytecode addresses them by +their index in the vars pools, so they are given either as a "varsPool" key of the +inputs file ({"strings": [...], "ints": [...]}) or as an encoded vars blob via --vars. + +The bytecode format tracks an unstable instruction set and is not a public interface. +`, + Args: cobra.ExactArgs(1), + RunE: func(cmd *cobra.Command, args []string) error { + err := bytecodeRun(args[0], opts) + if err != nil { + cmd.SilenceErrors = true + cmd.SilenceUsage = true + return err + } + + return nil + }, + } + + cmd.Flags().StringVar(&opts.InputsPath, "inputs", "", "Path of a json file containing the inputs") + cmd.Flags().StringVar(&opts.VarsPath, "vars", "", "Path of a file containing an encoded vars payload") + cmd.Flags().StringVarP(&opts.OutFormatOpt, "output-format", "o", OutputFormatPretty, "Set the output format. Available options: pretty, json.") + cmd.Flags().BoolVar(&opts.SkipVerify, "skip-verify", false, "Skip the static check of the bytecode. Only safe for a file this toolchain just produced.") + + return &cmd +} diff --git a/internal/cmd/bytecode_run_test.go b/internal/cmd/bytecode_run_test.go new file mode 100644 index 00000000..c09481a0 --- /dev/null +++ b/internal/cmd/bytecode_run_test.go @@ -0,0 +1,292 @@ +package cmd + +import ( + "context" + "math/big" + "os" + "path/filepath" + "testing" + + "github.com/formancehq/numscript/internal/funds" + "github.com/formancehq/numscript/internal/interpreter" + "github.com/formancehq/numscript/internal/ir" + "github.com/formancehq/numscript/internal/vm" + "github.com/stretchr/testify/require" +) + +const varsIR = ` + $asset = load_var(0) + set_current_asset($asset) + $amount = load_var(0) + $src = load_var(1) + $overdraft = load_var(1) + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) + $dest = load_var(2) + send_to_account(account: $dest) +` + +func TestDefaultBytecodePath(t *testing.T) { + require.Equal(t, "folder/x.numb", defaultBytecodePath("folder/x.ir")) + require.Equal(t, "folder/x.numb", defaultBytecodePath("folder/x")) + require.Equal(t, "folder/x.num.numb", defaultBytecodePath("folder/x.num")) + // stdin has no path to derive an output name from + require.Equal(t, "-", defaultBytecodePath("-")) +} + +// Reading the IR from stdin must assemble to the same program as reading it +// from a file, and default to writing the bytecode to stdout. +func TestAssembleFromStdin(t *testing.T) { + dir := t.TempDir() + + stdin, err := os.Create(filepath.Join(dir, "stdin")) + require.NoError(t, err) + _, err = stdin.WriteString(varsIR) + require.NoError(t, err) + require.NoError(t, stdin.Close()) + stdin, err = os.Open(filepath.Join(dir, "stdin")) + require.NoError(t, err) + + captured, err := os.Create(filepath.Join(dir, "stdout")) + require.NoError(t, err) + + realStdin, realStdout := os.Stdin, os.Stdout + os.Stdin, os.Stdout = stdin, captured + err = assemble("-", AssembleArgs{}) + os.Stdin, os.Stdout = realStdin, realStdout + require.NoError(t, stdin.Close()) + require.NoError(t, captured.Close()) + require.NoError(t, err) + + written, err := os.ReadFile(filepath.Join(dir, "stdout")) + require.NoError(t, err) + fromStdin, err := vm.DecodeProgram(written) + require.NoError(t, err) + + instrs, irErrs := ir.Parse(varsIR) + require.Empty(t, irErrs) + expected, err := ir.Assemble(instrs) + require.NoError(t, err) + require.Equal(t, expected.Instructions, fromStdin.Instructions) + + // nothing was written next to a "-" path + _, err = os.Stat("-.numb") + require.True(t, os.IsNotExist(err)) +} + +// The file assemble writes must decode back to exactly what the assembler +// produced in memory. +func TestAssembleWritesADecodableProgram(t *testing.T) { + dir := t.TempDir() + irPath := filepath.Join(dir, "prog.ir") + require.NoError(t, os.WriteFile(irPath, []byte(varsIR), 0o644)) + + require.NoError(t, assemble(irPath, AssembleArgs{})) + + written, err := os.ReadFile(filepath.Join(dir, "prog.numb")) + require.NoError(t, err) + decoded, err := vm.DecodeProgram(written) + require.NoError(t, err) + + instrs, irErrs := ir.Parse(varsIR) + require.Empty(t, irErrs) + require.NoError(t, ir.Typecheck(instrs)) + expected, err := ir.Assemble(instrs) + require.NoError(t, err) + + require.Equal(t, expected.Instructions, decoded.Instructions) + require.Equal(t, expected.MaxRegString, decoded.MaxRegString) + require.Equal(t, expected.MaxRegInt, decoded.MaxRegInt) + require.Equal(t, expected.MaxRegPortion, decoded.MaxRegPortion) + require.Equal(t, expected.MaxRegBool, decoded.MaxRegBool) + // compared by content, not with require.Equal on the whole Program: this + // program has no constants, and an empty pool assembles to a nil slice but + // decodes to an empty one (parseStringsPool's make([]T, 0)) + require.Empty(t, decoded.StringsPool) + require.Empty(t, decoded.IntsPool) +} + +func TestAssembleToStdoutLeavesNoFile(t *testing.T) { + dir := t.TempDir() + irPath := filepath.Join(dir, "prog.ir") + require.NoError(t, os.WriteFile(irPath, []byte(varsIR), 0o644)) + + // the bytecode is binary: capture it instead of letting it into the test log + captured, err := os.Create(filepath.Join(dir, "stdout")) + require.NoError(t, err) + realStdout := os.Stdout + os.Stdout = captured + err = assemble(irPath, AssembleArgs{OutputPath: "-"}) + os.Stdout = realStdout + require.NoError(t, captured.Close()) + require.NoError(t, err) + + _, err = os.Stat(filepath.Join(dir, "prog.numb")) + require.True(t, os.IsNotExist(err)) + + written, err := os.ReadFile(filepath.Join(dir, "stdout")) + require.NoError(t, err) + _, err = vm.DecodeProgram(written) + require.NoError(t, err) +} + +func TestLoadVarsFromPool(t *testing.T) { + vars, err := loadVars("in.json", BytecodeInputsFile{ + VarsPool: &VarsPoolFile{ + Strings: []string{"USD/2"}, + // wider than an int64, so a naive json number would have lost it + Ints: []string{"123456789012345678901234567890"}, + }, + }, "") + require.NoError(t, err) + + expected, _ := new(big.Int).SetString("123456789012345678901234567890", 10) + require.Equal(t, []string{"USD/2"}, vars.StringsPool) + require.Equal(t, []big.Int{*expected}, vars.IntsPool) +} + +func TestLoadVarsAbsentIsNil(t *testing.T) { + vars, err := loadVars("in.json", BytecodeInputsFile{}, "") + require.NoError(t, err) + require.Nil(t, vars) +} + +func TestLoadVarsRejectsBothSources(t *testing.T) { + _, err := loadVars("in.json", BytecodeInputsFile{VarsPool: &VarsPoolFile{}}, "vars.nvar") + require.ErrorContains(t, err, "cannot use --vars together with") +} + +// The --vars path is the leader/node wire format: an encoded vm.Vars blob has to +// decode and drive the program to the same result as the inline pool. +func TestLoadVarsFromEncodedFile(t *testing.T) { + dir := t.TempDir() + varsPath := filepath.Join(dir, "prog.nvar") + encoded := vm.Vars{ + StringsPool: []string{"USD/2", "src", "dest"}, + IntsPool: []big.Int{*big.NewInt(10), *big.NewInt(0)}, + }.Encode() + require.NoError(t, os.WriteFile(varsPath, encoded, 0o644)) + + vars, err := loadVars("in.json", BytecodeInputsFile{}, varsPath) + require.NoError(t, err) + + instrs, irErrs := ir.Parse(varsIR) + require.Empty(t, irErrs) + program, err := ir.Assemble(instrs) + require.NoError(t, err) + + store, err := newVmStore("in.json", BytecodeInputsFile{ + Balances: interpreter.Balances{ + {Account: "src", Asset: "USD/2", Amount: big.NewInt(100)}, + }, + }) + require.NoError(t, err) + + res, execErr := vm.Exec(context.Background(), vm.NewVm(program), vars, store) + require.Nil(t, execErr) + require.Len(t, res.Postings, 1) + require.Equal(t, funds.Posting{ + Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10), + }, res.Postings[0]) +} + +func TestLoadVarsRejectsAnUndecodableFile(t *testing.T) { + dir := t.TempDir() + varsPath := filepath.Join(dir, "prog.nvar") + require.NoError(t, os.WriteFile(varsPath, []byte("not a vars blob"), 0o644)) + + _, err := loadVars("in.json", BytecodeInputsFile{}, varsPath) + require.ErrorContains(t, err, "failed to decode vars file") +} + +// The store hands out balances the run state is free to mutate, so it must not +// alias the inputs. +func TestVmStoreReturnsACopyOfTheBalance(t *testing.T) { + amount := big.NewInt(100) + store, err := newVmStore("in.json", BytecodeInputsFile{ + Balances: interpreter.Balances{ + {Account: "src", Asset: "USD/2", Amount: amount}, + }, + }) + require.NoError(t, err) + + got, err := store.GetBalance(context.Background(), "src", "", "USD/2", "") + require.NoError(t, err) + require.Zero(t, got.Cmp(big.NewInt(100))) + + got.SetInt64(0) + require.Zero(t, amount.Cmp(big.NewInt(100))) +} + +func TestVmStoreUnknownAccountIsZeroNotAnError(t *testing.T) { + store, err := newVmStore("in.json", BytecodeInputsFile{}) + require.NoError(t, err) + + got, err := store.GetBalance(context.Background(), "nobody", "", "USD/2", "") + require.NoError(t, err) + require.Zero(t, got.Sign()) + + _, ok, err := store.GetMetadata(context.Background(), "nobody", "", "k") + require.NoError(t, err) + require.False(t, ok) +} + +func TestVmStoreSupportsScopedRows(t *testing.T) { + store, err := newVmStore("in.json", BytecodeInputsFile{ + Balances: interpreter.Balances{ + {Account: "src", Asset: "USD/2", Amount: big.NewInt(1), Scope: "reserve"}, + {Account: "src", Asset: "USD/2", Amount: big.NewInt(100)}, + }, + Meta: interpreter.AccountsMetadata{ + {Account: "src", Key: "k", Value: "scoped", Scope: "reserve"}, + {Account: "src", Key: "k", Value: "unscoped"}, + }, + }) + require.NoError(t, err) + + scopedBal, err := store.GetBalance(context.Background(), "src", "reserve", "USD/2", "") + require.NoError(t, err) + require.Zero(t, scopedBal.Cmp(big.NewInt(1))) + + unscopedBal, err := store.GetBalance(context.Background(), "src", "", "USD/2", "") + require.NoError(t, err) + require.Zero(t, unscopedBal.Cmp(big.NewInt(100))) + + scopedMeta, ok, err := store.GetMetadata(context.Background(), "src", "reserve", "k") + require.NoError(t, err) + require.True(t, ok) + require.Equal(t, "scoped", scopedMeta) + + unscopedMeta, ok, err := store.GetMetadata(context.Background(), "src", "", "k") + require.NoError(t, err) + require.True(t, ok) + require.Equal(t, "unscoped", unscopedMeta) +} + +// A .numb file is the one thing this toolchain reads that it did not produce +// itself, so bytecodeRun verifies before executing. Without that, a single +// flipped byte reaches Exec, which is entitled to assume it never sees one. +func TestBytecodeRunRejectsAMalformedFile(t *testing.T) { + dir := t.TempDir() + irPath := filepath.Join(dir, "prog.ir") + require.NoError(t, os.WriteFile(irPath, []byte(varsIR), 0o644)) + require.NoError(t, assemble(irPath, AssembleArgs{})) + + bytecodePath := filepath.Join(dir, "prog.numb") + bytecode, err := os.ReadFile(bytecodePath) + require.NoError(t, err) + + // point the last instruction's first register operand at a register the + // program never declared + program, err := vm.DecodeProgram(bytecode) + require.NoError(t, err) + require.NotEmpty(t, program.Instructions) + program.Instructions[len(program.Instructions)-1].A = 0xFE + require.NoError(t, os.WriteFile(bytecodePath, program.Encode(), 0o644)) + + inputsPath := filepath.Join(dir, "prog.numb.inputs.json") + require.NoError(t, os.WriteFile(inputsPath, []byte(`{"balances": []}`), 0o644)) + + err = bytecodeRun(bytecodePath, BytecodeRunArgs{OutFormatOpt: OutputFormatJson}) + require.ErrorContains(t, err, "is malformed") +} diff --git a/internal/cmd/root.go b/internal/cmd/root.go index 1a4d8e19..ebd8f3de 100644 --- a/internal/cmd/root.go +++ b/internal/cmd/root.go @@ -30,6 +30,15 @@ func Execute(options CliOptions) { rootCmd.AddCommand(getTestInitCmd()) rootCmd.AddCommand(getRunCmd()) + // The ir/bytecode tooling tracks an unstable instruction set, so it stays out + // of --help unless NUMSCRIPT_EXPERIMENTAL_CLI is set. It is always registered + // and runnable, like lsp and mcp. + hidden := os.Getenv("NUMSCRIPT_EXPERIMENTAL_CLI") == "" + for _, experimentalCmd := range []*cobra.Command{getAssembleCmd(), getBytecodeRunCmd()} { + experimentalCmd.Hidden = hidden + rootCmd.AddCommand(experimentalCmd) + } + if err := rootCmd.Execute(); err != nil { fmt.Println(err) os.Exit(1) diff --git a/internal/cmd/run.go b/internal/cmd/run.go index 8fef0907..37c4522d 100644 --- a/internal/cmd/run.go +++ b/internal/cmd/run.go @@ -29,6 +29,32 @@ type RunArgs struct { OutFormatOpt string } +// validateInputRows rejects a malformed inputs file before running anything: a +// balance list is a map keyed by (account, asset, color, scope) and a metadata +// list by (account, key, scope), so a repeated key is ambiguous. +func validateInputRows(inputsPath string, balances interpreter.Balances, meta interpreter.AccountsMetadata) error { + if dup, ok := balances.FirstDuplicate(); ok { + key := fmt.Sprintf("account=%q asset=%q", dup.Account, dup.Asset) + if dup.Color != "" { + key += fmt.Sprintf(" color=%q", dup.Color) + } + if dup.Scope != "" { + key += fmt.Sprintf(" scope=%q", dup.Scope) + } + return fmt.Errorf("invalid inputs file '%s': balances must not contain duplicate entries: duplicate entry for %s", inputsPath, key) + } + + if dup, ok := meta.FirstDuplicate(); ok { + key := fmt.Sprintf("account=%q key=%q", dup.Account, dup.Key) + if dup.Scope != "" { + key += fmt.Sprintf(" scope=%q", dup.Scope) + } + return fmt.Errorf("invalid inputs file '%s': metadata must not contain duplicate entries: duplicate entry for %s", inputsPath, key) + } + + return nil +} + func run(scriptPath string, opts RunArgs) error { numscriptContent, err := os.ReadFile(scriptPath) if err != nil { @@ -57,26 +83,8 @@ func run(scriptPath string, opts RunArgs) error { return fmt.Errorf("failed to parse inputs file '%s' as JSON: %w", inputsPath, err) } - // Reject a malformed inputs file before running anything: a balance list is a - // map keyed by (account, asset, color, scope), so a repeated key is ambiguous. - if dup, ok := inputs.Balances.FirstDuplicate(); ok { - key := fmt.Sprintf("account=%q asset=%q", dup.Account, dup.Asset) - if dup.Color != "" { - key += fmt.Sprintf(" color=%q", dup.Color) - } - if dup.Scope != "" { - key += fmt.Sprintf(" scope=%q", dup.Scope) - } - return fmt.Errorf("invalid inputs file '%s': balances must not contain duplicate entries: duplicate entry for %s", inputsPath, key) - } - - // Likewise, a metadata list is keyed by (account, key, scope). - if dup, ok := inputs.Meta.FirstDuplicate(); ok { - key := fmt.Sprintf("account=%q key=%q", dup.Account, dup.Key) - if dup.Scope != "" { - key += fmt.Sprintf(" scope=%q", dup.Scope) - } - return fmt.Errorf("invalid inputs file '%s': metadata must not contain duplicate entries: duplicate entry for %s", inputsPath, key) + if err := validateInputRows(inputsPath, inputs.Balances, inputs.Meta); err != nil { + return err } featureFlags := map[string]struct{}{} diff --git a/internal/compiler/bench_test.go b/internal/compiler/bench_test.go new file mode 100644 index 00000000..5c7db2b9 --- /dev/null +++ b/internal/compiler/bench_test.go @@ -0,0 +1,335 @@ +package compiler_test + +import ( + "context" + "math/big" + "testing" + + "github.com/formancehq/numscript/internal/compiler" + "github.com/formancehq/numscript/internal/funds" + "github.com/formancehq/numscript/internal/interpreter" + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/vm" +) + +// benchStore is a minimal vm.Store for the benchmarks. +type benchStore struct { + balances map[funds.PairKey]*big.Int +} + +func (s benchStore) GetBalance(ctx context.Context, account, scope, asset, color string) (*big.Int, error) { + if v, ok := s.balances[funds.PairKey{Account: account, Scope: scope, Asset: asset, Color: color}]; ok { + return v, nil + } + return new(big.Int), nil +} + +func (benchStore) GetMetadata(ctx context.Context, account, scope, key string) (string, bool, error) { + return "", false, nil +} + +type fundsStoreAdapter struct { + store vm.Store +} + +func (s fundsStoreAdapter) GetBalance( + account string, + scope string, + asset string, + color string, +) (*big.Int, error) { + return s.store.GetBalance(context.Background(), account, scope, asset, color) +} + +// Both benchmarks run the SAME program with the same starting balance; only the +// per-iteration RUN is measured (parse/compile/assemble happen once, up front). +const benchSrc = `send [USD/2 10] ( + source = @src + destination = @dest +)` + +// BenchmarkTreeWalker measures the tree-walking interpreter on a pre-parsed AST. +func BenchmarkTreeWalker(b *testing.B) { + parsed := parser.Parse(benchSrc) + if len(parsed.Errors) != 0 { + b.Fatalf("parse errors: %v", parsed.Errors) + } + store := interpreter.StaticStore{ + Balances: interpreter.Balances{ + {Account: "src", Asset: "USD/2", Amount: big.NewInt(100)}, + }, + } + ctx := context.Background() + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, err := interpreter.RunProgram(ctx, parsed.Value, nil, store, nil) + if err != nil { + b.Fatalf("run: %v", err) + } + } +} + +// BenchmarkRuntimeBaseline is the floor: it drives funds.RunState directly, +// performing exactly the funds operations the program lowers to — with no AST +// walk and no bytecode dispatch. It reuses one RunState (like the VM reuses its +// runstate) and hoists the constants (the compiler would pool them). The gap +// between this and BenchmarkCompiledVM is the VM's dispatch/register overhead; +// the gap to BenchmarkTreeWalker is the interpreter's front-end overhead. +func BenchmarkRuntimeBaseline(b *testing.B) { + store := fundsStoreAdapter{ + store: benchStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}, + } + + rs := funds.New(store) + + ten := big.NewInt(10) // the sent amount / pull cap + zero := big.NewInt(0) // bounded overdraft of 0 + pulled := new(big.Int) // reused output register + dest := "dest" + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + rs.Reset(store) + rs.SetCurrentAsset("USD/2") + _ = rs.Pull(pulled, "src", "", ten, zero, "") + _ = pulled.Cmp(ten) // CheckEnoughFunds + _ = rs.SendUncapped(&dest, "", nil) + _ = rs.GetPostings() + } +} + +// BenchmarkCompiledVM measures the compiled bytecode on the register VM, reusing +// a single Vm instance across iterations (its register banks are not realloc'd). +func BenchmarkCompiledVM(b *testing.B) { + parsed := parser.Parse(benchSrc) + if len(parsed.Errors) != 0 { + b.Fatalf("parse errors: %v", parsed.Errors) + } + _, program, err := compiler.Compile(parsed.Value, nil) + if err != nil { + b.Fatalf("compile: %v", err) + } + store := benchStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(100), + }} + + machine := vm.NewVm(program) // reused across iterations + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, err := vm.Exec(context.Background(), machine, nil, store) + if err != nil { + b.Fatalf("exec: %v", err) + } + } +} + +// --- Capped inorder script: `{ @a ; max [USD/2 5] from @b ; @c }` ----------- +// Same methodology as above, on a more representative script (inorder traversal, +// a `max` cap (a min, i.e. lt_int + copies), running total, and an early-exit +// jump). Balances: +// a=3, b=100 (capped to 5), c=100 → pulls 3 / 5 / 2. +const benchSrcCapped = `send [USD/2 10] ( + source = { + @a + max [USD/2 5] from @b + @c + } + destination = @dest +)` + +func BenchmarkTreeWalkerCapped(b *testing.B) { + parsed := parser.Parse(benchSrcCapped) + if len(parsed.Errors) != 0 { + b.Fatalf("parse errors: %v", parsed.Errors) + } + store := interpreter.StaticStore{ + Balances: interpreter.Balances{ + {Account: "a", Asset: "USD/2", Amount: big.NewInt(3)}, + {Account: "b", Asset: "USD/2", Amount: big.NewInt(100)}, + {Account: "c", Asset: "USD/2", Amount: big.NewInt(100)}, + }, + } + ctx := context.Background() + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, err := interpreter.RunProgram(ctx, parsed.Value, nil, store, nil) + if err != nil { + b.Fatalf("run: %v", err) + } + } +} + +func cappedStore() benchStore { + return benchStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2", Color: ""}: big.NewInt(3), + {Account: "b", Asset: "USD/2", Color: ""}: big.NewInt(100), + {Account: "c", Asset: "USD/2", Color: ""}: big.NewInt(100), + }} +} + +// BenchmarkRuntimeBaselineCapped is the floor: it drives funds.RunState +// directly, performing the funds ops the capped-inorder script lowers to (with +// the cap/running-total/early-exit arithmetic done inline on reused big.Ints) — +// no AST walk, no bytecode dispatch. RunState reused across iterations. +func BenchmarkRuntimeBaselineCapped(b *testing.B) { + store := fundsStoreAdapter{store: cappedStore()} + rs := funds.New(store) + + zero := big.NewInt(0) + ten := big.NewInt(10) + five := big.NewInt(5) + remaining := new(big.Int) + capB := new(big.Int) + pulled := new(big.Int) + total := new(big.Int) + dest := "dest" + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + rs.Reset(store) + rs.SetCurrentAsset("USD/2") + total.SetInt64(0) + remaining.Set(ten) // inorder cap = copy(amount) + + // @a (cap = remaining) + _ = rs.Pull(pulled, "a", "", remaining, zero, "") + total.Add(total, pulled) + remaining.Sub(remaining, pulled) + + if remaining.Sign() != 0 { // is_zero(remaining) + jmp_if_true + // max [USD/2 5] from @b -> cap = min(5, remaining) + if five.Cmp(remaining) < 0 { + capB.Set(five) + } else { + capB.Set(remaining) + } + _ = rs.Pull(pulled, "b", "", capB, zero, "") + total.Add(total, pulled) + remaining.Sub(remaining, pulled) + + if remaining.Sign() != 0 { + _ = rs.Pull(pulled, "c", "", remaining, zero, "") // @c (cap = remaining) + total.Add(total, pulled) + } + } + + _ = total.Cmp(ten) // check_enough_funds + _ = rs.SendUncapped(&dest, "", nil) + _ = rs.GetPostings() + } +} + +func BenchmarkCompiledVMCapped(b *testing.B) { + parsed := parser.Parse(benchSrcCapped) + if len(parsed.Errors) != 0 { + b.Fatalf("parse errors: %v", parsed.Errors) + } + _, program, err := compiler.Compile(parsed.Value, nil) + if err != nil { + b.Fatalf("compile: %v", err) + } + store := cappedStore() + + machine := vm.NewVm(program) // reused across iterations + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, err := vm.Exec(context.Background(), machine, nil, store) + if err != nil { + b.Fatalf("exec: %v", err) + } + } +} + +// --- Allotment scripts ------------------------------------------------------ +// Ported from feat/exp/optimize-vm so the before/after of decomposing the +// allotment split into pure ops is measurable on this branch. Same methodology. + +// benchCompiledVM is the shape the two benchmarks above open-code: compile once, +// reuse one Vm, measure only the run. +func benchCompiledVM(b *testing.B, src string, store benchStore) { + b.Helper() + + parsed := parser.Parse(src) + if len(parsed.Errors) != 0 { + b.Fatalf("parse errors: %v", parsed.Errors) + } + _, program, err := compiler.Compile(parsed.Value, nil) + if err != nil { + b.Fatalf("compile: %v", err) + } + + machine := vm.NewVm(program) // reused across iterations + + b.ReportAllocs() + b.ResetTimer() + for i := 0; i < b.N; i++ { + _, err := vm.Exec(context.Background(), machine, nil, store) + if err != nil { + b.Fatalf("exec: %v", err) + } + } +} + +// Fan-out allotment: 1 source -> {1/2 @a; 1/2 @b}. Exercises the allotment +// split and the queue drain across two capped sends. +const benchSrcAllotment = `send [USD/2 100] ( + source = @src + destination = { + 1/2 to @a + 1/2 to @b + } +)` + +func BenchmarkCompiledVMAllotment(b *testing.B) { + benchCompiledVM(b, benchSrcAllotment, benchStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(1000), + }}) +} + +// Thirds: the case where the flooring leftover is non-zero, so the fixup pass +// actually runs (100 -> 34/33/33). +const benchSrcAllotmentThirds = `send [USD/2 100] ( + source = @src + destination = { + 1/3 to @a + 1/3 to @b + remaining to @c + } +)` + +func BenchmarkCompiledVMAllotmentThirds(b *testing.B) { + benchCompiledVM(b, benchSrcAllotmentThirds, benchStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(1000), + }}) +} + +// Fan-in: {1/3 from @a; 1/3 from @b; 1/3 from @c} -> @dest. Allotment on the +// source side, no early-exit jump. +const benchSrcFanIn = `send [USD/2 30] ( + source = { + 1/3 from @a + 1/3 from @b + 1/3 from @c + } + destination = @dest +)` + +func BenchmarkCompiledVMFanIn(b *testing.B) { + benchCompiledVM(b, benchSrcFanIn, benchStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2", Color: ""}: big.NewInt(100), + {Account: "b", Asset: "USD/2", Color: ""}: big.NewInt(100), + {Account: "c", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}) +} diff --git a/internal/compiler/compile_error_test.go b/internal/compiler/compile_error_test.go new file mode 100644 index 00000000..26b26e20 --- /dev/null +++ b/internal/compiler/compile_error_test.go @@ -0,0 +1,243 @@ +package compiler + +// White-box tests asserting the concrete CompilerError produced for invalid +// programs. They call compileProgramToIR directly, since the public Compile +// stringifies the error and would lose the type. + +import ( + "testing" + + "github.com/formancehq/numscript/internal/flags" + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/typecheck" + "github.com/stretchr/testify/require" +) + +func TestE2E_RejectsUnboundVariable(t *testing.T) { + parsed := parser.Parse(`send [C 10] (source = $undeclared destination = @d)`) + require.Empty(t, parsed.Errors) + _, cErr := compileProgramToIR(parsed.Value, nil) + require.IsType(t, TypeError{}, cErr) + require.IsType(t, typecheck.UnboundVariable{}, cErr.(TypeError).Kind) +} + +func TestE2E_RejectsTypeMismatch(t *testing.T) { + parsed := parser.Parse(`vars { string $s } send [C 10] (source = $s destination = @d)`) + require.Empty(t, parsed.Errors) + _, cErr := compileProgramToIR(parsed.Value, nil) + require.IsType(t, TypeError{}, cErr) + require.IsType(t, typecheck.TypeMismatch{}, cErr.(TypeError).Kind) +} + +func TestE2E_RejectsMetaOutsideVarOrigin(t *testing.T) { + // meta() is only supported as a direct variable origin; nested in an + // expression it must be a compile error, not a panic. + parsed := parser.Parse(` + #![feature("experimental-mid-script-function-call")] + vars { + account $a + number $n = meta($a, "k") + 1 + } + send [C $n] (source = @world destination = @d) + `) + require.Empty(t, parsed.Errors) + _, cErr := compileProgramToIR(parsed.Value, nil) + require.IsType(t, InvalidMetaPosition{}, cErr) +} + +func TestE2E_RejectsNonCastableInterpVar(t *testing.T) { + // a monetary var has no string form: interpolating it must be a compile + // error (matching the interpreter's runtime CannotCastToString), not a panic. + parsed := parser.Parse(` + #![feature("experimental-account-interpolation")] + vars { monetary $m } + set_tx_meta("k", @acc:$m) + `) + require.Empty(t, parsed.Errors) + _, cErr := compileProgramToIR(parsed.Value, nil) + require.IsType(t, CannotCastToString{}, cErr) + require.Equal(t, typecheck.TypeMonetary, cErr.(CannotCastToString).Type) +} + +func TestE2E_AllotmentDuplicateRemaining(t *testing.T) { + parsed := parser.Parse(` + send [USD/2 100] ( + source = @world + destination = { + remaining to @a + remaining to @b + } + ) + `) + require.Empty(t, parsed.Errors) + _, cErr := compileProgramToIR(parsed.Value, nil) + require.IsType(t, DuplicateRemaining{}, cErr) +} + +// --- feature flags + +// each case is a construct gated behind a feature flag: compiling it without the +// flag must fail, and compiling it with the flag must get past the gate. +func TestFeatureFlagGating(t *testing.T) { + testCases := []struct { + name string + flag flags.FeatureFlag + src string + }{ + { + name: "oneof in source", + flag: flags.ExperimentalOneofFeatureFlag, + src: `send [C 10] ( + source = oneof { @a @b } + destination = @d + )`, + }, + { + name: "oneof in destination", + flag: flags.ExperimentalOneofFeatureFlag, + src: `send [C 10] ( + source = @world + destination = oneof { + max [C 3] to @a + remaining to @b + } + )`, + }, + { + name: "account interpolation", + flag: flags.ExperimentalAccountInterpolationFlag, + src: `vars { string $s } + send [C 10] (source = @world destination = @dest:$s)`, + }, + { + name: "mid-script function call", + flag: flags.ExperimentalMidScriptFunctionCall, + src: `send balance(@a, C) (source = @world destination = @d)`, + }, + { + name: "overdraft function", + flag: flags.ExperimentalOverdraftFunctionFeatureFlag, + src: `vars { monetary $m = overdraft(@a, C) } + send $m (source = @world destination = @d)`, + }, + { + name: "get_asset function", + flag: flags.ExperimentalGetAssetFunctionFeatureFlag, + src: `vars { monetary $m asset $a = get_asset($m) } + send [$a 10] (source = @world destination = @d)`, + }, + { + name: "get_amount function", + flag: flags.ExperimentalGetAmountFunctionFeatureFlag, + src: `vars { monetary $m number $n = get_amount($m) } + send [C $n] (source = @world destination = @d)`, + }, + { + name: "asset colors", + flag: flags.ExperimentalAssetColors, + src: `send [C 10] (source = @a \ "RED" destination = @d)`, + }, + { + name: "asset scaling", + flag: flags.AssetScaling, + src: `send [C 10] ( + source = @src with scaling through @swap + destination = @d + )`, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + parsed := parser.Parse(tc.src) + require.Empty(t, parsed.Errors) + + _, cErr := compileProgramToIR(parsed.Value, nil) + require.IsType(t, ExperimentalFeature{}, cErr) + require.Equal(t, tc.flag, cErr.(ExperimentalFeature).FlagName) + + // with the flag on, whatever comes back must not be about the flag + // (scaling still hits FeatureNotImplemented) + _, cErr = compileProgramToIR(parsed.Value, map[string]struct{}{tc.flag: {}}) + _, stillGated := cErr.(ExperimentalFeature) + require.False(t, stillGated, "still gated with the flag on: %v", cErr) + }) + } +} + +// a function call that *is* the variable's origin is not a mid-script call +func TestFnCallAsVarOriginIsNotMidScript(t *testing.T) { + parsed := parser.Parse(` + vars { monetary $m = balance(@a, C) } + send $m (source = @world destination = @d) + `) + require.Empty(t, parsed.Errors) + _, cErr := compileProgramToIR(parsed.Value, nil) + require.Nil(t, cErr) +} + +// ... but one nested inside the origin expression is +func TestNestedFnCallInVarOriginIsMidScript(t *testing.T) { + parsed := parser.Parse(` + vars { monetary $m = balance(@a, C) + balance(@b, C) } + send $m (source = @world destination = @d) + `) + require.Empty(t, parsed.Errors) + _, cErr := compileProgramToIR(parsed.Value, nil) + require.IsType(t, ExperimentalFeature{}, cErr) + require.Equal(t, flags.ExperimentalMidScriptFunctionCall, cErr.(ExperimentalFeature).FlagName) +} + +// #![feature(..)] in the source enables a flag the host didn't pass +func TestInSourceFeatureDeclaration(t *testing.T) { + parsed := parser.Parse(` + #![feature("experimental-oneof")] + send [C 10] ( + source = oneof { @a @b } + destination = @d + ) + `) + require.Empty(t, parsed.Errors) + _, cErr := compileProgramToIR(parsed.Value, nil) + require.Nil(t, cErr) +} + +func TestInSourceFeatureDeclarationRejectsUnknownFlag(t *testing.T) { + parsed := parser.Parse(` + #![feature("not-a-flag")] + send [C 10] (source = @world destination = @d) + `) + require.Empty(t, parsed.Errors) + _, cErr := compileProgramToIR(parsed.Value, nil) + require.IsType(t, InvalidFeature{}, cErr) + require.Equal(t, "not-a-flag", cErr.(InvalidFeature).Feature) +} + +// Every CompilerError must carry a human-readable message: CompilerError is +// parser.Ranged + compileError(), so a type missing Error() still satisfies the +// interface and Compile's fmt.Errorf("%v") would print the raw struct instead. +func TestCompilerErrorMessages(t *testing.T) { + testCases := []struct { + name string + err CompilerError + msg string + }{ + {"UnboundVar", UnboundVar{Var: "x"}, "the variable '$x' was not declared"}, + {"TypeError", TypeError{Kind: typecheck.UnboundVariable{Name: "x"}}, "The variable '$x' was not declared"}, + {"InvalidUncappedSource", InvalidUncappedSource{}, "cannot take all balance of an unbounded source"}, + {"DuplicateRemaining", DuplicateRemaining{}, "only one 'remaining' clause is allowed in an allotment expression"}, + {"InvalidMetaPosition", InvalidMetaPosition{}, "meta() is only allowed as a variable origin"}, + {"CannotCastToString", CannotCastToString{Type: typecheck.TypeMonetary}, "cannot cast a value of type monetary to string"}, + {"FeatureNotImplemented", FeatureNotImplemented{Feature: "scaling"}, "internal error: feature not implemented: scaling"}, + {"ExperimentalFeature", ExperimentalFeature{FlagName: flags.ExperimentalAssetColors}, "You need the 'experimental-asset-colors' feature flag to enable it"}, + {"InvalidFeature", InvalidFeature{Feature: "nope"}, "Invalid feature: nope"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + err, ok := tc.err.(error) + require.True(t, ok, "%T does not implement error", tc.err) + require.Contains(t, err.Error(), tc.msg) + }) + } +} diff --git a/internal/compiler/compiler.go b/internal/compiler/compiler.go new file mode 100644 index 00000000..9db110a2 --- /dev/null +++ b/internal/compiler/compiler.go @@ -0,0 +1,1553 @@ +package compiler + +import ( + "fmt" + "maps" + "math/big" + "slices" + + "github.com/formancehq/numscript/internal/builtins" + "github.com/formancehq/numscript/internal/flags" + "github.com/formancehq/numscript/internal/ir" + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/typecheck" + "github.com/formancehq/numscript/internal/vm" +) + +// Compile lowers a parsed program to the VarsEncoder that turns a json var +// payload into the vm.Vars the program expects, plus the vm.Program itself. +// +// featureFlags is the set of experimental features the host allows; a construct +// gated behind a flag that isn't in the set fails compilation. As in +// interpreter.RunProgram, the script's own #![feature(..)] declarations are +// unioned in. +func Compile(program parser.Program, featureFlags map[string]struct{}) (VarsEncoder, vm.Program, error) { + compiled, cErr := compileProgramToIR(program, featureFlags) + if cErr != nil { + return VarsEncoder{}, vm.Program{}, fmt.Errorf("%v", cErr) + } + + if err := ir.Typecheck(compiled.instructions); err != nil { + return VarsEncoder{}, vm.Program{}, err + } + + prog, err := ir.Assemble(compiled.instructions) + if err != nil { + return VarsEncoder{}, vm.Program{}, err + } + + return compiled.varsEncoder, prog, nil +} + +type compiledProgramIR struct { + instructions []ir.Instr + varsEncoder VarsEncoder +} + +type state struct { + ir.Builder + + vars map[string]value + exprTypes map[parser.ValueExpr]typecheck.Type + featureFlags map[string]struct{} + // set by compileSentValue before any source/destination is compiled; nil + // until then, so compileCapAmount can't silently assert against register 0. + currentAssetReg *ir.Reg + + nextIntVar int + nextStrVar int + varDecls []varDecl + + // holds worldAccount; see pullFromAccount + worldReg ir.Reg +} + +// The unbounded account. The VM knows nothing about it: a source account is a +// register, so the comparison is compiled, not built in. +const worldAccount = "world" + +func (st *state) checkFeatureFlag(rng parser.Range, flag flags.FeatureFlag) CompilerError { + if _, ok := st.featureFlags[flag]; ok { + return nil + } + return ExperimentalFeature{Range: rng, FlagName: flag} +} + +// pushInstructionWithDestErr is PushWithDest in the shape compileExpr returns. +func (st *state) pushInstructionWithDestErr(getInstr func(dest ir.Reg) ir.Instr) (ir.Reg, CompilerError) { + return st.PushWithDest(getInstr), nil +} + +func (st *state) compileAllot(amount ir.Reg, allotments []parser.AllotmentValue) ([]ir.Reg, CompilerError) { + n := len(allotments) + portions := make([]ir.Reg, n) + remainingIdx := -1 + for i, al := range allotments { + switch al := al.(type) { + case *parser.ValueExprAllotment: + p, err := st.compileExpr(al.Value) + if err != nil { + return nil, err + } + if !nonNegativePortionExpr(al.Value) { + st.Push(ir.AssertNonNegativePortion{Portion: p}) + } + portions[i] = p + case *parser.RemainingAllotment: + if remainingIdx != -1 { + return nil, DuplicateRemaining{Range: al.Range} + } + remainingIdx = i + default: + return nil, UnsupportedNode{Node: al} + } + } + + leftover := st.compilePortionOne() + for i := range allotments { + if i == remainingIdx { + continue + } + prev, pi := leftover, portions[i] + leftover = st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpSubPortion{}, Left: prev, Right: pi, Dest: dest} + }) + } + + st.Push(ir.AssertLeftover{Portion: leftover, Exact: remainingIdx == -1}) + if remainingIdx != -1 { + portions[remainingIdx] = leftover + } + + return st.compileAllotmentSplit(amount, portions), nil +} + +// nonNegativePortionExpr reports whether a portion expression cannot evaluate to a +// negative value: portion vars and meta() portions are range-checked when parsed, +// so only a division with a non-literal or negative operand can. +func nonNegativePortionExpr(expr parser.ValueExpr) bool { + switch expr := expr.(type) { + case *parser.PercentageLiteral: + return expr.Amount.Sign() >= 0 + case *parser.Variable: + return true + case *parser.BinaryInfix: + num, numOk := expr.Left.(*parser.NumberLiteral) + den, denOk := expr.Right.(*parser.NumberLiteral) + return expr.Operator == parser.InfixOperatorDiv && numOk && denOk && + num.Number.Sign() >= 0 && den.Number.Sign() > 0 + default: + return false + } +} + +// TODO properly review claude-generated compileAllotmentSplit + +// compileAllotmentSplit writes the amount split across the portions: one int +// register per portion, summing exactly to amount. It expects the portions to +// sum to 1, which is what the assert_leftover emitted by the caller establishes. +// +// Each share is floor(portion * amount); flooring loses strictly less than one +// unit per share, so the shortfall is under len(portions) and a single +// front-to-back pass handing out one unit each closes it. That order is +// observable — 100 by thirds is 34/33/33, not 33/33/34. +func (st *state) compileAllotmentSplit(amount ir.Reg, portions []ir.Reg) []ir.Reg { + n := len(portions) + dest := make([]ir.Reg, n) + + amountPortion := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpIntToPortion{}, Arg: amount, Dest: dest} + }) + + // total accumulates the floored shares; it starts as a copy of the first one + // rather than a zero, which saves a load + var total ir.Reg + for i, portion := range portions { + product := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpMulPortion{}, Left: portion, Right: amountPortion, Dest: dest} + }) + dest[i] = st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpPortionToInt{}, Arg: product, Dest: dest} + }) + + if i == 0 { + total = st.PushWithDest(func(t ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpIntCopy{}, Arg: dest[0], Dest: t} + }) + continue + } + st.Push(ir.BinaryOp{Op: ir.OpAddInt{}, Left: total, Right: dest[i], Dest: total}) + } + + // The shortfall is at most n-1, so the last share never receives a unit and + // its block would be dead. The jumps go forward to one shared exit, which is + // what lets this be a straight line: the assembler rejects backward jumps. + if n > 1 { + one := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{Value: *big.NewInt(1), Dest: dest} + }) + done := st.FreshLabel("allot_end") + + for i := 0; i < n-1; i++ { + short := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpLtInt{}, Left: total, Right: amount, Dest: dest} + }) + st.Push(ir.JmpIfFalse{Cond: short, Target: done}) + st.Push(ir.BinaryOp{Op: ir.OpAddInt{}, Left: dest[i], Right: one, Dest: dest[i]}) + st.Push(ir.BinaryOp{Op: ir.OpAddInt{}, Left: total, Right: one, Dest: total}) + } + + st.Push(ir.LabelMarker{Label: done}) + } + + return dest +} + +func (st *state) compileCapAmount(monExpr parser.ValueExpr) (ir.Reg, CompilerError) { + mon, err := st.compileMonetaryExpr(monExpr) + if err != nil { + return 0, err + } + if st.currentAssetReg == nil { + panic("compileCapAmount: no current asset (compileSentValue must run first)") + } + st.Push(ir.AssertSameAsset{Left: *st.currentAssetReg, Right: mon.Asset}) + return mon.Amount, nil +} + +func (st *state) compilePortionOne() ir.Reg { + one := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{Value: *big.NewInt(1), Dest: dest} + }) + return st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpMakePortion{}, Left: one, Right: one, Dest: dest} + }) +} + +// compileExpr compiles a non-monetary expression into the single register its +// type maps to. Monetary-typed expressions go to compileMonetaryExpr instead, +// since a monetary needs two registers. +func (st *state) compileExpr(expr parser.ValueExpr) (ir.Reg, CompilerError) { + if st.exprTypes[expr] == typecheck.TypeMonetary { + panic("compileExpr: monetary expression (use compileMonetaryExpr)") + } + + switch expr := expr.(type) { + case *parser.AssetLiteral: + return st.pushInstructionWithDestErr(func(dest ir.Reg) ir.Instr { + return ir.LoadStr{ + Value: expr.Asset, + Dest: dest, + } + }) + + case *parser.StringLiteral: + return st.pushInstructionWithDestErr(func(dest ir.Reg) ir.Instr { + return ir.LoadStr{ + Value: expr.String, + Dest: dest, + } + }) + + case *parser.NumberLiteral: + return st.pushInstructionWithDestErr(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{ + Value: *expr.Number, + Dest: dest, + } + }) + + case *parser.AccountInterpLiteral: + var parts []ir.Reg + hasVar := false + for _, part := range expr.Parts { + switch part := part.(type) { + case parser.AccountTextPart: + dest := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadStr{ + Value: part.Name, + Dest: dest, + } + }) + parts = append(parts, dest) + case *parser.Variable: + if err := st.checkFeatureFlag(part.Range, flags.ExperimentalAccountInterpolationFlag); err != nil { + return 0, err + } + hasVar = true + // reject before compiling, so a non-castable part doesn't reach + // compileExpr (which only handles non-monetary expressions) + t := st.exprTypes[part] + switch t { + case typecheck.TypeAccount, typecheck.TypeString, typecheck.TypeNumber: + default: + return 0, CannotCastToString{Range: part.GetRange(), Type: t} + } + var r ir.Reg + if t == typecheck.TypeAccount { + acc, err := st.compileAccountExpr(part) + if err != nil { + return 0, err + } + if acc.Scope != nil { + st.Push(ir.AssertUnscoped{Scope: *acc.Scope, Account: acc.Name}) + } + r = acc.Name + } else { + var err CompilerError + r, err = st.compileExpr(part) + if err != nil { + return 0, err + } + } + if t == typecheck.TypeNumber { + r = st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpIntToString{}, Arg: r, Dest: dest} + }) + } + parts = append(parts, r) + } + } + + acc := parts[0] + for _, part := range parts[1:] { + left, right := acc, part + acc = st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpAddString{}, Left: left, Right: right, Dest: dest} + }) + } + // an interpolated var can inject chars that make the name ill-formed; + // all-text literals are valid by construction, so skip the check + if hasVar { + st.Push(ir.AssertValidAccount{Account: acc}) + } + return acc, nil + + case *parser.Variable: + v, ok := st.vars[expr.Name] + if !ok { + return 0, UnboundVar{Range: expr.Range, Var: expr.Name} + } + return v.Reg, nil + + case *parser.PercentageLiteral: + // e.g. 50% -> portion 50/100; mk_portion reduces via SetFrac + ratio := expr.ToRatio() + numReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{Value: *ratio.Num(), Dest: dest} + }) + denReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{Value: *ratio.Denom(), Dest: dest} + }) + return st.pushInstructionWithDestErr(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpMakePortion{}, Left: numReg, Right: denReg, Dest: dest} + }) + + case *parser.BinaryInfix: + leftReg, err := st.compileExpr(expr.Left) + if err != nil { + return 0, err + } + rightReg, err := st.compileExpr(expr.Right) + if err != nil { + return 0, err + } + + switch expr.Operator { + case parser.InfixOperatorDiv: + return st.pushInstructionWithDestErr(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpMakePortion{}, Left: leftReg, Right: rightReg, Dest: dest} + }) + + case parser.InfixOperatorPlus: + return st.pushInstructionWithDestErr(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpAddInt{}, Left: leftReg, Right: rightReg, Dest: dest} + }) + + case parser.InfixOperatorMinus: + return st.pushInstructionWithDestErr(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpSubInt{}, Left: leftReg, Right: rightReg, Dest: dest} + }) + + default: + return 0, UnsupportedNode{Range: expr.Range, Node: expr.Operator} + } + + case *parser.Prefix: + switch expr.Operator { + case parser.PrefixOperatorMinus: + argReg, err := st.compileExpr(expr.Expr) + if err != nil { + return 0, err + } + return st.pushInstructionWithDestErr(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpNegInt{}, Arg: argReg, Dest: dest} + }) + + default: + return 0, UnsupportedNode{Range: expr.Range, Node: expr.Operator} + } + + case *parser.FnCall: + return st.compileFnCall(expr, false) + + default: + return 0, UnsupportedNode{Range: expr.GetRange(), Node: expr} + } +} + +// compileFnCall takes isVarOrigin to tell apart the two positions the interpreter +// distinguishes: a call that *is* a variable's origin expression, versus one +// nested anywhere else (which needs the mid-script-function-call flag). +func (st *state) compileFnCall(expr *parser.FnCall, isVarOrigin bool) (ir.Reg, CompilerError) { + if !isVarOrigin { + if err := st.checkFeatureFlag(expr.Range, flags.ExperimentalMidScriptFunctionCall); err != nil { + return 0, err + } + } + + switch expr.Caller.Name { + case builtins.GetAmount: + if err := st.checkFeatureFlag(expr.Range, flags.ExperimentalGetAmountFunctionFeatureFlag); err != nil { + return 0, err + } + mon, err := st.compileMonetaryExpr(expr.Args[0]) + if err != nil { + return 0, err + } + return mon.Amount, nil + + case builtins.GetAsset: + if err := st.checkFeatureFlag(expr.Range, flags.ExperimentalGetAssetFunctionFeatureFlag); err != nil { + return 0, err + } + mon, err := st.compileMonetaryExpr(expr.Args[0]) + if err != nil { + return 0, err + } + return mon.Asset, nil + + case builtins.Meta: + return 0, InvalidMetaPosition{Range: expr.Range} + + // scoped() only ever reaches here if some caller compiled a TypeAccount + // expression through the generic (single-register) path instead of + // compileAccountExpr — every real call site is routed through the latter, so + // this is a defensive backstop, not an expected path. + case builtins.Scoped: + return 0, InvalidScopedAccountPosition{Range: expr.Range} + + default: + return 0, UnsupportedNode{Range: expr.Range, Node: expr.Caller.Name} + } +} + +// compileMonetaryExpr compiles a monetary-typed expression into the (asset, +// amount) register pair. Which expressions reach here is decided by +// st.exprTypes; compileExpr rejects monetary-typed ones. +func (st *state) compileMonetaryExpr(expr parser.ValueExpr) (monetaryValue, CompilerError) { + switch expr := expr.(type) { + case *parser.MonetaryLiteral: + assetReg, err := st.compileExpr(expr.Asset) + if err != nil { + return monetaryValue{}, err + } + amtReg, err := st.compileExpr(expr.Amount) + if err != nil { + return monetaryValue{}, err + } + return monetaryValue{Asset: assetReg, Amount: amtReg}, nil + + case *parser.Variable: + v, ok := st.vars[expr.Name] + if !ok { + return monetaryValue{}, UnboundVar{Range: expr.Range, Var: expr.Name} + } + if v.Mon == nil { + panic("compileMonetaryExpr: $" + expr.Name + " is not a monetary") + } + return *v.Mon, nil + + case *parser.BinaryInfix: + left, err := st.compileMonetaryExpr(expr.Left) + if err != nil { + return monetaryValue{}, err + } + right, err := st.compileMonetaryExpr(expr.Right) + if err != nil { + return monetaryValue{}, err + } + st.Push(ir.AssertSameAsset{Left: left.Asset, Right: right.Asset}) + + var op ir.BinKind + switch expr.Operator { + case parser.InfixOperatorPlus: + op = ir.OpAddInt{} + case parser.InfixOperatorMinus: + op = ir.OpSubInt{} + default: + return monetaryValue{}, UnsupportedNode{Range: expr.Range, Node: expr.Operator} + } + amount := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: op, Left: left.Amount, Right: right.Amount, Dest: dest} + }) + // the assert above makes left vs right immaterial + return monetaryValue{Asset: left.Asset, Amount: amount}, nil + + case *parser.Prefix: + if expr.Operator != parser.PrefixOperatorMinus { + return monetaryValue{}, UnsupportedNode{Range: expr.Range, Node: expr.Operator} + } + arg, err := st.compileMonetaryExpr(expr.Expr) + if err != nil { + return monetaryValue{}, err + } + amount := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpNegInt{}, Arg: arg.Amount, Dest: dest} + }) + return monetaryValue{Asset: arg.Asset, Amount: amount}, nil + + case *parser.FnCall: + return st.compileMonetaryFnCall(expr, false) + + default: + return monetaryValue{}, UnsupportedNode{Range: expr.GetRange(), Node: expr} + } +} + +// compileMonetaryFnCall handles the builtins that return a monetary. isVarOrigin +// carries the same meaning as in compileFnCall. +func (st *state) compileMonetaryFnCall(expr *parser.FnCall, isVarOrigin bool) (monetaryValue, CompilerError) { + if !isVarOrigin { + if err := st.checkFeatureFlag(expr.Range, flags.ExperimentalMidScriptFunctionCall); err != nil { + return monetaryValue{}, err + } + } + + switch expr.Caller.Name { + case builtins.Balance: + acc, err := st.compileAccountExpr(expr.Args[0]) + if err != nil { + return monetaryValue{}, err + } + assetReg, err := st.compileExpr(expr.Args[1]) + if err != nil { + return monetaryValue{}, err + } + balReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.FetchBalance{Dest: dest, Account: acc.Name, Asset: assetReg, Scope: acc.Scope} + }) + st.Push(ir.AssertNonNegativeBalance{Balance: balReg, Account: acc.Name}) + return monetaryValue{Asset: assetReg, Amount: balReg}, nil + + case builtins.Overdraft: + if err := st.checkFeatureFlag(expr.Range, flags.ExperimentalOverdraftFunctionFeatureFlag); err != nil { + return monetaryValue{}, err + } + acc, err := st.compileAccountExpr(expr.Args[0]) + if err != nil { + return monetaryValue{}, err + } + assetReg, err := st.compileExpr(expr.Args[1]) + if err != nil { + return monetaryValue{}, err + } + balReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.FetchBalance{Dest: dest, Account: acc.Name, Asset: assetReg, Scope: acc.Scope} + }) + zeroReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{Value: *big.NewInt(0), Dest: dest} + }) + // overdraft = max(0, -balance) = -min(balance, 0) + minReg := st.minInt(balReg, zeroReg) + negReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpNegInt{}, Arg: minReg, Dest: dest} + }) + return monetaryValue{Asset: assetReg, Amount: negReg}, nil + + case builtins.Meta: + return monetaryValue{}, InvalidMetaPosition{Range: expr.Range} + + default: + return monetaryValue{}, UnsupportedNode{Range: expr.Range, Node: expr.Caller.Name} + } +} + +// compileAccountExpr compiles a TypeAccount expression into its (name, scope) +// pair, mirroring compileMonetaryExpr's role for TypeMonetary. Scope is nil +// unless the expression is (or resolves, through a var, to) a scoped() call. +func (st *state) compileAccountExpr(expr parser.ValueExpr) (accountValue, CompilerError) { + switch expr := expr.(type) { + case *parser.Variable: + v, ok := st.vars[expr.Name] + if !ok { + return accountValue{}, UnboundVar{Range: expr.Range, Var: expr.Name} + } + if v.Acc != nil { + return *v.Acc, nil + } + return accountValue{Name: v.Reg}, nil + + case *parser.FnCall: + // scoped() is the only builtin whose return type is TypeAccount, so any + // FnCall reaching here (checked by typecheck already) must be it. + return st.compileScopedFnCall(expr, false) + + default: + r, err := st.compileExpr(expr) + if err != nil { + return accountValue{}, err + } + return accountValue{Name: r}, nil + } +} + +// compileScopedFnCall compiles a scoped(account, scope) call into the (name, +// scope) pair. isVarOrigin carries the same meaning as in compileFnCall / +// compileMonetaryFnCall: true only when this call is itself a variable's whole +// origin expression. +func (st *state) compileScopedFnCall(expr *parser.FnCall, isVarOrigin bool) (accountValue, CompilerError) { + if !isVarOrigin { + if err := st.checkFeatureFlag(expr.Range, flags.ExperimentalMidScriptFunctionCall); err != nil { + return accountValue{}, err + } + } + if err := st.checkFeatureFlag(expr.Range, flags.ExperimentalScopedFunction); err != nil { + return accountValue{}, err + } + + inner, err := st.compileAccountExpr(expr.Args[0]) + if err != nil { + return accountValue{}, err + } + scopeReg, err := st.compileExpr(expr.Args[1]) + if err != nil { + return accountValue{}, err + } + st.Push(ir.AssertValidScope{Scope: scopeReg}) + + // scoped(scoped(x, "a"), "b") overwrites: the result is scoped "b", not "a". + return accountValue{Name: inner.Name, Scope: &scopeReg}, nil +} + +// compileColor returns nil when the source has no color clause: PullAccount with +// no color pulls the uncolored balance, same as an empty color string. +func (st *state) compileColor(colorExpr parser.ValueExpr) (*ir.Reg, CompilerError) { + if colorExpr == nil { + return nil, nil + } + if err := st.checkFeatureFlag(colorExpr.GetRange(), flags.ExperimentalAssetColors); err != nil { + return nil, err + } + reg, err := st.compileExpr(colorExpr) + if err != nil { + return nil, err + } + st.Push(ir.AssertValidColor{Color: reg}) + return ®, nil +} + +// pullFromAccount emits the pull of a source account, including the @world +// check. The account is a register — it can come from a var, an interpolation or +// metadata — so the check cannot be decided here and becomes a run-time branch: +// +// $eq = str_eq($account, $world) +// jmp_if_false($eq, #not_world) +// $pulled = pull_account(...) // no overdraft operand: unbounded +// jmp(#pull_end) +// #not_world +// $pulled = pull_account(..., overdraft: $od) +// #pull_end +// +// Both arms write the same dest, which the register typechecker allows because +// the type doesn't change. When overdraftReg is nil the source is unbounded for +// every account, so the two arms would be identical and the branch is skipped. +// A literal @world still gets the branch; collapsing it is a peephole's job +// (const-fold str_eq, then drop the dead arm). +func (st *state) pullFromAccount(acc accountValue, capReg, overdraftReg, colorReg *ir.Reg) ir.Reg { + pull := func(dest ir.Reg, overdraft *ir.Reg) ir.Instr { + return ir.PullAccount{ + Dest: dest, + Account: acc.Name, + Cap: capReg, + Overdraft: overdraft, + Color: colorReg, + Scope: acc.Scope, + } + } + + if overdraftReg == nil { + return st.PushWithDest(func(dest ir.Reg) ir.Instr { return pull(dest, nil) }) + } + + isWorld := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpStrEq{}, Left: acc.Name, Right: st.worldReg, Dest: dest} + }) + notWorldLabel := st.FreshLabel("not_world") + endLabel := st.FreshLabel("pull_end") + + st.Push(ir.JmpIfFalse{Cond: isWorld, Target: notWorldLabel}) + // an uncapped context reaches this arm with no cap and no overdraft, which is + // the InvalidUncappedSource case: taking *all* of an unbounded source + pulledReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { return pull(dest, nil) }) + st.Push(ir.Jmp{Target: endLabel}) + + st.Push(ir.LabelMarker{Label: notWorldLabel}) + st.Push(pull(pulledReg, overdraftReg)) + st.Push(ir.LabelMarker{Label: endLabel}) + + return pulledReg +} + +// minInt writes min(leftReg, rightReg) into a fresh register. There is no min +// opcode, so it is a comparison, a copy and a branch. Speculatively copying the +// left operand first saves the `jmp` the else arm would otherwise need: +// +// $min = int_copy($left) +// $lt = lt_int($left, $right) +// jmp_if_true($lt, #min_end) ; left is already the answer +// $min = int_copy($right) +// #min_end +// +// That form is only correct because the dest is freshly allocated: an aliased +// dest would clobber $right before the else arm reads it. +func (st *state) minInt(leftReg, rightReg ir.Reg) ir.Reg { + minReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpIntCopy{}, Arg: leftReg, Dest: dest} + }) + lt := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpLtInt{}, Left: leftReg, Right: rightReg, Dest: dest} + }) + + endLabel := st.FreshLabel("min_end") + st.Push(ir.JmpIfTrue{Cond: lt, Target: endLabel}) + st.Push(ir.UnaryOp{Op: ir.OpIntCopy{}, Arg: rightReg, Dest: minReg}) + st.Push(ir.LabelMarker{Label: endLabel}) + + return minReg +} + +func (st *state) maxInt(leftReg, rightReg ir.Reg) ir.Reg { + maxReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpIntCopy{}, Arg: leftReg, Dest: dest} + }) + gt := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpLtInt{}, Left: rightReg, Right: leftReg, Dest: dest} + }) + + endLabel := st.FreshLabel("max_end") + st.Push(ir.JmpIfTrue{Cond: gt, Target: endLabel}) + st.Push(ir.UnaryOp{Op: ir.OpIntCopy{}, Arg: rightReg, Dest: maxReg}) + st.Push(ir.LabelMarker{Label: endLabel}) + + return maxReg +} + +// The conditional jumps take a bool, so a quantity has to be projected onto one +// first — which is what stops a monetary amount from being used as a condition by +// accident (ir.Typecheck rejects it). +func (st *state) jmpIfAmountZero(amountReg ir.Reg, target ir.Label) { + isZero := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpIsZero{}, Arg: amountReg, Dest: dest} + }) + st.Push(ir.JmpIfTrue{Cond: isZero, Target: target}) +} + +// capReg is the register containing the current cap (or nil if context is uncapped) +// returns (when there's no err) the register where we store the pulled amount of this source +func (st *state) compileSource( + capReg *ir.Reg, + src parser.Source, +) (ir.Reg, CompilerError) { + switch src := src.(type) { + case *parser.SourceAccount: + acc, err := st.compileAccountExpr(src.ValueExpr) + if err != nil { + return 0, err + } + + colorReg, err := st.compileColor(src.Color) + if err != nil { + return 0, err + } + + overdraftReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{ + Value: *big.NewInt(0), + Dest: dest, + } + }) + + return st.pullFromAccount(acc, capReg, &overdraftReg, colorReg), nil + + case *parser.SourceOverdraft: + if src.Bounded == nil && capReg == nil { + return 0, InvalidUncappedSource{ + Range: src.GetRange(), + } + } + + acc, err := st.compileAccountExpr(src.Address) + if err != nil { + return 0, err + } + + colorReg, err := st.compileColor(src.Color) + if err != nil { + return 0, err + } + + var overdraftReg *ir.Reg + if src.Bounded != nil { + amtReg, err := st.compileCapAmount(*src.Bounded) + if err != nil { + return 0, err + } + overdraftReg = &amtReg + } + + return st.pullFromAccount(acc, capReg, overdraftReg, colorReg), nil + + case *parser.SourceCapped: + clauseCapIntReg, err := st.compileCapAmount(src.Cap) + if err != nil { + return 0, err + } + + // mirrors the interpreter's tryTakingUpTo/takeAll: NonNeg(min(amount, + // cap)). The clamp matters when the inner source is an allotment, which + // splits the cap before any pull gets a chance to clamp it. + innerCapReg := clauseCapIntReg + if capReg != nil { + innerCapReg = st.minInt(clauseCapIntReg, *capReg) + } + zeroReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{Value: *big.NewInt(0), Dest: dest} + }) + innerCapReg = st.maxInt(innerCapReg, zeroReg) + + return st.compileSource(&innerCapReg, src.From) + + case *parser.SourceInorder: + if capReg == nil { + inorderTotalReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{ + Value: *big.NewInt(0), + Dest: dest, + } + }) + for _, subSrc := range src.Sources { + innerPulledAmtReg, err := st.compileSource(nil, subSrc) + if err != nil { + return 0, err + } + // inorderTotalReg += innerPulledAmtReg + st.Push(ir.BinaryOp{ + Op: ir.OpAddInt{}, + Dest: inorderTotalReg, + Left: inorderTotalReg, + Right: innerPulledAmtReg, + }) + } + return inorderTotalReg, nil + } + + inorderTotalReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{ + Value: *big.NewInt(0), + Dest: dest, + } + }) + + inorderCap := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{ + Op: ir.OpIntCopy{}, + Arg: *capReg, + Dest: dest, + } + }) + + // every clause runs, even once the cap is exhausted: the interpreter + // evaluates later clauses' cap/account expressions (and pulls zero), so + // an early exit would skip failures it reports + for idx, subSrc := range src.Sources { + innerPulledAmtReg, err := st.compileSource(&inorderCap, subSrc) + if err != nil { + return 0, err + } + + // inorderTotalReg += innerPulledAmtReg + st.Push(ir.BinaryOp{ + Op: ir.OpAddInt{}, + Dest: inorderTotalReg, + Left: inorderTotalReg, + Right: innerPulledAmtReg, + }) + + isLast := idx == len(src.Sources)-1 + if !isLast { + // inorderCap -= innerPulledAmtReg + st.Push(ir.BinaryOp{ + Op: ir.OpSubInt{}, + Dest: inorderCap, + Left: inorderCap, + Right: innerPulledAmtReg, + }) + } + } + return inorderTotalReg, nil + + case *parser.SourceOneof: + if err := st.checkFeatureFlag(src.GetRange(), flags.ExperimentalOneofFeatureFlag); err != nil { + return 0, err + } + + if capReg == nil || len(src.Sources) == 1 { + return st.compileSource(capReg, src.Sources[0]) + } + + endLabel := st.FreshLabel("oneof_end") + + st.Push(ir.MarkPush{}) + + // allocated at first use, not up front, to keep registers numbered in + // emission order (see ir.Builder.PushWithDest) + var resultReg ir.Reg + + for index, subSrc := range src.Sources { + subPulledAmtReg, err := st.compileSource(capReg, subSrc) + if err != nil { + return 0, err + } + + if index == 0 { + resultReg = st.FreshReg() + } + + st.Push(ir.UnaryOp{ + Op: ir.OpIntCopy{}, + Arg: subPulledAmtReg, + Dest: resultReg, + }) + + isLast := index == len(src.Sources)-1 + if !isLast { + // PRE: bounded capReg + // $missing_amt = $cap - $pulled_amt + missingAmt := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{ + Op: ir.OpSubInt{}, + Left: *capReg, + Right: subPulledAmtReg, + Dest: dest, + } + }) + + st.jmpIfAmountZero(missingAmt, endLabel) + // this branch fell short: undo it and reopen for the next one. There + // is no rewind-without-closing, so a retry is a close plus a push — + // and after the rollback the new mark is identical to the closed one. + st.Push(ir.MarkEnd{Rewind: true}) + st.Push(ir.MarkPush{}) + } + } + + st.Push(ir.LabelMarker{Label: endLabel}) + // every path into endLabel — the jumps from a branch that covered the cap, + // and the fallthrough from the last branch — has exactly one region open, so + // a single commit here closes it once on all of them. Keeping it + // unconditional at the join is what keeps mark depth a function of position. + st.Push(ir.MarkEnd{Rewind: false}) + + return resultReg, nil + + case *parser.SourceAllotment: + // an allotment source splits the cap among sub-sources, so it needs one + if capReg == nil { + return 0, InvalidUncappedSource{Range: src.GetRange()} + } + allotments := make([]parser.AllotmentValue, len(src.Items)) + for i, item := range src.Items { + allotments[i] = item.Allotment + } + shares, err := st.compileAllot(*capReg, allotments) + if err != nil { + return 0, err + } + // pull exactly its share from each sub-source (tryTakingExact) + for i, item := range src.Items { + if _, err := st.compileSourceWithRequiredAmount(shares[i], item.From); err != nil { + return 0, err + } + } + return *capReg, nil + + case *parser.SourceWithScaling: + if err := st.checkFeatureFlag(src.GetRange(), flags.AssetScaling); err != nil { + return 0, err + } + return 0, FeatureNotImplemented{Range: src.GetRange(), Feature: "scaling"} + + default: + return 0, UnsupportedNode{Range: src.GetRange(), Node: src} + } +} + +// compileSourceWithRequiredAmount is the interpreter's tryTakingExact: pull up +// to capReg, then fail as missing funds unless exactly capReg was pulled. The +// cap handed to the source is clamped at zero first, as tryTakingUpTo does at +// entry. +func (st *state) compileSourceWithRequiredAmount( + capReg ir.Reg, + src parser.Source, +) (ir.Reg, CompilerError) { + zeroReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{Value: *big.NewInt(0), Dest: dest} + }) + clampedCapReg := st.maxInt(capReg, zeroReg) + got, err := st.compileSource(&clampedCapReg, src) + if err != nil { + return 0, err + } + st.Push(ir.CheckEnoughFunds{ + Got: got, + Needed: capReg, + }) + return got, nil +} + +func (st *state) compileDestination( + pulledAmtReg ir.Reg, + currentCap ir.Reg, + dest parser.Destination, +) CompilerError { + switch dest := dest.(type) { + case *parser.DestinationAllotment: + allotments := make([]parser.AllotmentValue, len(dest.Items)) + for i, item := range dest.Items { + allotments[i] = item.Allotment + } + // split the amount routed to this destination across the portions + shares, err := st.compileAllot(currentCap, allotments) + if err != nil { + return err + } + // send each computed share to its target (capped by that exact amount) + for i, item := range dest.Items { + if err := st.compileKeptOrDestination(item.To, pulledAmtReg, shares[i]); err != nil { + return err + } + } + return nil + + case *parser.DestinationOneof: + if err := st.checkFeatureFlag(dest.GetRange(), flags.ExperimentalOneofFeatureFlag); err != nil { + return err + } + + endLabel := st.FreshLabel("oneof_dest_end") + + clauseLabels := make([]ir.Label, len(dest.Clauses)) + for i, clause := range dest.Clauses { + clauseLabels[i] = st.FreshLabel("oneof_dest_clause") + + capAmtReg, err := st.compileCapAmount(clause.Cap) + if err != nil { + return err + } + minReg := st.minInt(currentCap, capAmtReg) + diff := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpSubInt{}, Left: currentCap, Right: minReg, Dest: dest} + }) + st.jmpIfAmountZero(diff, clauseLabels[i]) + } + + if err := st.compileKeptOrDestination(dest.Remaining, pulledAmtReg, currentCap); err != nil { + return err + } + st.Push(ir.Jmp{Target: endLabel}) + + for i, clause := range dest.Clauses { + st.Push(ir.LabelMarker{Label: clauseLabels[i]}) + if err := st.compileKeptOrDestination(clause.To, pulledAmtReg, currentCap); err != nil { + return err + } + st.Push(ir.Jmp{Target: endLabel}) + } + + st.Push(ir.LabelMarker{Label: endLabel}) + return nil + + case *parser.DestinationAccount: + acc, err := st.compileAccountExpr(dest.ValueExpr) + if err != nil { + return err + } + + var cap *ir.Reg + if pulledAmtReg != currentCap { + cap = ¤tCap + } + st.Push(ir.SendToAccount{ + Account: &acc.Name, + Cap: cap, + Scope: acc.Scope, + }) + + case *parser.DestinationInorder: + remaining := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpIntCopy{}, Arg: currentCap, Dest: dest} + }) + // mirrors internal/interpreter's sendTo, *parser.DestinationInorder case, + // laziness included: once remaining hits zero the loop breaks before + // evaluating the next clause's cap, and a clause (the trailing `remaining` + // one too) whose amount is zero has its destination never evaluated. + endLabel := st.FreshLabel("dest_inorder_end") + for _, clause := range dest.Clauses { + st.jmpIfAmountZero(remaining, endLabel) + + capAmtReg, err := st.compileCapAmount(clause.Cap) + if err != nil { + return err + } + // max(min(cap, remaining), 0): a negative `max` clause amount clamps + // to zero rather than erroring, matching the source-side `max ... from` + // clause and ledger's own behavior gap on this shape + // (oracle/DIVERGENCES.md #4). + zeroReg := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadInt{Value: *big.NewInt(0), Dest: dest} + }) + amtReg := st.maxInt(st.minInt(remaining, capAmtReg), zeroReg) + + skipLabel := st.FreshLabel("dest_inorder_skip") + st.jmpIfAmountZero(amtReg, skipLabel) + if err := st.compileKeptOrDestination(clause.To, pulledAmtReg, amtReg); err != nil { + return err + } + st.Push(ir.BinaryOp{Op: ir.OpSubInt{}, Dest: remaining, Left: remaining, Right: amtReg}) + st.Push(ir.LabelMarker{Label: skipLabel}) + } + st.Push(ir.LabelMarker{Label: endLabel}) + + remSkipLabel := st.FreshLabel("dest_inorder_rem_skip") + st.jmpIfAmountZero(remaining, remSkipLabel) + if err := st.compileKeptOrDestination(dest.Remaining, pulledAmtReg, remaining); err != nil { + return err + } + st.Push(ir.LabelMarker{Label: remSkipLabel}) + return nil + + default: + return UnsupportedNode{Range: dest.GetRange(), Node: dest} + } + + return nil +} + +func (st *state) compileKeptOrDestination( + keptOrDest parser.KeptOrDestination, + pulledAmtReg ir.Reg, + currentCap ir.Reg, +) CompilerError { + switch keptOrDest := keptOrDest.(type) { + case *parser.DestinationTo: + return st.compileDestination(pulledAmtReg, currentCap, keptOrDest.Destination) + + case *parser.DestinationKept: + var cap *ir.Reg + if pulledAmtReg != currentCap { + cap = ¤tCap + } + st.Push(ir.SendToAccount{ + Account: nil, + Cap: cap, + }) + return nil + + default: + return UnsupportedNode{Node: keptOrDest} + } +} + +func (st *state) compileSentValue( + sentValue parser.SentValue, + source parser.Source, +) (ir.Reg, CompilerError) { + switch sentValue := sentValue.(type) { + case *parser.SentValueLiteral: + mon, err := st.compileMonetaryExpr(sentValue.Monetary) + if err != nil { + return 0, err + } + st.Push(ir.AssertNonNegativeAmount{Amount: mon.Amount}) + st.Push(ir.SetCurrentAsset{ + Asset: mon.Asset, + }) + st.currentAssetReg = &mon.Asset + + return st.compileSourceWithRequiredAmount(mon.Amount, source) + + case *parser.SentValueAll: + assetReg, err := st.compileExpr(sentValue.Asset) + if err != nil { + return 0, err + } + st.Push(ir.SetCurrentAsset{ + Asset: assetReg, + }) + st.currentAssetReg = &assetReg + return st.compileSource(nil, source) + + default: + return 0, UnsupportedNode{Range: sentValue.GetRange(), Node: sentValue} + } + +} + +func (st *state) compileStatements(stmt parser.Statement) CompilerError { + switch stmt := stmt.(type) { + case *parser.SendStatement: + pulledAmtReg, err := st.compileSentValue(stmt.SentValue, stmt.Source) + if err != nil { + return err + } + + err = st.compileDestination(pulledAmtReg, pulledAmtReg, stmt.Destination) + if err != nil { + return err + } + + return nil + + case *parser.SaveStatement: + var assetReg ir.Reg + var amountReg *ir.Reg + switch sv := stmt.SentValue.(type) { + case *parser.SentValueLiteral: + mon, err := st.compileMonetaryExpr(sv.Monetary) + if err != nil { + return err + } + st.Push(ir.AssertNonNegativeAmount{Amount: mon.Amount}) + assetReg = mon.Asset + amountReg = &mon.Amount + case *parser.SentValueAll: + r, err := st.compileExpr(sv.Asset) + if err != nil { + return err + } + assetReg = r + default: + return UnsupportedNode{Range: stmt.SentValue.GetRange(), Node: stmt.SentValue} + } + + acc, err := st.compileAccountExpr(stmt.Account) + if err != nil { + return err + } + st.Push(ir.Save{Account: acc.Name, Asset: assetReg, Amount: amountReg, Scope: acc.Scope}) + return nil + case *parser.FnCall: + switch stmt.Caller.Name { + case builtins.SetTxMeta: + key, err := st.compileExpr(stmt.Args[0]) + if err != nil { + return err + } + value, err := st.compileMetaValue(stmt.Args[1]) + if err != nil { + return err + } + st.Push(ir.SetTxMeta{Key: key, Value: value}) + return nil + + case builtins.SetAccountMeta: + acc, err := st.compileAccountExpr(stmt.Args[0]) + if err != nil { + return err + } + key, err := st.compileExpr(stmt.Args[1]) + if err != nil { + return err + } + value, err := st.compileMetaValue(stmt.Args[2]) + if err != nil { + return err + } + st.Push(ir.SetAccountMeta{Account: acc.Name, Key: key, Value: value, Scope: acc.Scope}) + return nil + + default: + return UnsupportedNode{Range: stmt.Range, Node: stmt.Caller.Name} + } + + default: + return UnsupportedNode{Range: stmt.GetRange(), Node: stmt} + } +} + +// compileMetaValue compiles a value into a string register (metadata is stored +// stringified). Strings/accounts/assets already live in string registers; +// numbers go through int_to_string. +func (st *state) compileMetaValue(expr parser.ValueExpr) (ir.Reg, CompilerError) { + if st.exprTypes[expr] == typecheck.TypeMonetary { + mon, err := st.compileMonetaryExpr(expr) + if err != nil { + return 0, err + } + return st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{ + Op: ir.OpMonetaryToString{}, + Left: mon.Asset, + Right: mon.Amount, + Dest: dest, + } + }), nil + } + + // a scoped account can be the *subject* of a metadata write (compiled via + // compileAccountExpr elsewhere), but never the stored *value* — mirrors the + // interpreter's CannotStoreScopedAccountInMeta, checked here at compile time + // since scopedness is static. + if st.exprTypes[expr] == typecheck.TypeAccount { + acc, err := st.compileAccountExpr(expr) + if err != nil { + return 0, err + } + if acc.Scope != nil { + return 0, CannotStoreScopedAccountInMeta{Range: expr.GetRange()} + } + return acc.Name, nil + } + + r, err := st.compileExpr(expr) + if err != nil { + return 0, err + } + + switch st.exprTypes[expr] { + case typecheck.TypeString, typecheck.TypeAsset: + return r, nil + case typecheck.TypeNumber: + return st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpIntToString{}, Arg: r, Dest: dest} + }), nil + case typecheck.TypePortion: + return st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{Op: ir.OpPortionToString{}, Arg: r, Dest: dest} + }), nil + default: + return 0, CannotCastToString{Range: expr.GetRange(), Type: st.exprTypes[expr]} + } +} + +func compileProgramToIR(program parser.Program, featureFlags map[string]struct{}) (compiledProgramIR, CompilerError) { + tc := typecheck.Check(program) + if len(tc.Errors) > 0 { + return compiledProgramIR{}, TypeError{Range: tc.Errors[0].Range, Kind: tc.Errors[0].Kind} + } + + flagSet := maps.Clone(featureFlags) + if flagSet == nil { + flagSet = make(map[string]struct{}, len(program.Flags)) + } + for _, flag := range program.Flags { + if !slices.Contains(flags.AllFlags, flag.String) { + return compiledProgramIR{}, InvalidFeature{Range: flag.Range, Feature: flag.String} + } + flagSet[flag.String] = struct{}{} + } + + st := state{vars: map[string]value{}, exprTypes: tc.ExprTypes, featureFlags: flagSet} + + // loaded once, up front, so that it dominates every pullFromAccount branch + // regardless of the jumps those branches sit between + st.worldReg = st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadStr{Value: worldAccount, Dest: dest} + }) + + if program.Vars != nil { + for _, decl := range program.Vars.Declarations { + if err := st.compileVarDeclaration(decl); err != nil { + return compiledProgramIR{}, err + } + } + } + + for _, stmt := range program.Statements { + if err := st.compileStatements(stmt); err != nil { + return compiledProgramIR{}, err + } + } + + return compiledProgramIR{ + instructions: st.Instrs(), + varsEncoder: VarsEncoder{ + decls: st.varDecls, + nStr: st.nextStrVar, + nInt: st.nextIntVar, + }, + }, nil +} + +func (st *state) compileVarDeclaration(decl parser.VarDeclaration) CompilerError { + if decl.Origin == nil { + return st.compileExternalVar(decl) + } + if decl.Type.Name == typecheck.TypeMonetary { + if fnCall, ok := (*decl.Origin).(*parser.FnCall); ok { + if fnCall.Caller.Name == builtins.Meta { + return st.compileMetaVar(decl, fnCall) + } + mon, err := st.compileMonetaryFnCall(fnCall, true) + if err != nil { + return err + } + st.vars[decl.Name.Name] = monValue(mon) + return nil + } + mon, err := st.compileMonetaryExpr(*decl.Origin) + if err != nil { + return err + } + st.vars[decl.Name.Name] = monValue(mon) + return nil + } + + if decl.Type.Name == typecheck.TypeAccount { + if fnCall, ok := (*decl.Origin).(*parser.FnCall); ok { + // meta() can produce any declared type, account included (e.g. + // `account $seller = meta($sale, "seller")`) — a value read from + // metadata is always a plain, unscoped account name. + if fnCall.Caller.Name == builtins.Meta { + return st.compileMetaVar(decl, fnCall) + } + // scoped() is the only other builtin returning TypeAccount; a call + // that is the whole origin expression isn't a mid-script call. + acc, err := st.compileScopedFnCall(fnCall, true) + if err != nil { + return err + } + st.vars[decl.Name.Name] = accValue(acc) + return nil + } + acc, err := st.compileAccountExpr(*decl.Origin) + if err != nil { + return err + } + st.vars[decl.Name.Name] = accValue(acc) + return nil + } + + var r ir.Reg + var err CompilerError + if fnCall, ok := (*decl.Origin).(*parser.FnCall); ok { + // meta() is only supported as a variable origin, statically dispatched on + // the declared type; elsewhere compileFnCall reports InvalidMetaPosition. + if fnCall.Caller.Name == builtins.Meta { + return st.compileMetaVar(decl, fnCall) + } + // a call that is the whole origin expression isn't a mid-script call + r, err = st.compileFnCall(fnCall, true) + } else { + r, err = st.compileExpr(*decl.Origin) + } + if err != nil { + return err + } + st.vars[decl.Name.Name] = scalarValue(r) + return nil +} + +func (st *state) compileMetaVar(decl parser.VarDeclaration, fnCall *parser.FnCall) CompilerError { + acc, err := st.compileAccountExpr(fnCall.Args[0]) + if err != nil { + return err + } + key, err := st.compileExpr(fnCall.Args[1]) + if err != nil { + return err + } + + // monetary is the one meta type whose single store read yields two values, so + // it has its own two-destination instruction rather than a MetaType. + if decl.Type.Name == typecheck.TypeMonetary { + destAsset := st.FreshReg() + destAmount := st.FreshReg() + st.Push(ir.MetaMonetary{ + DestAsset: destAsset, + DestAmount: destAmount, + Account: acc.Name, + Key: key, + Scope: acc.Scope, + }) + st.vars[decl.Name.Name] = monValue(monetaryValue{Asset: destAsset, Amount: destAmount}) + return nil + } + + var typ ir.MetaType + switch decl.Type.Name { + case typecheck.TypeString, typecheck.TypeAccount, typecheck.TypeAsset: + typ = ir.MetaStr{} + case typecheck.TypeNumber: + typ = ir.MetaInt{} + case typecheck.TypePortion: + typ = ir.MetaPortion{} + default: + return UnsupportedNode{Range: decl.Type.Range, Node: decl.Type.Name} + } + + st.vars[decl.Name.Name] = scalarValue(st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.MetaVar{Dest: dest, Account: acc.Name, Key: key, Typ: typ, Scope: acc.Scope} + })) + return nil +} + +// TODO review AI blob +func (st *state) compileExternalVar(decl parser.VarDeclaration) CompilerError { + name := decl.Name.Name + st.varDecls = append(st.varDecls, varDecl{name: name, typ: decl.Type.Name}) + + switch decl.Type.Name { + case typecheck.TypeNumber: + st.vars[name] = scalarValue(st.loadIntVar()) + + case typecheck.TypeString, typecheck.TypeAsset, typecheck.TypeAccount: + st.vars[name] = scalarValue(st.loadStrVar()) + + case typecheck.TypePortion: + num := st.loadIntVar() + den := st.loadIntVar() + st.vars[name] = scalarValue(st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.BinaryOp{Op: ir.OpMakePortion{}, Left: num, Right: den, Dest: dest} + })) + + case typecheck.TypeMonetary: + // the vars payload already carries a monetary as two scalars, so the pair + // is the value — nothing to assemble + asset := st.loadStrVar() + amount := st.loadIntVar() + st.vars[name] = monValue(monetaryValue{Asset: asset, Amount: amount}) + + default: + return UnsupportedNode{Range: decl.Type.Range, Node: decl.Type.Name} + } + + return nil +} + +func (st *state) loadIntVar() ir.Reg { + index := uint16(st.nextIntVar) + st.nextIntVar++ + return st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadVar{Dest: dest, Typ: ir.VarInt{}, Index: index} + }) +} + +func (st *state) loadStrVar() ir.Reg { + index := uint16(st.nextStrVar) + st.nextStrVar++ + return st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.LoadVar{Dest: dest, Typ: ir.VarStr{}, Index: index} + }) +} diff --git a/internal/compiler/compiler_error.go b/internal/compiler/compiler_error.go new file mode 100644 index 00000000..a923f5c7 --- /dev/null +++ b/internal/compiler/compiler_error.go @@ -0,0 +1,164 @@ +package compiler + +import ( + "fmt" + + "github.com/formancehq/numscript/internal/flags" + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/typecheck" +) + +type ( + CompilerError interface { + parser.Ranged + compileError() + } + + UnboundVar struct { + parser.Range + Var string + } + + TypeError struct { + parser.Range + Kind typecheck.ErrorKind + } + + InvalidUncappedSource struct { + parser.Range + } + + DuplicateRemaining struct { + parser.Range + } + + // InvalidMetaPosition is reported when meta() appears anywhere other than as + // a top-level variable origin (the only place it's supported). + InvalidMetaPosition struct { + parser.Range + } + + // CannotCastToString is reported for an interpolation part whose type has no + // string form (monetary, asset, portion). + CannotCastToString struct { + parser.Range + Type typecheck.Type + } + + // CannotStoreScopedAccountInMeta is reported when a scoped account (the + // result of scoped()) is used as the *value* stored by set_tx_meta or + // set_account_meta. Mirrors the interpreter's runtime error of the same name, + // but caught at compile time since the compiler already knows an expression's + // scopedness statically. + CannotStoreScopedAccountInMeta struct { + parser.Range + } + + // InvalidScopedAccountPosition is reported when scoped() is reached from a + // position that only wants a plain string/generic value (e.g. account + // interpolation, or any other non-account context) — a defensive check that + // should be unreachable given the compiler's other call sites already route + // account-typed expressions through compileAccountExpr. + InvalidScopedAccountPosition struct { + parser.Range + } + + // FeatureNotImplemented is returned (never panicked) when the compiler meets a + // construct it does not support yet — e.g. colors or scoped accounts — so the + // host gets an error instead of a crash. + FeatureNotImplemented struct { + parser.Range + Feature string + } + + // ExperimentalFeature is reported when the script uses a construct gated + // behind a feature flag that wasn't enabled. Mirrors the interpreter's + // interpreter.ExperimentalFeature. + ExperimentalFeature struct { + parser.Range + FlagName flags.FeatureFlag + } + + // InvalidFeature is reported when a #![feature(..)] declaration names a flag + // that doesn't exist. + InvalidFeature struct { + parser.Range + Feature string + } + + // UnsupportedNode is the compiler's defensive backstop for a switch over an + // AST node, operator, or builtin name that's meant to be exhaustive (every + // case the parser/typecheck can currently produce is handled) but isn't + // enforced as such by the Go compiler — e.g. a switch on a string-typed + // operator, or a fixed set of builtin names. Reaching it means the compiler + // has drifted out of sync with the parser or typecheck, not something a + // script can trigger, so the message carries a debug dump rather than + // user-actionable text. + UnsupportedNode struct { + parser.Range + Node any + } +) + +func (UnboundVar) compileError() {} +func (TypeError) compileError() {} +func (InvalidUncappedSource) compileError() {} +func (DuplicateRemaining) compileError() {} +func (InvalidMetaPosition) compileError() {} +func (CannotCastToString) compileError() {} +func (CannotStoreScopedAccountInMeta) compileError() {} +func (InvalidScopedAccountPosition) compileError() {} +func (FeatureNotImplemented) compileError() {} +func (ExperimentalFeature) compileError() {} +func (InvalidFeature) compileError() {} +func (UnsupportedNode) compileError() {} + +func (e FeatureNotImplemented) Error() string { + return "internal error: feature not implemented: " + e.Feature +} +func (e UnboundVar) Error() string { + return fmt.Sprintf("the variable '$%s' was not declared", e.Var) +} +func (InvalidUncappedSource) Error() string { + return "cannot take all balance of an unbounded source" +} +func (DuplicateRemaining) Error() string { + return "only one 'remaining' clause is allowed in an allotment expression" +} +func (e TypeError) Error() string { return e.Kind.Message() } +func (InvalidMetaPosition) Error() string { + return "meta() is only allowed as a variable origin" +} +func (e CannotCastToString) Error() string { + return "cannot cast a value of type " + string(e.Type) + " to string" +} +func (CannotStoreScopedAccountInMeta) Error() string { + return "cannot store a scoped account as a metadata value" +} +func (InvalidScopedAccountPosition) Error() string { + return "a scoped account cannot be used here" +} +func (e ExperimentalFeature) Error() string { + return fmt.Sprintf("this feature is experimental. You need the '%s' feature flag to enable it", e.FlagName) +} +func (e InvalidFeature) Error() string { + return fmt.Sprintf("Invalid feature: %s", e.Feature) +} +func (e UnsupportedNode) Error() string { + return fmt.Sprintf("internal error: unsupported node %#v", e.Node) +} + +var ( + _ CompilerError = (*UnboundVar)(nil) + _ CompilerError = (*TypeError)(nil) + _ CompilerError = (*InvalidUncappedSource)(nil) + _ CompilerError = (*DuplicateRemaining)(nil) + _ CompilerError = (*InvalidMetaPosition)(nil) + _ CompilerError = (*CannotCastToString)(nil) + _ CompilerError = (*CannotStoreScopedAccountInMeta)(nil) + _ CompilerError = (*InvalidScopedAccountPosition)(nil) + _ CompilerError = (*FeatureNotImplemented)(nil) + _ CompilerError = (*ExperimentalFeature)(nil) + _ CompilerError = (*InvalidFeature)(nil) + _ CompilerError = (*UnsupportedNode)(nil) +) diff --git a/internal/compiler/compiler_example_test.go b/internal/compiler/compiler_example_test.go new file mode 100644 index 00000000..eb5559ac --- /dev/null +++ b/internal/compiler/compiler_example_test.go @@ -0,0 +1,85 @@ +package compiler_test + +import ( + "context" + "math/big" + "testing" + + "github.com/formancehq/numscript" + "github.com/stretchr/testify/require" +) + +func TestCompilerExample(t *testing.T) { + script := ` + vars { + account $acc + } + + send [USD/2 10] ( + source = $acc + destination = @dest + ) + ` + + varsEncoder, compiledProgram, compilationErr := numscript.Compile(script) + require.NoError(t, compilationErr) // e.g. parsing errors or type errors or any other kind of compile-time errors + + { + // The compiledProgram represents the compiled version of the program. + // We can serialize it into a []byte sequence and decode it back to the same data structure. + // the serialised []byte format is meant to be used to send it over the wire + bytecode := compiledProgram.Encode() // <- cast to []byte + + decodedCompiledProgram, decodingErr := numscript.DecodeCompiledProgram(bytecode) // <- decode it back + require.NoError(t, decodingErr) + require.Equal(t, decodedCompiledProgram, compiledProgram) + } + + // the vars encoder must be stored by the leader, so that it can encode the vars payload + // in a way that can be consumed by the vm + vars, err := varsEncoder.Encode(map[string]string{ + "acc": "src_account", + }) + require.NoError(t, err) + + { + // just like the compiledProgram. the Vars can be serialised and deserialised into/from []byte + serialisedVars := vars.Encode() // <- []byte to be sent over the wire from leader to nodes + + decodedVars, decodingErr := numscript.DecodeVars(serialisedVars) // <- turning []byte into Vars + require.NoError(t, decodingErr) + require.Equal(t, decodedVars, vars) + } + + // We can initialise the vm by passing the numscript.CompiledProgram value. + // Not only it's valid to re-use the same instance of the VM from many script runs, + // it's actually best to keep that in memory instead of the keeping the program and re-creating the vm each time + // this way we can avoid allocating/deallocating the registers and vm state each time + vm := numscript.NewVm(compiledProgram) + + // mock store (repr'd as map) + store := testStore{ + "src_account": 100, + } + + result, execErr := numscript.ExecVm(context.Background(), vm, &vars, store) + require.NoError(t, execErr) // e.g. missing funds, or any other runtime error + require.Equal(t, []numscript.Posting{ + { + Source: "src_account", + Destination: "dest", + Asset: "USD/2", + Amount: big.NewInt(10), + }, + }, result.Postings) +} + +type testStore map[string]int64 + +func (s testStore) GetBalance(ctx context.Context, account, scope, asset, color string) (*big.Int, error) { + return big.NewInt(s[account]), nil +} + +func (testStore) GetMetadata(ctx context.Context, account, scope, key string) (string, bool, error) { + return "", false, nil +} diff --git a/internal/compiler/compiler_test.go b/internal/compiler/compiler_test.go new file mode 100644 index 00000000..08d35d35 --- /dev/null +++ b/internal/compiler/compiler_test.go @@ -0,0 +1,981 @@ +package compiler + +import ( + "testing" + + "github.com/formancehq/numscript/internal/ir" + "github.com/formancehq/numscript/internal/parser" + "github.com/gkampitakis/go-snaps/snaps" + "github.com/stretchr/testify/require" +) + +func getCompiledOutput(t *testing.T, source string) string { + t.Helper() + program := parser.Parse(source) + require.Empty(t, program.Errors) + compiled, err := compileProgramToIR(program.Value, nil) + require.Nil(t, err) + + out := "\n" + ir.Dump(compiled.instructions) + + // every snapshot below doubles as a round-trip test of the textual format + instrs, errs := ir.Parse(out) + require.Empty(t, errs, "the dump does not parse back") + require.Equal(t, out, "\n"+ir.Dump(instrs), "the dump does not round-trip") + + return out +} + +func TestSimpleProgram(t *testing.T) { + out := getCompiledOutput(t, ` + send [USD/2 10] ( + source = @src + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 10 + assert_non_negative_amount($r2) + set_current_asset($r1) + $r3 = 0 + $r4 = int_copy($r2) + $r5 = lt_int($r3, $r2) + jmp_if_true($r5, #max_end_0) + $r4 = int_copy($r3) +#max_end_0 + $r6 = "src" + $r7 = 0 + $r8 = str_eq($r6, $r0) + jmp_if_false($r8, #not_world_1) + $r9 = pull_account(account: $r6, cap: $r4) + jmp(#pull_end_2) +#not_world_1 + $r9 = pull_account(account: $r6, cap: $r4, overdraft: $r7) +#pull_end_2 + check_enough_funds($r9, $r2) + $r10 = "dest" + send_to_account(account: $r10) +`)) +} + +func TestIntAddition(t *testing.T) { + out := getCompiledOutput(t, ` + send [USD/2 4 + 6] ( + source = @src + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 4 + $r3 = 6 + $r4 = $r2 + $r3 + assert_non_negative_amount($r4) + set_current_asset($r1) + $r5 = 0 + $r6 = int_copy($r4) + $r7 = lt_int($r5, $r4) + jmp_if_true($r7, #max_end_0) + $r6 = int_copy($r5) +#max_end_0 + $r8 = "src" + $r9 = 0 + $r10 = str_eq($r8, $r0) + jmp_if_false($r10, #not_world_1) + $r11 = pull_account(account: $r8, cap: $r6) + jmp(#pull_end_2) +#not_world_1 + $r11 = pull_account(account: $r8, cap: $r6, overdraft: $r9) +#pull_end_2 + check_enough_funds($r11, $r4) + $r12 = "dest" + send_to_account(account: $r12) +`)) +} + +func TestIntSubtraction(t *testing.T) { + out := getCompiledOutput(t, ` + send [USD/2 16 - 6] ( + source = @src + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 16 + $r3 = 6 + $r4 = $r2 - $r3 + assert_non_negative_amount($r4) + set_current_asset($r1) + $r5 = 0 + $r6 = int_copy($r4) + $r7 = lt_int($r5, $r4) + jmp_if_true($r7, #max_end_0) + $r6 = int_copy($r5) +#max_end_0 + $r8 = "src" + $r9 = 0 + $r10 = str_eq($r8, $r0) + jmp_if_false($r10, #not_world_1) + $r11 = pull_account(account: $r8, cap: $r6) + jmp(#pull_end_2) +#not_world_1 + $r11 = pull_account(account: $r8, cap: $r6, overdraft: $r9) +#pull_end_2 + check_enough_funds($r11, $r4) + $r12 = "dest" + send_to_account(account: $r12) +`)) +} + +func TestMonetaryAddition(t *testing.T) { + out := getCompiledOutput(t, ` + vars { + monetary $a = [USD/2 3] + monetary $b = [USD/2 7] + } + send $a + $b ( + source = @src + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 3 + $r3 = "USD/2" + $r4 = 7 + assert_same_asset($r1, $r3) + $r5 = $r2 + $r4 + assert_non_negative_amount($r5) + set_current_asset($r1) + $r6 = 0 + $r7 = int_copy($r5) + $r8 = lt_int($r6, $r5) + jmp_if_true($r8, #max_end_0) + $r7 = int_copy($r6) +#max_end_0 + $r9 = "src" + $r10 = 0 + $r11 = str_eq($r9, $r0) + jmp_if_false($r11, #not_world_1) + $r12 = pull_account(account: $r9, cap: $r7) + jmp(#pull_end_2) +#not_world_1 + $r12 = pull_account(account: $r9, cap: $r7, overdraft: $r10) +#pull_end_2 + check_enough_funds($r12, $r5) + $r13 = "dest" + send_to_account(account: $r13) +`)) +} + +func TestMonetarySubtraction(t *testing.T) { + out := getCompiledOutput(t, ` + vars { + monetary $a = [USD/2 30] + monetary $b = [USD/2 20] + } + send $a - $b ( + source = @src + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 30 + $r3 = "USD/2" + $r4 = 20 + assert_same_asset($r1, $r3) + $r5 = $r2 - $r4 + assert_non_negative_amount($r5) + set_current_asset($r1) + $r6 = 0 + $r7 = int_copy($r5) + $r8 = lt_int($r6, $r5) + jmp_if_true($r8, #max_end_0) + $r7 = int_copy($r6) +#max_end_0 + $r9 = "src" + $r10 = 0 + $r11 = str_eq($r9, $r0) + jmp_if_false($r11, #not_world_1) + $r12 = pull_account(account: $r9, cap: $r7) + jmp(#pull_end_2) +#not_world_1 + $r12 = pull_account(account: $r9, cap: $r7, overdraft: $r10) +#pull_end_2 + check_enough_funds($r12, $r5) + $r13 = "dest" + send_to_account(account: $r13) +`)) +} + +func TestGetAmount(t *testing.T) { + out := getCompiledOutput(t, ` + #![feature("experimental-get-amount-function")] + vars { + monetary $m = [USD/2 42] + number $n = get_amount($m) + } + send [USD/2 $n] ( + source = @src + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 42 + $r3 = "USD/2" + assert_non_negative_amount($r2) + set_current_asset($r3) + $r4 = 0 + $r5 = int_copy($r2) + $r6 = lt_int($r4, $r2) + jmp_if_true($r6, #max_end_0) + $r5 = int_copy($r4) +#max_end_0 + $r7 = "src" + $r8 = 0 + $r9 = str_eq($r7, $r0) + jmp_if_false($r9, #not_world_1) + $r10 = pull_account(account: $r7, cap: $r5) + jmp(#pull_end_2) +#not_world_1 + $r10 = pull_account(account: $r7, cap: $r5, overdraft: $r8) +#pull_end_2 + check_enough_funds($r10, $r2) + $r11 = "dest" + send_to_account(account: $r11) +`)) +} + +func TestGetAsset(t *testing.T) { + out := getCompiledOutput(t, ` + #![feature("experimental-get-asset-function")] + vars { + monetary $m = [USD/2 42] + asset $a = get_asset($m) + } + send [$a 10] ( + source = @src + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 42 + $r3 = 10 + assert_non_negative_amount($r3) + set_current_asset($r1) + $r4 = 0 + $r5 = int_copy($r3) + $r6 = lt_int($r4, $r3) + jmp_if_true($r6, #max_end_0) + $r5 = int_copy($r4) +#max_end_0 + $r7 = "src" + $r8 = 0 + $r9 = str_eq($r7, $r0) + jmp_if_false($r9, #not_world_1) + $r10 = pull_account(account: $r7, cap: $r5) + jmp(#pull_end_2) +#not_world_1 + $r10 = pull_account(account: $r7, cap: $r5, overdraft: $r8) +#pull_end_2 + check_enough_funds($r10, $r3) + $r11 = "dest" + send_to_account(account: $r11) +`)) +} + +func TestPrefixMinusMonetary(t *testing.T) { + out := getCompiledOutput(t, ` + vars { + monetary $neg_mon = [USD/2 -10] + monetary $pos_mon = -$neg_mon + } + send $pos_mon ( + source = @src + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 10 + $r3 = neg_int($r2) + $r4 = neg_int($r3) + assert_non_negative_amount($r4) + set_current_asset($r1) + $r5 = 0 + $r6 = int_copy($r4) + $r7 = lt_int($r5, $r4) + jmp_if_true($r7, #max_end_0) + $r6 = int_copy($r5) +#max_end_0 + $r8 = "src" + $r9 = 0 + $r10 = str_eq($r8, $r0) + jmp_if_false($r10, #not_world_1) + $r11 = pull_account(account: $r8, cap: $r6) + jmp(#pull_end_2) +#not_world_1 + $r11 = pull_account(account: $r8, cap: $r6, overdraft: $r9) +#pull_end_2 + check_enough_funds($r11, $r4) + $r12 = "dest" + send_to_account(account: $r12) +`)) +} + +func TestBalance(t *testing.T) { + out := getCompiledOutput(t, ` + vars { + monetary $bal = balance(@src, USD/2) + } + send $bal ( + source = @src + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "src" + $r2 = "USD/2" + $r3 = balance($r1, $r2) + assert_non_negative_balance($r3, $r1) + assert_non_negative_amount($r3) + set_current_asset($r2) + $r4 = 0 + $r5 = int_copy($r3) + $r6 = lt_int($r4, $r3) + jmp_if_true($r6, #max_end_0) + $r5 = int_copy($r4) +#max_end_0 + $r7 = "src" + $r8 = 0 + $r9 = str_eq($r7, $r0) + jmp_if_false($r9, #not_world_1) + $r10 = pull_account(account: $r7, cap: $r5) + jmp(#pull_end_2) +#not_world_1 + $r10 = pull_account(account: $r7, cap: $r5, overdraft: $r8) +#pull_end_2 + check_enough_funds($r10, $r3) + $r11 = "dest" + send_to_account(account: $r11) +`)) +} + +func TestAccountInterpolation(t *testing.T) { + out := getCompiledOutput(t, ` + #![feature("experimental-account-interpolation")] + vars { + string $id = "alice" + } + send [USD/2 10] ( + source = @world + destination = @users:$id:wallet + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "alice" + $r2 = "USD/2" + $r3 = 10 + assert_non_negative_amount($r3) + set_current_asset($r2) + $r4 = 0 + $r5 = int_copy($r3) + $r6 = lt_int($r4, $r3) + jmp_if_true($r6, #max_end_0) + $r5 = int_copy($r4) +#max_end_0 + $r7 = "world" + $r8 = 0 + $r9 = str_eq($r7, $r0) + jmp_if_false($r9, #not_world_1) + $r10 = pull_account(account: $r7, cap: $r5) + jmp(#pull_end_2) +#not_world_1 + $r10 = pull_account(account: $r7, cap: $r5, overdraft: $r8) +#pull_end_2 + check_enough_funds($r10, $r3) + $r11 = "users" + $r12 = ":" + $r13 = ":" + $r14 = "wallet" + $r15 = add_string($r11, $r12) + $r16 = add_string($r15, $r1) + $r17 = add_string($r16, $r13) + $r18 = add_string($r17, $r14) + assert_valid_account($r18) + send_to_account(account: $r18) +`)) +} + +func TestAccountInterpolationInt(t *testing.T) { + out := getCompiledOutput(t, ` + #![feature("experimental-account-interpolation")] + vars { + number $n = 42 + } + send [USD/2 10] ( + source = @world + destination = @account:$n + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = 42 + $r2 = "USD/2" + $r3 = 10 + assert_non_negative_amount($r3) + set_current_asset($r2) + $r4 = 0 + $r5 = int_copy($r3) + $r6 = lt_int($r4, $r3) + jmp_if_true($r6, #max_end_0) + $r5 = int_copy($r4) +#max_end_0 + $r7 = "world" + $r8 = 0 + $r9 = str_eq($r7, $r0) + jmp_if_false($r9, #not_world_1) + $r10 = pull_account(account: $r7, cap: $r5) + jmp(#pull_end_2) +#not_world_1 + $r10 = pull_account(account: $r7, cap: $r5, overdraft: $r8) +#pull_end_2 + check_enough_funds($r10, $r3) + $r11 = "account" + $r12 = ":" + $r13 = int_to_string($r1) + $r14 = add_string($r11, $r12) + $r15 = add_string($r14, $r13) + assert_valid_account($r15) + send_to_account(account: $r15) +`)) +} + +func TestInorder(t *testing.T) { + out := getCompiledOutput(t, ` + send [USD/2 10] ( + source = { + @a + @b + @c + } + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 10 + assert_non_negative_amount($r2) + set_current_asset($r1) + $r3 = 0 + $r4 = int_copy($r2) + $r5 = lt_int($r3, $r2) + jmp_if_true($r5, #max_end_0) + $r4 = int_copy($r3) +#max_end_0 + $r6 = 0 + $r7 = int_copy($r4) + $r8 = "a" + $r9 = 0 + $r10 = str_eq($r8, $r0) + jmp_if_false($r10, #not_world_1) + $r11 = pull_account(account: $r8, cap: $r7) + jmp(#pull_end_2) +#not_world_1 + $r11 = pull_account(account: $r8, cap: $r7, overdraft: $r9) +#pull_end_2 + $r6 += $r11 + $r7 -= $r11 + $r12 = "b" + $r13 = 0 + $r14 = str_eq($r12, $r0) + jmp_if_false($r14, #not_world_3) + $r15 = pull_account(account: $r12, cap: $r7) + jmp(#pull_end_4) +#not_world_3 + $r15 = pull_account(account: $r12, cap: $r7, overdraft: $r13) +#pull_end_4 + $r6 += $r15 + $r7 -= $r15 + $r16 = "c" + $r17 = 0 + $r18 = str_eq($r16, $r0) + jmp_if_false($r18, #not_world_5) + $r19 = pull_account(account: $r16, cap: $r7) + jmp(#pull_end_6) +#not_world_5 + $r19 = pull_account(account: $r16, cap: $r7, overdraft: $r17) +#pull_end_6 + $r6 += $r19 + check_enough_funds($r6, $r2) + $r20 = "dest" + send_to_account(account: $r20) +`)) +} + +func TestInorderWithCap(t *testing.T) { + out := getCompiledOutput(t, ` + send [USD/2 10] ( + source = { + @a + max [USD/2 5] from @b + @c + } + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 10 + assert_non_negative_amount($r2) + set_current_asset($r1) + $r3 = 0 + $r4 = int_copy($r2) + $r5 = lt_int($r3, $r2) + jmp_if_true($r5, #max_end_0) + $r4 = int_copy($r3) +#max_end_0 + $r6 = 0 + $r7 = int_copy($r4) + $r8 = "a" + $r9 = 0 + $r10 = str_eq($r8, $r0) + jmp_if_false($r10, #not_world_1) + $r11 = pull_account(account: $r8, cap: $r7) + jmp(#pull_end_2) +#not_world_1 + $r11 = pull_account(account: $r8, cap: $r7, overdraft: $r9) +#pull_end_2 + $r6 += $r11 + $r7 -= $r11 + $r12 = "USD/2" + $r13 = 5 + assert_same_asset($r1, $r12) + $r14 = int_copy($r13) + $r15 = lt_int($r13, $r7) + jmp_if_true($r15, #min_end_3) + $r14 = int_copy($r7) +#min_end_3 + $r16 = 0 + $r17 = int_copy($r14) + $r18 = lt_int($r16, $r14) + jmp_if_true($r18, #max_end_4) + $r17 = int_copy($r16) +#max_end_4 + $r19 = "b" + $r20 = 0 + $r21 = str_eq($r19, $r0) + jmp_if_false($r21, #not_world_5) + $r22 = pull_account(account: $r19, cap: $r17) + jmp(#pull_end_6) +#not_world_5 + $r22 = pull_account(account: $r19, cap: $r17, overdraft: $r20) +#pull_end_6 + $r6 += $r22 + $r7 -= $r22 + $r23 = "c" + $r24 = 0 + $r25 = str_eq($r23, $r0) + jmp_if_false($r25, #not_world_7) + $r26 = pull_account(account: $r23, cap: $r7) + jmp(#pull_end_8) +#not_world_7 + $r26 = pull_account(account: $r23, cap: $r7, overdraft: $r24) +#pull_end_8 + $r6 += $r26 + check_enough_funds($r6, $r2) + $r27 = "dest" + send_to_account(account: $r27) +`)) +} + +func TestDestInorder(t *testing.T) { + out := getCompiledOutput(t, ` + send [USD/2 10] ( + source = @world + destination = { + max [USD/2 4] to @d1 + remaining to @d2 + } + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 10 + assert_non_negative_amount($r2) + set_current_asset($r1) + $r3 = 0 + $r4 = int_copy($r2) + $r5 = lt_int($r3, $r2) + jmp_if_true($r5, #max_end_0) + $r4 = int_copy($r3) +#max_end_0 + $r6 = "world" + $r7 = 0 + $r8 = str_eq($r6, $r0) + jmp_if_false($r8, #not_world_1) + $r9 = pull_account(account: $r6, cap: $r4) + jmp(#pull_end_2) +#not_world_1 + $r9 = pull_account(account: $r6, cap: $r4, overdraft: $r7) +#pull_end_2 + check_enough_funds($r9, $r2) + $r10 = int_copy($r9) + $r11 = is_zero($r10) + jmp_if_true($r11, #dest_inorder_end_3) + $r12 = "USD/2" + $r13 = 4 + assert_same_asset($r1, $r12) + $r14 = 0 + $r15 = int_copy($r10) + $r16 = lt_int($r10, $r13) + jmp_if_true($r16, #min_end_4) + $r15 = int_copy($r13) +#min_end_4 + $r17 = int_copy($r15) + $r18 = lt_int($r14, $r15) + jmp_if_true($r18, #max_end_5) + $r17 = int_copy($r14) +#max_end_5 + $r19 = is_zero($r17) + jmp_if_true($r19, #dest_inorder_skip_6) + $r20 = "d1" + send_to_account(account: $r20, cap: $r17) + $r10 -= $r17 +#dest_inorder_skip_6 +#dest_inorder_end_3 + $r21 = is_zero($r10) + jmp_if_true($r21, #dest_inorder_rem_skip_7) + $r22 = "d2" + send_to_account(account: $r22, cap: $r10) +#dest_inorder_rem_skip_7 +`)) +} + +func TestSourceOneofSimple(t *testing.T) { + out := getCompiledOutput(t, ` + #![feature("experimental-oneof")] + send [USD/2 10] ( + source = oneof { + @a + @b + @c + } + destination = @dest + ) + `) + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 10 + assert_non_negative_amount($r2) + set_current_asset($r1) + $r3 = 0 + $r4 = int_copy($r2) + $r5 = lt_int($r3, $r2) + jmp_if_true($r5, #max_end_0) + $r4 = int_copy($r3) +#max_end_0 + mark_push() + $r6 = "a" + $r7 = 0 + $r8 = str_eq($r6, $r0) + jmp_if_false($r8, #not_world_2) + $r9 = pull_account(account: $r6, cap: $r4) + jmp(#pull_end_3) +#not_world_2 + $r9 = pull_account(account: $r6, cap: $r4, overdraft: $r7) +#pull_end_3 + $r10 = int_copy($r9) + $r11 = $r4 - $r9 + $r12 = is_zero($r11) + jmp_if_true($r12, #oneof_end_1) + mark_rewind() + mark_push() + $r13 = "b" + $r14 = 0 + $r15 = str_eq($r13, $r0) + jmp_if_false($r15, #not_world_4) + $r16 = pull_account(account: $r13, cap: $r4) + jmp(#pull_end_5) +#not_world_4 + $r16 = pull_account(account: $r13, cap: $r4, overdraft: $r14) +#pull_end_5 + $r10 = int_copy($r16) + $r17 = $r4 - $r16 + $r18 = is_zero($r17) + jmp_if_true($r18, #oneof_end_1) + mark_rewind() + mark_push() + $r19 = "c" + $r20 = 0 + $r21 = str_eq($r19, $r0) + jmp_if_false($r21, #not_world_6) + $r22 = pull_account(account: $r19, cap: $r4) + jmp(#pull_end_7) +#not_world_6 + $r22 = pull_account(account: $r19, cap: $r4, overdraft: $r20) +#pull_end_7 + $r10 = int_copy($r22) +#oneof_end_1 + mark_commit() + check_enough_funds($r10, $r2) + $r23 = "dest" + send_to_account(account: $r23) +`)) +} + +func TestSourceOneofBounded(t *testing.T) { + + out := getCompiledOutput(t, ` + #![feature("experimental-oneof")] + send [USD/2 10] ( + source = oneof { + @a + @b + } + destination = @dest + ) + `) + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 10 + assert_non_negative_amount($r2) + set_current_asset($r1) + $r3 = 0 + $r4 = int_copy($r2) + $r5 = lt_int($r3, $r2) + jmp_if_true($r5, #max_end_0) + $r4 = int_copy($r3) +#max_end_0 + mark_push() + $r6 = "a" + $r7 = 0 + $r8 = str_eq($r6, $r0) + jmp_if_false($r8, #not_world_2) + $r9 = pull_account(account: $r6, cap: $r4) + jmp(#pull_end_3) +#not_world_2 + $r9 = pull_account(account: $r6, cap: $r4, overdraft: $r7) +#pull_end_3 + $r10 = int_copy($r9) + $r11 = $r4 - $r9 + $r12 = is_zero($r11) + jmp_if_true($r12, #oneof_end_1) + mark_rewind() + mark_push() + $r13 = "b" + $r14 = 0 + $r15 = str_eq($r13, $r0) + jmp_if_false($r15, #not_world_4) + $r16 = pull_account(account: $r13, cap: $r4) + jmp(#pull_end_5) +#not_world_4 + $r16 = pull_account(account: $r13, cap: $r4, overdraft: $r14) +#pull_end_5 + $r10 = int_copy($r16) +#oneof_end_1 + mark_commit() + check_enough_funds($r10, $r2) + $r17 = "dest" + send_to_account(account: $r17) +`)) +} + +func TestDestOneof(t *testing.T) { + out := getCompiledOutput(t, ` + #![feature("experimental-oneof")] + send [USD/2 10] ( + source = @world + destination = oneof { + max [USD/2 4] to @a + remaining to @b + } + ) + `) + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "USD/2" + $r2 = 10 + assert_non_negative_amount($r2) + set_current_asset($r1) + $r3 = 0 + $r4 = int_copy($r2) + $r5 = lt_int($r3, $r2) + jmp_if_true($r5, #max_end_0) + $r4 = int_copy($r3) +#max_end_0 + $r6 = "world" + $r7 = 0 + $r8 = str_eq($r6, $r0) + jmp_if_false($r8, #not_world_1) + $r9 = pull_account(account: $r6, cap: $r4) + jmp(#pull_end_2) +#not_world_1 + $r9 = pull_account(account: $r6, cap: $r4, overdraft: $r7) +#pull_end_2 + check_enough_funds($r9, $r2) + $r10 = "USD/2" + $r11 = 4 + assert_same_asset($r1, $r10) + $r12 = int_copy($r9) + $r13 = lt_int($r9, $r11) + jmp_if_true($r13, #min_end_5) + $r12 = int_copy($r11) +#min_end_5 + $r14 = $r9 - $r12 + $r15 = is_zero($r14) + jmp_if_true($r15, #oneof_dest_clause_4) + $r16 = "b" + send_to_account(account: $r16) + jmp(#oneof_dest_end_3) +#oneof_dest_clause_4 + $r17 = "a" + send_to_account(account: $r17) + jmp(#oneof_dest_end_3) +#oneof_dest_end_3 +`)) +} + +func TestColoredSource(t *testing.T) { + out := getCompiledOutput(t, ` + #![feature("experimental-asset-colors")] + send [COIN 10] ( + source = @src \ "RED" + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "COIN" + $r2 = 10 + assert_non_negative_amount($r2) + set_current_asset($r1) + $r3 = 0 + $r4 = int_copy($r2) + $r5 = lt_int($r3, $r2) + jmp_if_true($r5, #max_end_0) + $r4 = int_copy($r3) +#max_end_0 + $r6 = "src" + $r7 = "RED" + assert_valid_color($r7) + $r8 = 0 + $r9 = str_eq($r6, $r0) + jmp_if_false($r9, #not_world_1) + $r10 = pull_account(account: $r6, cap: $r4, color: $r7) + jmp(#pull_end_2) +#not_world_1 + $r10 = pull_account(account: $r6, cap: $r4, overdraft: $r8, color: $r7) +#pull_end_2 + check_enough_funds($r10, $r2) + $r11 = "dest" + send_to_account(account: $r11) +`)) +} + +func TestColoredOverdraftSource(t *testing.T) { + out := getCompiledOutput(t, ` + #![feature("experimental-asset-colors")] + send [COIN 10] ( + source = @src \ "RED" allowing unbounded overdraft + destination = @dest + ) + `) + + snaps.MatchInlineSnapshot(t, out, snaps.Inline(` + $r0 = "world" + $r1 = "COIN" + $r2 = 10 + assert_non_negative_amount($r2) + set_current_asset($r1) + $r3 = 0 + $r4 = int_copy($r2) + $r5 = lt_int($r3, $r2) + jmp_if_true($r5, #max_end_0) + $r4 = int_copy($r3) +#max_end_0 + $r6 = "src" + $r7 = "RED" + assert_valid_color($r7) + $r8 = pull_account(account: $r6, cap: $r4, color: $r7) + check_enough_funds($r8, $r2) + $r9 = "dest" + send_to_account(account: $r9) +`)) +} + +// TestCompileIsDeterministic: compiling the same script repeatedly yields +// byte-identical bytecode. Several registers die on the same instruction here +// (each ADD is the last use of both operands), which is where the register +// allocator used to push freed slots in map order and assemble a different — +// equivalent — program on each run. +func TestCompileIsDeterministic(t *testing.T) { + t.Parallel() + + script := `vars { + number $a + number $b + number $c + number $d + number $e + number $f + monetary $x + monetary $y +} + +send [COIN ($a + $b) + ($c + $d) + ($e + $f)] ( + source = @world + destination = @dst +) + +send $x (source = @world destination = @a) +send $y (source = @world destination = @b) +` + parsed := parser.Parse(script) + require.Empty(t, parsed.Errors) + + _, first, err := Compile(parsed.Value, nil) + require.NoError(t, err) + want := first.Encode() + + for range 200 { + _, program, err := Compile(parsed.Value, nil) + require.NoError(t, err) + require.Equal(t, want, program.Encode()) + } +} diff --git a/internal/compiler/e2e_test.go b/internal/compiler/e2e_test.go new file mode 100644 index 00000000..7f9af02a --- /dev/null +++ b/internal/compiler/e2e_test.go @@ -0,0 +1,383 @@ +package compiler_test + +import ( + "context" + "math/big" + "testing" + + "github.com/formancehq/numscript/internal/compiler" + "github.com/formancehq/numscript/internal/funds" + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/vm" + "github.com/stretchr/testify/require" +) + +// e2eStore is a minimal vm.Store for the end-to-end test. +type e2eStore struct { + balances map[funds.PairKey]*big.Int + metadata map[e2eMetaKey]string +} + +// e2eMetaKey identifies one metadata slot: account, scope and key. +type e2eMetaKey struct { + account string + scope string + key string +} + +func (s e2eStore) GetBalance(ctx context.Context, account, scope, asset, color string) (*big.Int, error) { + if v, ok := s.balances[funds.PairKey{Account: account, Scope: scope, Asset: asset, Color: color}]; ok { + return v, nil + } + return new(big.Int), nil +} + +func (s e2eStore) GetMetadata(ctx context.Context, account, scope, key string) (string, bool, error) { + v, ok := s.metadata[e2eMetaKey{account: account, scope: scope, key: key}] + return v, ok, nil +} + +// TestE2E_AllotmentOverSum: portions summing to > 1 must error (leftover < 0). +func TestE2E_AllotmentOverSum(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + 2/3 to @a + 2/3 to @b + } + ) + ` + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + require.Equal(t, vm.InvalidAllotmentSum{ActualSum: *big.NewRat(4, 3)}, execErr) +} + +func TestE2E_NegativeAllotmentPortion(t *testing.T) { + for name, src := range map[string]string{ + "source": ` + send [USD/2 90] ( + source = { + -1/3 from @s1 + remaining from @s2 + } + destination = @dest + ) + `, + "destination": ` + send [USD/2 90] ( + source = @world + destination = { + -1/3 to @a + remaining to @b + } + ) + `, + "portions summing to one": ` + send [USD/2 90] ( + source = @world + destination = { + 4/3 to @a + -1/3 to @b + } + ) + `, + } { + t.Run(name, func(t *testing.T) { + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "s1", Asset: "USD/2", Color: ""}: big.NewInt(500), + {Account: "s2", Asset: "USD/2", Color: ""}: big.NewInt(500), + }}) + require.Equal(t, vm.NegativePortionError{Portion: *big.NewRat(-1, 3)}, execErr) + }) + } +} + +// TestE2E_AllotmentUnderSum: without a `remaining` clause the portions must sum +// to exactly 1, so 1/3 + 1/3 = 2/3 must error. +func TestE2E_AllotmentUnderSum(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + 1/3 to @a + 1/3 to @b + } + ) + ` + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + require.Equal(t, vm.InvalidAllotmentSum{ActualSum: *big.NewRat(2, 3)}, execErr) +} + +func TestE2E_MonetarySubtractionAssetMismatch(t *testing.T) { + src := ` + vars { + monetary $a = [USD/2 30] + monetary $b = [EUR/2 20] + } + send $a - $b ( + source = @src + destination = @dest + ) + ` + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + _, execErr := vm.Exec(context.Background(), vm.NewVm(program), nil, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}) + require.Equal(t, vm.AssetMismatchError{Expected: "USD/2", Got: "EUR/2"}, execErr) +} + +func TestE2E_MonetaryAdditionAssetMismatch(t *testing.T) { + src := ` + vars { + monetary $a = [USD/2 3] + monetary $b = [EUR/2 7] + } + send $a + $b ( + source = @src + destination = @dest + ) + ` + + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + + store := e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(100), + }} + + _, execErr := vm.Exec(context.Background(), vm.NewVm(program), nil, store) + require.Equal(t, vm.AssetMismatchError{Expected: "USD/2", Got: "EUR/2"}, execErr) +} + +// not a spec fixture: specs can only expect missing funds or a negative amount. +// The interpreter rejects this with CannotCastScopedAccountToString. +func TestE2E_AccountInterpolationOfScopedAccountIsRejected(t *testing.T) { + src := ` + #![feature("experimental-account-interpolation", "experimental-scoped-function")] + vars { + account $a = scoped(@src, "reserve") + } + send [USD/2 10] ( + source = @world + destination = @dest:$a + ) + ` + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + _, execErr := vm.Exec(context.Background(), vm.NewVm(program), nil, e2eStore{}) + require.Equal(t, vm.CannotCastScopedAccountToString{Account: "src", Scope: "reserve"}, execErr) +} + +func TestE2E_CapAssetMismatch(t *testing.T) { + src := ` + send [USD/2 100] ( + source = max [EUR/2 5] from @a + destination = @dest + ) + ` + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + require.Equal(t, vm.AssetMismatchError{Expected: "USD/2", Got: "EUR/2"}, execErr) +} + +func TestE2E_OverdraftAssetMismatch(t *testing.T) { + src := ` + send [USD/2 42] ( + source = @a allowing overdraft up to [EUR/2 5] + destination = @dest + ) + ` + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + require.Equal(t, vm.AssetMismatchError{Expected: "USD/2", Got: "EUR/2"}, execErr) +} + +func TestE2E_BalanceNegativeErrors(t *testing.T) { + src := ` + vars { monetary $b = balance(@acc, USD/2) } + send $b (source = @world destination = @dest) + ` + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "acc", Asset: "USD/2", Color: ""}: big.NewInt(-1), + }}) + require.Equal(t, vm.NegativeBalanceError{Account: "acc", Amount: *big.NewInt(-1)}, execErr) +} + +func TestE2E_DivideByZero(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + 1/0 to @a + remaining kept + } + ) + ` + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + require.Equal(t, vm.DivideByZeroError{Numerator: *big.NewInt(1)}, execErr) +} + +func TestE2E_InvalidColor(t *testing.T) { + src := ` + #![feature("experimental-asset-colors")] + send [COIN 10] ( + source = @src \ "not a color" + destination = @dest + ) + ` + + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, e2eStore{}) + require.Equal(t, vm.InvalidColor{Color: "not a color"}, execErr) +} + +// countingStore is an e2eStore that records how many balances it was asked for. +type countingStore struct { + e2eStore + balanceCalls int +} + +func (s *countingStore) GetBalance(ctx context.Context, account, scope, asset, color string) (*big.Int, error) { + s.balanceCalls++ + return s.e2eStore.GetBalance(ctx, account, scope, asset, color) +} + +// The compiled world arm has no overdraft operand, which is what makes the pull +// unbounded and therefore free of Store round-trips. numscript_test.go asserts +// the same for the interpreter. +func TestE2E_WorldSourceReadsNoBalance(t *testing.T) { + src := `send [USD/2 100] (source = @world destination = @dest)` + + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + store := &countingStore{} + machine := vm.NewVm(program) + res, execErr := vm.Exec(context.Background(), machine, nil, store) + require.Nil(t, execErr) + + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(100)}, + }, res.Postings) + require.Zero(t, store.balanceCalls, "a world source must not read any balance") +} + +// the run-time branch, not the literal, is what decides it +func TestE2E_DynamicWorldSourceReadsNoBalance(t *testing.T) { + src := ` + vars { account $src } + send [USD/2 100] (source = $src destination = @dest) + ` + + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + enc, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + vars, err := enc.Encode(map[string]string{"src": "world"}) + require.NoError(t, err) + store := &countingStore{} + machine := vm.NewVm(program) + res, execErr := vm.Exec(context.Background(), machine, &vars, store) + require.Nil(t, execErr) + + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(100)}, + }, res.Postings) + require.Zero(t, store.balanceCalls, "a world source must not read any balance") +} + +// A send-all needs a bounded source to know how much "all" is, and @world is +// unbounded. The specs format has no expectation field for this error, so it is +// asserted here; the interpreter's twin is TestInvalidUnboundedWorldInSendAll. +func TestE2E_SendAllFromWorldErrors(t *testing.T) { + src := `send [USD/2 *] (source = @world destination = @dest)` + + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, e2eStore{}) + + require.Equal(t, vm.InvalidUncappedSource{Account: "world"}, execErr) +} + +// same, but world is only known at run time, so the compiler cannot reject it +func TestE2E_SendAllFromDynamicWorldErrors(t *testing.T) { + src := ` + vars { account $src } + send [USD/2 *] (source = $src destination = @dest) + ` + + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + enc, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + vars, err := enc.Encode(map[string]string{"src": "world"}) + require.NoError(t, err) + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, &vars, e2eStore{}) + + require.Equal(t, vm.InvalidUncappedSource{Account: "world"}, execErr) +} + +func requirePostingsEqual(t *testing.T, want, got []funds.Posting) { + t.Helper() + require.Len(t, got, len(want)) + for i := range want { + w, g := want[i], got[i] + require.Equal(t, w.Source, g.Source, "posting[%d].Source", i) + require.Equal(t, w.Destination, g.Destination, "posting[%d].Destination", i) + require.Equal(t, w.Asset, g.Asset, "posting[%d].Asset", i) + require.Equal(t, w.Color, g.Color, "posting[%d].Color", i) + require.Zero(t, g.Amount.Cmp(w.Amount), "posting[%d].Amount: got %s want %s", i, g.Amount, w.Amount) + } +} diff --git a/internal/compiler/fuzz_mutate_test.go b/internal/compiler/fuzz_mutate_test.go new file mode 100644 index 00000000..95eab9c5 --- /dev/null +++ b/internal/compiler/fuzz_mutate_test.go @@ -0,0 +1,111 @@ +package compiler_test + +import ( + "context" + "math/big" + "testing" + + "github.com/formancehq/numscript/internal/compiler" + "github.com/formancehq/numscript/internal/funds" + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/vm" +) + +// A script covering most of the instruction set: a oneof source (marks, and the +// jumps around them), a capped source (lt_int plus copies), an allotment +// destination (the portion ops and the leftover fixup), a balance() read, and a +// set_account_meta. Deliberately no vars, so the mutants run against nil. +const mutateBaseSrc = `send [USD/2 10] ( + source = { + max [USD/2 5] from @a + oneof { + @b + @c + } + @d + } + destination = { + 1/2 to @e + remaining to @f + } +) +set_account_meta(@e, "k", "v") +` + +// FuzzMutatedBytecode is the converse of FuzzExec: instead of random bytes it +// starts from bytecode the compiler really emitted and corrupts it, so the +// mutants stay close enough to valid that they get past the cheap checks and +// reach the interesting ones. +// +// The property is the same either way — if the verifier accepts it, Exec must +// not crash on it. +func FuzzMutatedBytecode(f *testing.F) { + parsed := parser.Parse(mutateBaseSrc) + if len(parsed.Errors) != 0 { + f.Fatalf("parse: %v", parsed.Errors) + } + featureFlags := map[string]struct{}{"experimental-oneof": {}} + _, base, cErr := compiler.Compile(parsed.Value, featureFlags) + if cErr != nil { + f.Fatalf("compile: %v", cErr) + } + if err := vm.Verify(base); err != nil { + f.Fatalf("the unmutated program must verify: %v", err) + } + + store := e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2"}: big.NewInt(3), + {Account: "b", Asset: "USD/2"}: big.NewInt(4), + {Account: "c", Asset: "USD/2"}: big.NewInt(4), + {Account: "d", Asset: "USD/2"}: big.NewInt(100), + }} + + // The base script reads no variables, but a mutation can turn any byte into + // an Op_LoadVar*, so the mutants have to be verified against the vars they + // will actually be given. Verify alone does not look at them. + vars := &vm.Vars{ + StringsPool: []string{"v"}, + IntsPool: []big.Int{*big.NewInt(1)}, + } + + f.Add([]byte{0, 0}) + f.Add([]byte{4, 255}) + f.Add([]byte{1, 9, 8, 2, 12, 0}) + + f.Fuzz(func(t *testing.T, data []byte) { + if len(base.Instructions) == 0 { + return + } + + flat := make([]byte, len(base.Instructions)*4) + for k, ins := range base.Instructions { + flat[k*4], flat[k*4+1], flat[k*4+2], flat[k*4+3] = ins.Opcode, ins.A, ins.B, ins.C + } + // data is read as (offset, value) pairs. The offset is a uint16 rather + // than the byte the original used, which could only ever reach the first + // 256 bytes — this program is longer than that. + for i := 0; i+2 < len(data); i += 3 { + off := (int(data[i])<<8 | int(data[i+1])) % len(flat) + flat[off] = data[i+2] + } + + instrs := make([]vm.Instruction, len(base.Instructions)) + for k := range instrs { + instrs[k] = vm.Instruction{Opcode: flat[k*4], A: flat[k*4+1], B: flat[k*4+2], C: flat[k*4+3]} + } + + prog := base + prog.Instructions = instrs + + if _, err := vm.VerifyWithVars(prog, vars); err != nil { + return // the mutation broke something static: nothing left to prove + } + + defer func() { + if r := recover(); r != nil { + t.Fatalf("verified program panicked in Exec: %v", r) + } + }() + _, _ = vm.Exec(context.Background(), vm.NewVm(prog), vars, store) + }) +} diff --git a/internal/compiler/scripts_test.go b/internal/compiler/scripts_test.go new file mode 100644 index 00000000..7a4f4d6e --- /dev/null +++ b/internal/compiler/scripts_test.go @@ -0,0 +1,176 @@ +package compiler_test + +import ( + "context" + "encoding/json" + "maps" + "math/big" + "path/filepath" + "slices" + "testing" + + "github.com/formancehq/numscript/internal/compiler" + "github.com/formancehq/numscript/internal/funds" + "github.com/formancehq/numscript/internal/interpreter" + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/specs_format" + "github.com/formancehq/numscript/internal/vm" + + "github.com/stretchr/testify/require" +) + +const scriptsFolder = "../interpreter/testdata/script-tests" + +// scriptsBlacklist lists spec files the compiler+VM can't run yet. What's left +// is asset-scaling, which the compiler has no lowering for. Delete entries as +// features land, until it's empty. +var scriptsBlacklist = []string{ + "experimental/asset-scaling/no-solution.num", + "experimental/asset-scaling/scaling-all-allotment.num", + "experimental/asset-scaling/scaling-allotment.num", + "experimental/asset-scaling/scaling-kept.num", + "experimental/asset-scaling/scaling-prefetch-midscript-balance.num", + "experimental/asset-scaling/scaling-send-all.num", + "experimental/asset-scaling/scaling-with-oneof.num", + "experimental/asset-scaling/scaling.num", + "experimental/asset-scaling/update-swap-account-balance.num", +} + +func TestCompilerScripts(t *testing.T) { + rawSpecs, err := specs_format.ReadSpecsFiles([]string{scriptsFolder}) + require.NoError(t, err) + + for _, rawSpec := range rawSpecs { + rel, err := filepath.Rel(scriptsFolder, rawSpec.NumscriptPath) + require.NoError(t, err) + + t.Run(rel, func(t *testing.T) { + if slices.Contains(scriptsBlacklist, rel) { + t.Skip("blacklisted: not supported yet") + } + + var specs specs_format.Specs + require.NoError(t, json.Unmarshal(rawSpec.SpecsFileContent, &specs)) + + defer func() { + if r := recover(); r != nil { + t.Errorf("panic: %v", r) + } + }() + + runScriptSpec(t, specs, rawSpec.NumscriptContent) + }) + } +} + +func runScriptSpec(t *testing.T, specs specs_format.Specs, src string) { + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + + featureFlags := make(map[string]struct{}, len(specs.FeatureFlags)) + for _, flag := range specs.FeatureFlags { + featureFlags[flag] = struct{}{} + } + + enc, program, cErr := compiler.Compile(parsed.Value, featureFlags) + require.Nil(t, cErr) + + hasFocused := slices.ContainsFunc(specs.TestCases, func(tc specs_format.TestCase) bool { + return tc.Focus + }) + + for _, tc := range specs.TestCases { + if tc.Skip || (hasFocused && !tc.Focus) { + continue + } + caseVars := map[string]string{} + maps.Copy(caseVars, specs.Vars) + maps.Copy(caseVars, tc.Vars) + vars, encErr := enc.Encode(caseVars) + require.NoError(t, encErr, "case %q: encode vars", tc.It) + + balances := specs_format.MergeBalances(specs.Balances, tc.Balances) + + machine := vm.NewVm(program) + store := scriptStore(balances, specs.Meta, tc.Meta) + res, execErr := vm.Exec(context.Background(), machine, &vars, store) + + if tc.ExpectMissingFunds { + require.IsType(t, vm.MissingFundsError{}, execErr, "case %q", tc.It) + continue + } + if tc.ExpectNegativeAmount { + require.IsType(t, vm.NegativeAmountError{}, execErr, "case %q", tc.It) + continue + } + require.Nil(t, execErr, "case %q: unexpected error: %v", tc.It, execErr) + + if tc.ExpectPostings != nil { + require.Equal(t, tc.ExpectPostings, res.Postings, "case %q: expect.postings", tc.It) + } + + if tc.ExpectEndBalances != nil { + got := specs_format.EndBalances(res.Postings, balances) + require.True(t, interpreter.CompareBalances(tc.ExpectEndBalances, got), + "case %q: expect.endBalances: want %v, got %v", tc.It, tc.ExpectEndBalances, got) + } + + if tc.ExpectEndBalancesInclude != nil { + got := specs_format.EndBalances(res.Postings, balances) + require.True(t, interpreter.CompareBalancesIncluding(tc.ExpectEndBalancesInclude, got), + "case %q: expect.endBalances.include: want %v to be included in %v", tc.It, tc.ExpectEndBalancesInclude, got) + } + + if tc.ExpectMovements != nil { + got := specs_format.GetMovements(res.Postings) + require.True(t, specs_format.CompareMovements(tc.ExpectMovements, got), + "case %q: expect.movements: want %v, got %v", tc.It, tc.ExpectMovements, got) + } + + if tc.ExpectTxMeta != nil { + require.Equal(t, txMetaAsStrings(tc.ExpectTxMeta), res.Metadata, "case %q: expect.txMetadata", tc.It) + } + + if tc.ExpectAccountsMeta != nil { + require.Equal(t, accountsMetaAsStrings(tc.ExpectAccountsMeta), res.AccountsMetadata, + "case %q: expect.metadata", tc.It) + } + } +} + +func txMetaAsStrings(rows specs_format.ExpectedTxMeta) map[string]string { + out := map[string]string{} + for _, row := range rows { + out[row.Key] = row.Value + } + return out +} + +func accountsMetaAsStrings(rows interpreter.SetAccountsMetadata) funds.AccountsMetadata { + out := make(funds.AccountsMetadata, 0, len(rows)) + for _, row := range rows { + out = append(out, funds.AccountMetadataEntry{ + Account: row.Account, + Scope: row.Scope, + Key: row.Key, + Value: row.Value, + }) + } + return out +} + +func scriptStore(balances interpreter.Balances, metaOuter, metaInner interpreter.AccountsMetadata) e2eStore { + m := map[funds.PairKey]*big.Int{} + for _, b := range balances { + m[funds.PairKey{Account: b.Account, Scope: b.Scope, Asset: b.Asset, Color: b.Color}] = b.Amount + } + + meta := map[e2eMetaKey]string{} + for _, src := range []interpreter.AccountsMetadata{metaOuter, metaInner} { + for _, row := range src { + meta[e2eMetaKey{account: row.Account, scope: row.Scope, key: row.Key}] = row.Value + } + } + + return e2eStore{balances: m, metadata: meta} +} diff --git a/internal/compiler/testdata/fuzz/FuzzMutatedBytecode/db56cb61b1ca3263 b/internal/compiler/testdata/fuzz/FuzzMutatedBytecode/db56cb61b1ca3263 new file mode 100644 index 00000000..ba8a98c4 --- /dev/null +++ b/internal/compiler/testdata/fuzz/FuzzMutatedBytecode/db56cb61b1ca3263 @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("08\x13") diff --git a/internal/compiler/value.go b/internal/compiler/value.go new file mode 100644 index 00000000..22a43fe2 --- /dev/null +++ b/internal/compiler/value.go @@ -0,0 +1,38 @@ +package compiler + +import "github.com/formancehq/numscript/internal/ir" + +// monetaryValue is a monetary-typed expression after codegen. The VM has no +// monetary register, so a monetary travels as two: the asset in a string +// register and the amount in an int register. +// +// Named fields rather than a returned (asset, amount) pair so that transposing +// the two is a build error instead of an ir.Typecheck error. +type monetaryValue struct { + Asset ir.Reg // str + Amount ir.Reg // int +} + +// accountValue is an account-typed expression after codegen. Scope is a second +// register alongside the name, nil when the expression is provably unscoped (an +// account literal, a plain var, or any account not produced by scoped()) — the +// same nilable-operand idiom PullAccount already uses for Color/Overdraft. +type accountValue struct { + Name ir.Reg // str + Scope *ir.Reg // str +} + +// value is a compiled expression of any type. Mon is set exactly for +// monetary-typed expressions, Acc for account-typed ones, Reg for every other +// type. +type value struct { + Reg ir.Reg + Mon *monetaryValue + Acc *accountValue +} + +func scalarValue(r ir.Reg) value { return value{Reg: r} } + +func monValue(m monetaryValue) value { return value{Mon: &m} } + +func accValue(a accountValue) value { return value{Acc: &a} } diff --git a/internal/compiler/vars_e2e_test.go b/internal/compiler/vars_e2e_test.go new file mode 100644 index 00000000..5f575c6a --- /dev/null +++ b/internal/compiler/vars_e2e_test.go @@ -0,0 +1,154 @@ +package compiler_test + +import ( + "context" + "math/big" + "testing" + + "github.com/formancehq/numscript/internal/compiler" + "github.com/formancehq/numscript/internal/funds" + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/vm" + "github.com/stretchr/testify/require" +) + +func TestE2E_ExternalVars(t *testing.T) { + src := ` + vars { + account $dest + monetary $m + } + send $m ( + source = @world + destination = $dest + ) + ` + + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + + enc, program, err := compiler.Compile(parsed.Value, nil) + require.NoError(t, err) + + vars, err := enc.Encode(map[string]string{ + "dest": "alice", + "m": "USD/2 100", + }) + require.NoError(t, err) + + machine := vm.NewVm(program) + res, execErr := vm.Exec(context.Background(), machine, &vars, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + require.Nil(t, execErr) + + want := []funds.Posting{ + {Source: "world", Destination: "alice", Asset: "USD/2", Amount: big.NewInt(100)}, + } + requirePostingsEqual(t, want, res.Postings) +} + +func TestE2E_InvalidInterpolatedAccount(t *testing.T) { + src := ` + #![feature("experimental-account-interpolation")] + vars { string $status } + set_tx_meta("k", @user:$status) + ` + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + + enc, program, err := compiler.Compile(parsed.Value, nil) + require.NoError(t, err) + + vars, err := enc.Encode(map[string]string{"status": "!invalid acc.."}) + require.NoError(t, err) + + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, &vars, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + require.Equal(t, vm.InvalidAccountName{Name: "user:!invalid acc.."}, execErr) +} + +func compileEncoder(t *testing.T, src string) compiler.VarsEncoder { + t.Helper() + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + enc, _, err := compiler.Compile(parsed.Value, nil) + require.NoError(t, err) + return enc +} + +// A var of each type decomposes into its int/string slots, in declaration order. +func TestVarsEncoder_AllTypes(t *testing.T) { + enc := compileEncoder(t, ` + vars { + number $n + account $acc + portion $p + monetary $m + asset $a + string $s + } + send [COIN 0] (source = @world destination = @world) + `) + + vars, err := enc.Encode(map[string]string{ + "n": "42", + "acc": "alice", + "p": "1/4", + "m": "USD/2 100", + "a": "EUR", + "s": "hello", + }) + require.NoError(t, err) + + // str slots: acc, m.asset, a, s int slots: n, p.num, p.den, m.amount + require.Equal(t, []string{"alice", "USD/2", "EUR", "hello"}, vars.StringsPool) + require.Equal(t, []big.Int{ + *big.NewInt(42), *big.NewInt(1), *big.NewInt(4), *big.NewInt(100), + }, vars.IntsPool) +} + +func TestVarsEncoder_Errors(t *testing.T) { + enc := compileEncoder(t, ` + vars { number $n account $acc } + send [COIN 0] (source = @world destination = @world) + `) + + _, err := enc.Encode(map[string]string{"n": "1"}) + require.ErrorContains(t, err, "missing variable: $acc") + + _, err = enc.Encode(map[string]string{"n": "not-a-number", "acc": "alice"}) + require.ErrorContains(t, err, "variable $n") +} + +// Every var type validates its raw value, and the error names the variable. +func TestVarsEncoder_ErrorsPerType(t *testing.T) { + testCases := []struct { + typ string + raw string + msg string + }{ + {"number", "4.2", `invalid number: "4.2"`}, + {"account", "not an account", `invalid account: "not an account"`}, + {"asset", "usd", `invalid asset: "usd"`}, + {"portion", "nope", "invalid format"}, + {"portion", "200%", "between 0% and 100%"}, + {"monetary", "USD/2", `invalid monetary: "USD/2"`}, + {"monetary", "usd 1", `invalid asset: "usd"`}, + {"string", "anything goes", ""}, + } + + for _, tc := range testCases { + t.Run(tc.typ+" "+tc.raw, func(t *testing.T) { + enc := compileEncoder(t, ` + vars { `+tc.typ+` $v } + send [COIN 0] (source = @world destination = @world) + `) + _, err := enc.Encode(map[string]string{"v": tc.raw}) + if tc.msg == "" { + require.NoError(t, err) + return + } + require.ErrorContains(t, err, "variable $v") + require.ErrorContains(t, err, tc.msg) + }) + } +} diff --git a/internal/compiler/vars_encoder.go b/internal/compiler/vars_encoder.go new file mode 100644 index 00000000..b5e2be70 --- /dev/null +++ b/internal/compiler/vars_encoder.go @@ -0,0 +1,89 @@ +package compiler + +import ( + "fmt" + "math/big" + + "github.com/formancehq/numscript/internal/funds" + "github.com/formancehq/numscript/internal/typecheck" + "github.com/formancehq/numscript/internal/vm" +) + +type VarsEncoder struct { + decls []varDecl + nStr int + nInt int +} + +type varDecl struct { + name string + typ typecheck.Type +} + +// TODO review AI blob +func (e VarsEncoder) Encode(vars map[string]string) (vm.Vars, error) { + strs := make([]string, 0, e.nStr) + ints := make([]big.Int, 0, e.nInt) + + for _, d := range e.decls { + raw, ok := vars[d.name] + if !ok { + return vm.Vars{}, fmt.Errorf("missing variable: $%s", d.name) + } + + var err error + strs, ints, err = appendVar(strs, ints, d.typ, raw) + if err != nil { + return vm.Vars{}, fmt.Errorf("variable $%s: %w", d.name, err) + } + } + + return vm.Vars{StringsPool: strs, IntsPool: ints, Version: vm.CurrentBytecodeVersion}, nil +} + +// TODO review AI blob +func appendVar(strs []string, ints []big.Int, typ typecheck.Type, raw string) ([]string, []big.Int, error) { + switch typ { + case typecheck.TypeNumber: + n, ok := funds.ParseNumber(raw) + if !ok { + return strs, ints, fmt.Errorf("invalid number: %q", raw) + } + ints = append(ints, *n) + + case typecheck.TypeString: + strs = append(strs, raw) + + case typecheck.TypeAccount: + if !funds.ValidateAccount(raw) { + return strs, ints, fmt.Errorf("invalid account: %q", raw) + } + strs = append(strs, raw) + + case typecheck.TypeAsset: + if !funds.ValidateAsset(raw) { + return strs, ints, fmt.Errorf("invalid asset: %q", raw) + } + strs = append(strs, raw) + + case typecheck.TypePortion: + r, err := funds.ParsePortion(raw) + if err != nil { + return strs, ints, err + } + ints = append(ints, *r.Num(), *r.Denom()) + + case typecheck.TypeMonetary: + asset, amount, err := funds.ParseMonetary(raw) + if err != nil { + return strs, ints, err + } + strs = append(strs, asset) + ints = append(ints, *amount) + + default: + panic("unexpected var type: " + typ) + } + + return strs, ints, nil +} diff --git a/internal/compiler/verify_corpus_test.go b/internal/compiler/verify_corpus_test.go new file mode 100644 index 00000000..3544ff6a --- /dev/null +++ b/internal/compiler/verify_corpus_test.go @@ -0,0 +1,51 @@ +package compiler_test + +import ( + "encoding/json" + "path/filepath" + "slices" + "testing" + + "github.com/formancehq/numscript/internal/compiler" + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/specs_format" + "github.com/formancehq/numscript/internal/vm" + + "github.com/stretchr/testify/require" +) + +// Compile does not call vm.Verify — the VM assumes its own compiler's output is +// well formed, and paying for a full static pass on every compile would make +// that assumption cost something. This test is what earns the assumption: every +// script in the corpus must produce bytecode the verifier accepts. +func TestCompiledCorpusPassesVerify(t *testing.T) { + rawSpecs, err := specs_format.ReadSpecsFiles([]string{scriptsFolder}) + require.NoError(t, err) + + for _, rawSpec := range rawSpecs { + rel, err := filepath.Rel(scriptsFolder, rawSpec.NumscriptPath) + require.NoError(t, err) + + t.Run(rel, func(t *testing.T) { + if slices.Contains(scriptsBlacklist, rel) { + t.Skip("blacklisted: not supported yet") + } + + var specs specs_format.Specs + require.NoError(t, json.Unmarshal(rawSpec.SpecsFileContent, &specs)) + + featureFlags := make(map[string]struct{}, len(specs.FeatureFlags)) + for _, flag := range specs.FeatureFlags { + featureFlags[flag] = struct{}{} + } + + parsed := parser.Parse(rawSpec.NumscriptContent) + require.Empty(t, parsed.Errors) + + _, program, cErr := compiler.Compile(parsed.Value, featureFlags) + require.Nil(t, cErr) + + require.NoError(t, vm.Verify(program)) + }) + } +} diff --git a/internal/difftest/compare.go b/internal/difftest/compare.go index 8e4f12e5..c5722d71 100644 --- a/internal/difftest/compare.go +++ b/internal/difftest/compare.go @@ -24,12 +24,24 @@ func mismatch(format string, args ...any) Verdict { return Verdict{Mismatch: true, Reason: fmt.Sprintf(format, args...)} } -// Compare normalizes and diffs two engines' results for the same script. -// aLabel/bLabel appear in mismatch messages only. +// Compare normalizes and diffs two engines' results for the same script, with +// the oracle on the b-side. aLabel/bLabel appear in mismatch messages only. // // Error strings are never compared, only whether an error occurred and at which // stage: wording legitimately differs between implementations. func Compare(aRes, bRes SideResult, aLabel, bLabel string) Verdict { + return compare(aRes, bRes, aLabel, bLabel, true) +} + +// CompareEngines is Compare for the vm against the interpreter. Both live in +// this repo and must agree exactly, so only the vm's capacity bounds are +// tolerated: every tolerance that exists for a legacy-machine behavior is a +// mismatch here. +func CompareEngines(vmRes, newRes SideResult) Verdict { + return compare(vmRes, newRes, "vm", "new interpreter", false) +} + +func compare(aRes, bRes SideResult, aLabel, bLabel string, oracle bool) Verdict { // Checked before anything else, and symmetrically: an engine breaking its // own contract is never an expected outcome, and every tolerance below is // about the two engines legitimately disagreeing. Ordering matters — the @@ -45,6 +57,9 @@ func Compare(aRes, bRes SideResult, aLabel, bLabel string) Verdict { bCompileFailed := bRes.CompileErr != "" if bCompileFailed && !aCompileFailed { + if !oracle { + return mismatch("%s rejected a script %s compiled: compileErr=%q", bLabel, aLabel, bRes.CompileErr) + } // Expected, not a mismatch: internal/gen's cleanup pass is best-effort, not // a guarantee — it does not track unboundedness propagating up through // nested inorder blocks, so it can still emit a script the b-side rejects. @@ -52,6 +67,20 @@ func Compare(aRes, bRes SideResult, aLabel, bLabel string) Verdict { return tolerated("b-side compile rejection") } if aCompileFailed && !bCompileFailed { + if aRes.RegisterOverflow { + // A known capacity bound, not a semantic rejection: the vm's + // one-byte register operands cap how many values can be live at + // once, and a pathological generated script can exceed that. It + // fails closed at compile time and nothing is compared, so it is + // tolerated by name and counted — any other a-side rejection of a + // script the b-side ran stays the interesting direction below. + return tolerated("vm register capacity") + } + if aRes.ProgramTooLarge { + // The encoding's other capacity bound (uint16 jump targets), with + // the same reasoning and the same fail-closed behavior. + return tolerated("vm program size") + } // The interesting direction: the generator stays within the b-side's // grammar and a-side should be a strict superset, so a-side rejecting what // b-side compiled is a genuine divergence. @@ -75,6 +104,9 @@ func Compare(aRes, bRes SideResult, aLabel, bLabel string) Verdict { // sides failing to compile. return ok() } + if !oracle { + return mismatch("%s rejected a script %s resolved: resolveErr=%q", bLabel, aLabel, bRes.ResolveErr) + } // Same shape as the compile-stage tolerance above, one stage later: // the generator's cleanup pass does not track everything the oracle's // resolve stage refuses (e.g. binding `@world` to an account variable @@ -110,7 +142,7 @@ func Compare(aRes, bRes SideResult, aLabel, bLabel string) Verdict { // The reverse is DIVERGENCES.md #4, where b-side rejects the script and // a-side silently zeroes the clause and runs short. That one must still be // a mismatch -- see TestMissingFundsClassificationMismatchStillCaught. - if aNegativeAmount(aRes) && bRes.MissingFunds { + if oracle && aNegativeAmount(aRes) && bRes.MissingFunds { return tolerated("negative amount vs missing funds") } return mismatch( @@ -127,7 +159,7 @@ func Compare(aRes, bRes SideResult, aLabel, bLabel string) Verdict { // keeps going. Anything else here is a real bug -- one engine moved // money the other refused to, or vice versa -- and must not be // absorbed. - if bRunFailed && !aRunFailed && bRes.NegativeMaxReject { + if oracle && bRunFailed && !aRunFailed && bRes.NegativeMaxReject { return tolerated("negative max clause") } return mismatch( diff --git a/internal/difftest/difftest.go b/internal/difftest/difftest.go index 3a4a058e..189bd5cc 100644 --- a/internal/difftest/difftest.go +++ b/internal/difftest/difftest.go @@ -12,60 +12,72 @@ import ( "github.com/formancehq/numscript/internal/gen" ) -// Case is one generated script plus both engines' results and the verdict -// between them. +// Case is one generated script plus all three engines' results and the +// pairwise verdicts between them. type Case struct { Script string Vars map[string]string Shape gen.Shape New SideResult Oracle SideResult + VM SideResult // OracleVsNew compares this repo's interpreter against the vendored // legacy machine, which is the only independent ground truth available: // the oracle is a separate implementation, maintained elsewhere, that - // this repo's engine must agree with. + // this repo's engine must agree with. OracleVsVM extends the same check + // to the compiler+VM engine; NewVsVM has no independent ground truth + // (both engines live in this repo), so a mismatch there means the + // interpreter and the compiler+VM disagree with each other, not just + // with the legacy oracle. // - // The harness is built to carry more engines than this — Compare is - // deliberately engine-agnostic and takes labels. When internal/vm lands, - // re-add a VM leg here (runVM, Case.VM, OracleVsVM and NewVsVM) so the - // compiler+VM is checked against both the oracle and the interpreter. + // The VM sits on the a-side of its two verdicts so that Compare's b-side + // tolerances keep their meaning: an oracle-side rejection stays an + // expected outcome, while the VM rejecting a script another engine ran is + // a mismatch. OracleVsNew Verdict + OracleVsVM Verdict + NewVsVM Verdict } -// RunOne generates one program from rng, runs it against both engines, and -// compares the results. +// RunOne generates one program from rng, runs it against all three engines, +// and compares the results pairwise. // // A script the generator marked oracle-incompatible (numscript-only shapes: // oneof, colors, division portions, wrong-asset caps) never reaches the legacy -// machine: its oracle leg is recorded as a named tolerance so the sweep counts -// it. Until the compiler+VM leg lands, such a script is only executed by the -// interpreter (a panic still fails the fuzz target); the VM leg is what will -// compare those shapes engine-against-engine. +// machine: its oracle legs are recorded as a named tolerance so the sweep +// counts them, and only NewVsVM is compared — which is the point of those +// shapes, since both engines live in this repo and must agree exactly. func RunOne(ctx context.Context, rng *rand.Rand) Case { g := gen.Generate(rng) newRes := runNew(ctx, g.Script, g.Vars, g.Balances, g.Metadata, g.Flags) + vmRes := runVM(ctx, g.Script, g.Vars, g.Balances, g.Metadata, g.Flags) c := Case{ Script: g.Script, Vars: g.Vars, Shape: g.Shape, New: newRes, + VM: vmRes, + + NewVsVM: CompareEngines(vmRes, newRes), } if g.OracleCompatible { c.Oracle = runOracle(ctx, g.Script, g.Vars, g.Balances, g.Metadata) c.OracleVsNew = Compare(newRes, c.Oracle, "new interpreter", "oracle") + c.OracleVsVM = Compare(vmRes, c.Oracle, "vm", "oracle") } else { c.OracleVsNew = tolerated("numscript-only script, oracle skipped") + c.OracleVsVM = tolerated("numscript-only script, oracle skipped") } return c } -// Legs returns the pairwise verdicts with stable names, for callers that -// report per leg. One leg today; the compiler+VM adds two more. +// Legs returns the three pairwise verdicts with stable names, for callers +// that report per leg. func (c Case) Legs() []struct { Name string Verdict Verdict @@ -75,10 +87,13 @@ func (c Case) Legs() []struct { Verdict Verdict }{ {"oracle vs new", c.OracleVsNew}, + {"oracle vs vm", c.OracleVsVM}, + {"new vs vm", c.NewVsVM}, } } -// AnyMismatch reports whether the comparison found a divergence. +// AnyMismatch reports whether any of the three pairwise verdicts found a +// divergence. func (c Case) AnyMismatch() bool { - return c.OracleVsNew.Mismatch + return c.OracleVsNew.Mismatch || c.OracleVsVM.Mismatch || c.NewVsVM.Mismatch } diff --git a/internal/difftest/difftest_test.go b/internal/difftest/difftest_test.go index 04c6df43..7d327257 100644 --- a/internal/difftest/difftest_test.go +++ b/internal/difftest/difftest_test.go @@ -24,11 +24,11 @@ func FuzzDiff(f *testing.F) { rng := gen.RandFromBytes(data) c := difftest.RunOne(context.Background(), rng) - for _, v := range []difftest.Verdict{c.OracleVsNew} { - if v.Mismatch { + for _, leg := range c.Legs() { + if leg.Verdict.Mismatch { t.Fatalf( - "divergence: %s\n\nvars: %v\n\nscript:\n%s", - v.Reason, c.Vars, c.Script, + "divergence (%s): %s\n\nvars: %v\n\nscript:\n%s", + leg.Name, leg.Verdict.Reason, c.Vars, c.Script, ) } } @@ -76,3 +76,69 @@ func TestCompareStillToleratesAPlainRejection(t *testing.T) { t.Fatalf("a plain b-side rejection should be tolerated, got %+v", v) } } + +// An a-side compile rejection is a mismatch (the vm refusing a script another +// engine ran), with two named exceptions, both fail-closed capacity bounds of +// the bytecode encoding that a pathological generated script can exceed: the +// register bank (ir.ErrRegisterBankOverflow) and the instruction count a jump +// target can address (ir.ErrProgramTooLarge). Tolerated and counted; every +// other a-side rejection stays flagged. +func TestCompareToleratesOnlyTheCapacityRejections(t *testing.T) { + overflow := difftest.SideResult{CompileErr: "register bank overflow: ...", RegisterOverflow: true} + if v := difftest.Compare(overflow, difftest.SideResult{}, "vm", "oracle"); v.Mismatch || v.Tolerated != "vm register capacity" { + t.Fatalf("expected the register-capacity tolerance, got %+v", v) + } + + tooLarge := difftest.SideResult{CompileErr: "program too large: ...", ProgramTooLarge: true} + if v := difftest.Compare(tooLarge, difftest.SideResult{}, "vm", "oracle"); v.Mismatch || v.Tolerated != "vm program size" { + t.Fatalf("expected the program-size tolerance, got %+v", v) + } + + plain := difftest.SideResult{CompileErr: "no such feature"} + if v := difftest.Compare(plain, difftest.SideResult{}, "vm", "oracle"); !v.Mismatch { + t.Fatalf("a plain a-side compile rejection must stay a mismatch, got %+v", v) + } +} + +// The vm and the interpreter must agree exactly: every tolerance that exists for +// a legacy-machine behavior is a mismatch between them, and only the vm's +// capacity bounds stay tolerated. +func TestCompareEnginesHasNoOracleTolerances(t *testing.T) { + ran := difftest.SideResult{} + + mismatches := []struct { + name string + vm, nw difftest.SideResult + }{ + {"interpreter-only compile rejection", ran, difftest.SideResult{CompileErr: "rejected"}}, + {"interpreter-only resolve rejection", ran, difftest.SideResult{ResolveErr: "rejected"}}, + {"negative amount vs missing funds", + difftest.SideResult{RunErr: "negative amount", NegativeAmount: true}, + difftest.SideResult{RunErr: "missing funds", MissingFunds: true}}, + {"negative max clause", ran, difftest.SideResult{RunErr: "negative max", NegativeMaxReject: true}}, + } + for _, tc := range mismatches { + t.Run(tc.name, func(t *testing.T) { + if v := difftest.CompareEngines(tc.vm, tc.nw); !v.Mismatch { + t.Fatalf("expected a mismatch, got %+v", v) + } + }) + } + + t.Run("vm capacity bounds stay tolerated", func(t *testing.T) { + overflow := difftest.SideResult{CompileErr: "register bank overflow: ...", RegisterOverflow: true} + if v := difftest.CompareEngines(overflow, ran); v.Mismatch || v.Tolerated != "vm register capacity" { + t.Fatalf("expected the register-capacity tolerance, got %+v", v) + } + tooLarge := difftest.SideResult{CompileErr: "program too large: ...", ProgramTooLarge: true} + if v := difftest.CompareEngines(tooLarge, ran); v.Mismatch || v.Tolerated != "vm program size" { + t.Fatalf("expected the program-size tolerance, got %+v", v) + } + }) + + t.Run("both failing at different stages is not a mismatch", func(t *testing.T) { + if v := difftest.CompareEngines(difftest.SideResult{RunErr: "boom"}, difftest.SideResult{ResolveErr: "boom"}); v.Mismatch || v.Tolerated != "" { + t.Fatalf("expected agreement, got %+v", v) + } + }) +} diff --git a/internal/difftest/numscript_only_shapes_test.go b/internal/difftest/numscript_only_shapes_test.go new file mode 100644 index 00000000..d86c89dc --- /dev/null +++ b/internal/difftest/numscript_only_shapes_test.go @@ -0,0 +1,334 @@ +package difftest + +import ( + "context" + "math/big" + "strings" + "testing" + + "github.com/formancehq/numscript/internal/flags" + "github.com/formancehq/numscript/internal/gen" +) + +// TestNumscriptOnlyShapeAgreements pins NewVsVM agreement on shapes the oracle +// cannot parse, so only the two in-repo engines can check each other: oneof, +// colored sources, account interpolation, mid-script calls, division-expression +// portions and mixed-asset caps. The sweep reaches most of these since the +// generator learned them, but each case here is a deterministic witness of a +// specific behavior — several were live NewVsVM bugs when first probed +// (2026-09-25): the eager/lazy clause-evaluation family, the negative +// allotment share family, and the oneof covered-check under a negative cap. +// +// The specs fixture corpus (internal/interpreter/testdata/script-tests) pins +// the same families wherever the expectation is postings or a typed error; +// cases whose expected outcome is "both engines reject, same classification" +// live here, since the specs format cannot express generic errors. +func TestNumscriptOnlyShapeAgreements(t *testing.T) { + testCases := []struct { + name string + script string + flags []string + vars map[string]string + balances map[gen.BalanceKey]*big.Int + metadata map[gen.MetaKey]string + // bothReject: the agreement is that neither engine commits anything. + bothReject bool + }{ + { + // The interpreter evaluates every source-inorder clause even after + // the cap is exhausted; the vm once jumped out early and committed + // where the interpreter errors on the wrong-asset cap. + name: "exhausted source-inorder clause still evaluates its cap", + bothReject: true, + balances: map[gen.BalanceKey]*big.Int{{Account: "a", Asset: "COIN"}: big.NewInt(100)}, + script: `send [COIN 10] ( + source = { + max [COIN 10] from @a + max [EUR 1] from @b + } + destination = @d +)`, + }, + { + // The negative portion is rejected before the oneof runs, so the + // over-100% allotment in its second branch is never evaluated + // (sweep seed 571). + name: "negative allotment portion rejected before its oneof runs", + flags: []string{flags.ExperimentalOneofFeatureFlag}, + bothReject: true, + vars: map[string]string{"n": "-1"}, + balances: map[gen.BalanceKey]*big.Int{ + {Account: "a", Asset: "COIN"}: big.NewInt(500), + {Account: "b", Asset: "COIN"}: big.NewInt(500), + }, + script: `vars { + number $n +} + +send [COIN 90] ( + source = { + $n/3 from oneof { + @a + { + 3/2 from @b + remaining from @b + } + } + remaining from @b + } + destination = @d +)`, + }, + { + name: "oneof source: first branch short, second covers", + flags: []string{flags.ExperimentalOneofFeatureFlag}, + balances: map[gen.BalanceKey]*big.Int{ + {Account: "a", Asset: "COIN"}: big.NewInt(5), + {Account: "b", Asset: "COIN"}: big.NewInt(50), + }, + script: `send [COIN 10] ( + source = oneof { @a @b } + destination = @d +)`, + }, + { + name: "oneof source: all branches short fails as missing funds", + flags: []string{flags.ExperimentalOneofFeatureFlag}, + bothReject: true, + balances: map[gen.BalanceKey]*big.Int{ + {Account: "a", Asset: "COIN"}: big.NewInt(5), + {Account: "b", Asset: "COIN"}: big.NewInt(7), + }, + script: `send [COIN 10] ( + source = oneof { @a @b } + destination = @d +)`, + }, + { + // The first branch's partial pulls must be rolled back before the + // second branch runs, or source attribution differs. + name: "oneof source: nested inorder branch rolls back", + flags: []string{flags.ExperimentalOneofFeatureFlag}, + balances: map[gen.BalanceKey]*big.Int{ + {Account: "a1", Asset: "COIN"}: big.NewInt(4), + {Account: "a2", Asset: "COIN"}: big.NewInt(4), + {Account: "b1", Asset: "COIN"}: big.NewInt(30), + }, + script: `send [COIN 10] ( + source = oneof { + { @a1 @a2 } + { @b1 } + } + destination = @d +)`, + }, + { + name: "oneof source under send-all takes the first branch only", + flags: []string{flags.ExperimentalOneofFeatureFlag}, + balances: map[gen.BalanceKey]*big.Int{ + {Account: "a", Asset: "COIN"}: big.NewInt(5), + {Account: "b", Asset: "COIN"}: big.NewInt(50), + }, + script: `send [COIN *] ( + source = oneof { @a @b } + destination = @d +)`, + }, + { + name: "oneof destination: clause choice by cap", + flags: []string{flags.ExperimentalOneofFeatureFlag}, + script: `send [COIN 10] ( + source = @world + destination = oneof { + max [COIN 5] to @a + max [COIN 10] to @b + remaining to @c + } +)`, + }, + { + name: "oneof destination: no clause covers, remaining takes all", + flags: []string{flags.ExperimentalOneofFeatureFlag}, + script: `send [COIN 10] ( + source = @world + destination = oneof { + max [COIN 5] to @a + max [COIN 9] to @b + remaining to @c + } +)`, + }, + { + name: "colored funds pulled from world stay colored downstream", + flags: []string{flags.ExperimentalAssetColors}, + script: `send [COIN 30] ( + source = @world \ "RED" + destination = @a +) + +send [COIN 10] ( + source = @a \ "RED" + destination = @b +)`, + }, + { + // A colored pull must not see the uncolored store balance: the + // account has 25 RED and 7 uncolored, and only the 25 move. + name: "colored pull reads only the colored balance", + flags: []string{flags.ExperimentalAssetColors}, + balances: map[gen.BalanceKey]*big.Int{ + {Account: "b", Asset: "COIN"}: big.NewInt(7), + }, + script: `send [COIN 30] ( + source = @world \ "RED" + destination = @b +) + +send [COIN *] ( + source = @b \ "RED" + destination = @c +)`, + }, + { + name: "invalid color via var rejected by both", + flags: []string{flags.ExperimentalAssetColors}, + bothReject: true, + vars: map[string]string{"c": "not-valid-color"}, + balances: map[gen.BalanceKey]*big.Int{{Account: "a", Asset: "COIN"}: big.NewInt(25)}, + script: `vars { + string $c +} + +send [COIN 20] ( + source = @a \ $c + destination = @b +)`, + }, + { + name: "interpolated account from a number var", + flags: []string{flags.ExperimentalAccountInterpolationFlag}, + vars: map[string]string{"id": "42"}, + script: `vars { + number $id +} + +send [COIN 10] ( + source = @world + destination = @users:$id +)`, + }, + { + name: "invalid interpolated account rejected by both", + flags: []string{flags.ExperimentalAccountInterpolationFlag}, + bothReject: true, + vars: map[string]string{"id": "no spaces!"}, + script: `vars { + string $id +} + +send [COIN 10] ( + source = @world + destination = @users:$id +)`, + }, + { + name: "mid-script balance reflects earlier statements", + flags: []string{flags.ExperimentalMidScriptFunctionCall}, + balances: map[gen.BalanceKey]*big.Int{ + {Account: "dest", Asset: "COIN"}: big.NewInt(7), + }, + script: `send [COIN 50] ( + source = @world + destination = @dest +) + +set_tx_meta("bal", balance(@dest, COIN))`, + }, + { + name: "mid-script balance as a send amount", + flags: []string{flags.ExperimentalMidScriptFunctionCall}, + balances: map[gen.BalanceKey]*big.Int{ + {Account: "a", Asset: "COIN"}: big.NewInt(30), + }, + script: `send [COIN 20] ( + source = @world + destination = @a +) + +send balance(@a, COIN) ( + source = @a + destination = @b +)`, + }, + { + name: "negative mid-script balance rejected by both", + flags: []string{flags.ExperimentalMidScriptFunctionCall}, + bothReject: true, + balances: map[gen.BalanceKey]*big.Int{{Account: "a", Asset: "COIN"}: big.NewInt(-5)}, + script: `set_tx_meta("bal", balance(@a, COIN))`, + }, + { + name: "amounts beyond int64 round-trip the whole pipeline", + script: `send [COIN 123456789012345678901234567890] ( + source = @world + destination = @a +)`, + }, + { + // long strings through the vars payload, the string pools and the + // program/vars codecs (runVM byte-stability check included) + name: "long account name and long meta value round-trip encoding", + vars: map[string]string{ + "acc": "acc-" + strings.Repeat("x", 400) + ":seg-" + strings.Repeat("y", 300), + }, + balances: map[gen.BalanceKey]*big.Int{ + {Account: "acc-" + strings.Repeat("x", 400) + ":seg-" + strings.Repeat("y", 300), Asset: "COIN"}: big.NewInt(40), + }, + script: `vars { + account $acc +} + +send [COIN 25] ( + source = $acc + destination = @d +) + +set_tx_meta("note", "` + strings.Repeat("z", 900) + `")`, + }, + { + name: "negative monetary var send amount rejected by both", + bothReject: true, + vars: map[string]string{"m": "COIN -5"}, + balances: map[gen.BalanceKey]*big.Int{{Account: "a", Asset: "COIN"}: big.NewInt(100)}, + script: `vars { + monetary $m +} + +send $m ( + source = @a + destination = @b +)`, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + newRes := runNew(ctx, tc.script, tc.vars, tc.balances, tc.metadata, tc.flags) + vmRes := runVM(ctx, tc.script, tc.vars, tc.balances, tc.metadata, tc.flags) + + if v := CompareEngines(vmRes, newRes); v.Mismatch { + t.Fatalf("vm mismatch: %s\nvm: %+v\nnew: %+v", v.Reason, vmRes, newRes) + } + if tc.bothReject { + if !newRes.Failed() || !vmRes.Failed() { + t.Fatalf("expected both engines to reject\nnew: %+v\nvm: %+v", newRes, vmRes) + } + return + } + if newRes.Failed() || vmRes.Failed() { + t.Fatalf("expected both engines to complete\nnew: %+v\nvm: %+v", newRes, vmRes) + } + }) + } +} diff --git a/internal/difftest/regression_test.go b/internal/difftest/regression_test.go index 2ed9c32d..f5d41a98 100644 --- a/internal/difftest/regression_test.go +++ b/internal/difftest/regression_test.go @@ -214,11 +214,17 @@ send [COIN *] ( ctx := context.Background() newRes := runNew(ctx, tc.script, tc.vars, tc.balances, nil, nil) oracleRes := runOracle(ctx, tc.script, tc.vars, tc.balances, nil) + vmRes := runVM(ctx, tc.script, tc.vars, tc.balances, nil, nil) - v := Compare(newRes, oracleRes, "new interpreter", "oracle") - if v.Mismatch { + if v := Compare(newRes, oracleRes, "new interpreter", "oracle"); v.Mismatch { t.Fatalf("mismatch: %s\nnew: %+v\noracle: %+v", v.Reason, newRes, oracleRes) } + if v := Compare(vmRes, oracleRes, "vm", "oracle"); v.Mismatch { + t.Fatalf("mismatch: %s\nvm: %+v\noracle: %+v", v.Reason, vmRes, oracleRes) + } + if v := CompareEngines(vmRes, newRes); v.Mismatch { + t.Fatalf("mismatch: %s\nvm: %+v\nnew: %+v", v.Reason, vmRes, newRes) + } }) } } @@ -256,6 +262,15 @@ func TestDestinationSideNegativeMaxTolerated(t *testing.T) { if oracleRes.RunErr == "" || oracleRes.MissingFunds { t.Fatalf("expected the oracle to reject the negative max for a non-missing-funds reason; got %+v", oracleRes) } + + // The vm clamps like the interpreter, so the same tolerance fires on its leg. + vmRes := runVM(ctx, script, nil, nil, nil, nil) + if vmRes.Failed() { + t.Fatalf("expected the vm to clamp and succeed; got %+v", vmRes) + } + if v := Compare(vmRes, oracleRes, "vm", "oracle"); v.Mismatch || v.Tolerated != "negative max clause" { + t.Fatalf("expected the negative max clause tolerance on the vm leg, got %+v\nvm: %+v", v, vmRes) + } } // TestSourceSideNegativeMaxClauseTolerated locks in an accepted gap: a negative @@ -293,6 +308,12 @@ func TestSourceSideNegativeMaxClauseTolerated(t *testing.T) { if newRes.Failed() { t.Fatalf("expected the interpreter to succeed; got new=%+v", newRes) } + + // The vm must side with the interpreter on the gap. + vmRes := runVM(ctx, script, nil, nil, nil, nil) + if v := CompareEngines(vmRes, newRes); v.Mismatch { + t.Fatalf("the vm does not side with the interpreter: %s\nvm: %+v\nnew: %+v", v.Reason, vmRes, newRes) + } } // TestMissingFundsClassificationMismatchStillCaught is the companion to @@ -318,6 +339,15 @@ func TestMissingFundsClassificationMismatchStillCaught(t *testing.T) { if v := Compare(newRes, oracleRes, "new interpreter", "oracle"); !v.Mismatch { t.Fatalf("expected new-vs-oracle to be flagged as a mismatch, got none") } + + // Same classification on the vm, so its oracle leg is flagged too. + vmRes := runVM(ctx, script, nil, nil, nil, nil) + if !vmRes.MissingFunds { + t.Fatalf("expected the vm to fail specifically due to missing funds; got vm=%+v", vmRes) + } + if v := Compare(vmRes, oracleRes, "vm", "oracle"); !v.Mismatch { + t.Fatalf("expected vm-vs-oracle to be flagged as a mismatch, got none") + } } // TestKnownOpenDivergences pins the numscript/ledger disagreements that are @@ -358,6 +388,7 @@ func TestKnownOpenDivergences(t *testing.T) { ctx := context.Background() newRes := runNew(ctx, tc.script, tc.vars, tc.balances, nil, nil) oracleRes := runOracle(ctx, tc.script, tc.vars, tc.balances, nil) + vmRes := runVM(ctx, tc.script, tc.vars, tc.balances, nil, nil) v := Compare(newRes, oracleRes, "new interpreter", "oracle") if !v.Mismatch { @@ -365,6 +396,12 @@ func TestKnownOpenDivergences(t *testing.T) { tc.why, newRes, oracleRes) } t.Logf("still diverging, as expected (%s): %s", tc.why, v.Reason) + + // The divergence is numscript-vs-ledger; within numscript the two + // engines must still agree on it. + if v := CompareEngines(vmRes, newRes); v.Mismatch { + t.Fatalf("the vm does not side with the interpreter: %s\nvm: %+v\nnew: %+v", v.Reason, vmRes, newRes) + } }) } } diff --git a/internal/difftest/run_new.go b/internal/difftest/run_new.go index cc22bac7..513f30ee 100644 --- a/internal/difftest/run_new.go +++ b/internal/difftest/run_new.go @@ -61,6 +61,19 @@ type SideResult struct { // divergence (oracle/DIVERGENCES.md #4) without tolerating every other // one. NegativeMaxReject bool + // RegisterOverflow is only meaningful when CompileErr is set on the vm's + // side: true iff the compiler refused the script because it needs more + // simultaneously-live registers than the bytecode encoding's one-byte + // operands can address (ir.ErrRegisterBankOverflow). A known capacity + // bound, not a semantic rejection: Compare tolerates it by name so any + // other vm-side rejection of a script another engine ran stays a + // mismatch. + RegisterOverflow bool + // ProgramTooLarge is RegisterOverflow's sibling for the encoding's other + // capacity bound: an instruction stream outgrowing what a uint16 jump + // target can address (ir.ErrProgramTooLarge). Same treatment: tolerated by + // name and counted. + ProgramTooLarge bool // InternalErr is set when an engine broke its own contract, as opposed to // rejecting the script. Deliberately not CompileErr: Compare tolerates one // side rejecting what the other accepted, so a self-inconsistency reported as diff --git a/internal/difftest/run_vm.go b/internal/difftest/run_vm.go new file mode 100644 index 00000000..07116619 --- /dev/null +++ b/internal/difftest/run_vm.go @@ -0,0 +1,129 @@ +package difftest + +import ( + "bytes" + "context" + "errors" + "math/big" + + "github.com/formancehq/numscript" + "github.com/formancehq/numscript/internal/gen" + "github.com/formancehq/numscript/internal/ir" + "github.com/formancehq/numscript/internal/vm" +) + +// vmStore is a vm.Store over the same (account, asset) -> amount / +// (account, key) -> value maps runNew/runOracle already build their stores +// from. internal/gen presets only uncolored, unscoped balances and metadata, +// so a query for any other color or scope answers zero/absent — exactly what +// runNew's StaticStore does, since its rows all carry the empty color and +// scope. Answering the uncolored balance regardless of color (as this store +// once did) hands the vm phantom colored funds the interpreter doesn't see. +type vmStore struct { + balances map[gen.BalanceKey]*big.Int + metadata map[gen.MetaKey]string +} + +func (s vmStore) GetBalance(_ context.Context, account, scope, asset, color string) (*big.Int, error) { + if scope == "" && color == "" { + if amount, ok := s.balances[gen.BalanceKey{Account: account, Asset: asset}]; ok { + return new(big.Int).Set(amount), nil + } + } + return new(big.Int), nil +} + +func (s vmStore) GetMetadata(_ context.Context, account, scope, key string) (string, bool, error) { + if scope != "" { + return "", false, nil + } + value, ok := s.metadata[gen.MetaKey{Account: account, Key: key}] + return value, ok, nil +} + +func runVM(ctx context.Context, script string, vars map[string]string, balances map[gen.BalanceKey]*big.Int, metadata map[gen.MetaKey]string, featureFlags []string) SideResult { + flagSet := make(map[string]struct{}, len(featureFlags)) + for _, f := range featureFlags { + flagSet[f] = struct{}{} + } + varsEncoder, program, err := numscript.CompileWithFeatureFlags(script, flagSet) + if err != nil { + return SideResult{ + CompileErr: err.Error(), + RegisterOverflow: errors.Is(err, ir.ErrRegisterBankOverflow), + ProgramTooLarge: errors.Is(err, ir.ErrProgramTooLarge), + } + } + + // Encode/DecodeProgram must round-trip on everything the compiler emits; + // nothing else in the differential loop exercises the codec, and generated + // scripts reach value shapes (multi-word big.Ints, long pools) the codec + // tests don't. Byte-stability (re-encoding the decoded program reproduces + // the bytes) is the cheap full-structure check; behavioral equivalence is + // pinned separately by the corpus. + encoded := program.Encode() + decoded, decErr := numscript.DecodeCompiledProgram(encoded) + if decErr != nil { + return SideResult{InternalErr: "DecodeProgram failed on Encode output: " + decErr.Error()} + } + if !bytes.Equal(decoded.Encode(), encoded) { + return SideResult{InternalErr: "Encode/Decode/Encode is not byte-stable"} + } + + // Binding the vars is the compile-to-run boundary, like the machine's + // SetVarsFromJSON: report it as the resolve stage so Compare flags the vm + // rejecting bindings the other engines accepted. + encodedVars, err := varsEncoder.Encode(vars) + if err != nil { + return SideResult{ResolveErr: err.Error()} + } + + // The verifier is opt-in and the compiler is trusted not to need it, so this + // is not defending the run — it is checking that claim on every generated + // script. A failure means the compiler emitted bytecode the VM's own static + // rules reject, which is a bug in this repo rather than a disagreement with + // the oracle, hence InternalErr and not CompileErr. + // + // Worth having here specifically because internal/gen reaches shapes the + // hand-written corpus doesn't, and it does so on inputs nobody chose. + if _, err := vm.VerifyWithVars(program, &encodedVars); err != nil { + return SideResult{InternalErr: "compiled program failed verification: " + err.Error()} + } + + store := vmStore{balances: balances, metadata: metadata} + + // The public entry point, so this leg exercises exactly what an integrator + // calls — including its metadata contract. + execResult, execErr := numscript.ExecVm(ctx, numscript.NewVm(program), &encodedVars, store) + if execErr != nil { + var missingFunds vm.MissingFundsError + var negativeAmount vm.NegativeAmountError + return SideResult{ + RunErr: execErr.Error(), + MissingFunds: errors.As(execErr, &missingFunds), + NegativeAmount: errors.As(execErr, &negativeAmount), + } + } + + postings := make([]Posting, 0, len(execResult.Postings)) + for _, p := range execResult.Postings { + postings = append(postings, Posting{ + Source: p.Source, + Destination: p.Destination, + Asset: p.Asset, + Color: p.Color, + Amount: p.Amount, + }) + } + + txMeta := make(map[string]string, len(execResult.Metadata)) + for k, v := range execResult.Metadata { + txMeta[k] = v + } + accountMeta := make(map[string]string, len(execResult.AccountsMetadata)) + for _, row := range execResult.AccountsMetadata { + accountMeta[metaKey(row.Account, row.Key)] = row.Value + } + + return SideResult{Postings: postings, TxMeta: txMeta, AccountMeta: accountMeta} +} diff --git a/internal/difftest/sweep_test.go b/internal/difftest/sweep_test.go index 55cbe162..2ae3a1db 100644 --- a/internal/difftest/sweep_test.go +++ b/internal/difftest/sweep_test.go @@ -50,20 +50,22 @@ func TestDifferentialSweep(t *testing.T) { for seed := range sweepSeeds { c := difftest.RunOne(context.Background(), rand.New(rand.NewSource(int64(seed)))) r.add(c) - if c.OracleVsNew.Tolerated != "" { - tolerated[c.OracleVsNew.Tolerated]++ + for _, leg := range c.Legs() { + if leg.Verdict.Tolerated != "" { + tolerated[leg.Name+": "+leg.Verdict.Tolerated]++ + } + if !leg.Verdict.Mismatch { + continue + } + k := leg.Name + ": " + divergenceClass(leg.Verdict.Reason) + cl, seen := classes[k] + if !seen { + cl = &class{firstSeed: seed, sampleWhy: leg.Verdict.Reason, sampleVars: c.Vars, sampleSrc: c.Script} + classes[k] = cl + order = append(order, k) + } + cl.count++ } - if !c.OracleVsNew.Mismatch { - continue - } - k := divergenceClass(c.OracleVsNew.Reason) - cl, seen := classes[k] - if !seen { - cl = &class{firstSeed: seed, sampleWhy: c.OracleVsNew.Reason, sampleVars: c.Vars, sampleSrc: c.Script} - classes[k] = cl - order = append(order, k) - } - cl.count++ } t.Log(r.table()) @@ -78,7 +80,7 @@ func TestDifferentialSweep(t *testing.T) { fmt.Fprintf(&sb, "%d divergence class(es) over %d seeds:\n", len(classes), sweepSeeds) for _, k := range order { cl := classes[k] - fmt.Fprintf(&sb, " %-44s %4d scripts (first: seed %d)\n", k, cl.count, cl.firstSeed) + fmt.Fprintf(&sb, " %-60s %4d scripts (first: seed %d)\n", k, cl.count, cl.firstSeed) } sb.WriteString("\nSee internal/oracle/DIVERGENCES.md.\n") @@ -119,7 +121,7 @@ func (r *reach) add(c difftest.Case) { } sh := c.Shape r.strategy[sh.Strategy]++ - if !c.New.Failed() && !c.Oracle.Failed() { + if !c.New.Failed() && !c.Oracle.Failed() && !c.VM.Failed() { r.completed++ } if c.Oracle.CompileErr != "" { @@ -136,7 +138,7 @@ func (r *reach) add(c difftest.Case) { count(&r.sameResource, sh.SaveOverdraftSameResource) count(&r.inOrder, sh.SaveOverdraftSameResourceInOrder) count(&r.overdraws, sh.SaveOverdrawsInitial) - count(&r.inOrderDiverg, sh.SaveOverdraftSameResourceInOrder && c.OracleVsNew.Mismatch) + count(&r.inOrderDiverg, sh.SaveOverdraftSameResourceInOrder && c.AnyMismatch()) count(&r.allotRemaining, sh.HasAllotmentRemaining) count(&r.portionVar, sh.HasPortionVar) count(&r.chainedOrigin, sh.HasChainedOrigin) @@ -154,7 +156,7 @@ func (r *reach) table() string { row("strategy: uniform", r.strategy[gen.StrategyUniform]) row("strategy: scenario + random", r.strategy[gen.StrategyScenarioMixed]) row("strategy: scenario only", r.strategy[gen.StrategyScenarioOnly]) - row("both engines ran to completion", r.completed) + row("all three engines ran to completion", r.completed) row("oracle rejected at compile time", r.oracleReject) row("scripts containing a save", r.save) row("containing a bounded overdraft", r.bounded) diff --git a/internal/difftest/testdata/fuzz/FuzzDiff/02464db166cce72c b/internal/difftest/testdata/fuzz/FuzzDiff/02464db166cce72c new file mode 100644 index 00000000..b162decd --- /dev/null +++ b/internal/difftest/testdata/fuzz/FuzzDiff/02464db166cce72c @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte(" /\xa3\xb6") diff --git a/internal/difftest/uncovered_shapes_test.go b/internal/difftest/uncovered_shapes_test.go index aac12854..0571d666 100644 --- a/internal/difftest/uncovered_shapes_test.go +++ b/internal/difftest/uncovered_shapes_test.go @@ -263,6 +263,7 @@ send [COIN *] ( ctx := context.Background() newRes := runNew(ctx, tc.script, tc.vars, tc.balances, tc.metadata, nil) oracleRes := runOracle(ctx, tc.script, tc.vars, tc.balances, tc.metadata) + vmRes := runVM(ctx, tc.script, tc.vars, tc.balances, tc.metadata, nil) v := Compare(newRes, oracleRes, "new interpreter", "oracle") if v.Mismatch { @@ -271,14 +272,17 @@ send [COIN *] ( if v.Tolerated != "" { t.Fatalf("nothing was compared (tolerated: %s)\nnew: %+v\noracle: %+v", v.Tolerated, newRes, oracleRes) } + if v := CompareEngines(vmRes, newRes); v.Mismatch { + t.Fatalf("vm mismatch: %s\nvm: %+v\nnew: %+v", v.Reason, vmRes, newRes) + } if tc.bothReject { - if !newRes.Failed() || !oracleRes.Failed() { - t.Fatalf("expected both engines to reject\nnew: %+v\noracle: %+v", newRes, oracleRes) + if !newRes.Failed() || !oracleRes.Failed() || !vmRes.Failed() { + t.Fatalf("expected all three engines to reject\nnew: %+v\noracle: %+v\nvm: %+v", newRes, oracleRes, vmRes) } return } - if newRes.Failed() || oracleRes.Failed() { - t.Fatalf("expected both engines to complete\nnew: %+v\noracle: %+v", newRes, oracleRes) + if newRes.Failed() || oracleRes.Failed() || vmRes.Failed() { + t.Fatalf("expected all three engines to complete\nnew: %+v\noracle: %+v\nvm: %+v", newRes, oracleRes, vmRes) } }) } @@ -315,4 +319,13 @@ func TestAllotmentFullSumPlusRemainingRejectedByOracleOnly(t *testing.T) { if newRes.Failed() { t.Fatalf("expected the interpreter to run the script; got %+v", newRes) } + + // The vm runs it like the interpreter: the remaining clause receives zero. + vmRes := runVM(context.Background(), script, nil, nil, nil, nil) + if vmRes.Failed() { + t.Fatalf("expected the vm to run the script; got %+v", vmRes) + } + if v := CompareEngines(vmRes, newRes); v.Mismatch { + t.Fatalf("vm mismatch: %s\nvm: %+v\nnew: %+v", v.Reason, vmRes, newRes) + } } diff --git a/internal/funds/funds.go b/internal/funds/funds.go index a729e49b..fd48ee19 100644 --- a/internal/funds/funds.go +++ b/internal/funds/funds.go @@ -25,6 +25,8 @@ package funds import ( "errors" "math/big" + "slices" + "strings" ) // ErrNegativePosting is returned by ForcePosting for a negative amount. A @@ -188,15 +190,31 @@ type AccountBalance struct { // entry outside the family can hold only a write delta whose base was never // prewarmed, and against a zero-backed Store loadBase would stamp it loaded, // masking the starting balance from every later read of that asset. +// +// Entries are visited in (asset, color) order, not map order, so both the +// returned slice and the sequence of Store reads — hence which error surfaces +// when several reads fail — are the same on every run. func (s *RunState) AccountBalances(account, scope, baseAsset string) ([]AccountBalance, error) { - var out []AccountBalance - for key, e := range s.balances { + var keys []PairKey + for key := range s.balances { if key.Account != account || key.Scope != scope { continue } if base, _ := GetBaseAndScale(key.Asset); base != baseAsset { continue } + keys = append(keys, key) + } + slices.SortFunc(keys, func(a, b PairKey) int { + if c := strings.Compare(a.Asset, b.Asset); c != 0 { + return c + } + return strings.Compare(a.Color, b.Color) + }) + + out := make([]AccountBalance, 0, len(keys)) + for _, key := range keys { + e := s.balances[key] if err := s.loadBase(key, e); err != nil { return nil, err } diff --git a/internal/funds/funds_test.go b/internal/funds/funds_test.go index afe2f400..90257689 100644 --- a/internal/funds/funds_test.go +++ b/internal/funds/funds_test.go @@ -1116,3 +1116,34 @@ func TestEndToEnd_TwoSourcesSplitAcrossDestinations(t *testing.T) { } } } + +// TestAccountBalances_DeterministicOrder: entries come back sorted by +// (asset, color) on every run, independent of the balances map's iteration +// order. +func TestAccountBalances_DeterministicOrder(t *testing.T) { + balances := map[funds.PairKey]*big.Int{ + {"A", "", "EUR/2", ""}: big.NewInt(1), + {"A", "", "EUR", "RED"}: big.NewInt(2), + {"A", "", "EUR", ""}: big.NewInt(3), + {"A", "", "EUR/4", ""}: big.NewInt(4), + {"A", "", "EUR/2", "RED"}: big.NewInt(5), + {"B", "", "EUR", ""}: big.NewInt(6), // other account + {"A", "", "USD", ""}: big.NewInt(7), // other family + } + want := []funds.AccountBalance{ + {Asset: "EUR", Color: "", Amount: big.NewInt(3)}, + {Asset: "EUR", Color: "RED", Amount: big.NewInt(2)}, + {Asset: "EUR/2", Color: "", Amount: big.NewInt(1)}, + {Asset: "EUR/2", Color: "RED", Amount: big.NewInt(5)}, + {Asset: "EUR/4", Color: "", Amount: big.NewInt(4)}, + } + + for range 50 { + rs, _ := newRS(nil) + rs.Prewarm(balances) + + got, err := rs.AccountBalances("A", "", "EUR") + require.NoError(t, err) + require.Equal(t, want, got) + } +} diff --git a/internal/funds/values.go b/internal/funds/values.go index b8298267..c8c2a09c 100644 --- a/internal/funds/values.go +++ b/internal/funds/values.go @@ -33,6 +33,15 @@ func ValidateAsset(v string) bool { return assetNameRegex.MatchString(v) } func ValidateColor(v string) bool { return v == "" || colorNameRegex.MatchString(v) } func ValidateScope(v string) bool { return scopeNameRegex.MatchString(v) } +// ValidatePosting reports whether a posting has a non-negative amount and +// well-formed accounts, scopes, asset and color. +func ValidatePosting(p Posting) bool { + return p.Amount != nil && p.Amount.Sign() >= 0 && + ValidateAccount(p.Source) && ValidateScope(p.SourceScope) && + ValidateAccount(p.Destination) && ValidateScope(p.DestinationScope) && + ValidateAsset(p.Asset) && ValidateColor(p.Color) +} + // ParseNumber parses a base-10 integer (arbitrary precision). func ParseNumber(s string) (*big.Int, bool) { return new(big.Int).SetString(s, 10) diff --git a/internal/funds/values_test.go b/internal/funds/values_test.go index 0899d3c7..c977fdb3 100644 --- a/internal/funds/values_test.go +++ b/internal/funds/values_test.go @@ -52,6 +52,32 @@ func TestValidate(t *testing.T) { } } +func TestValidatePosting(t *testing.T) { + valid := Posting{Source: "a", Destination: "b:c", Amount: big.NewInt(0), Asset: "USD/2"} + require.True(t, ValidatePosting(valid)) + + scoped := valid + scoped.SourceScope, scoped.DestinationScope, scoped.Color = "s1", "s_2", "RED" + require.True(t, ValidatePosting(scoped)) + + for name, mutate := range map[string]func(*Posting){ + "negative amount": func(p *Posting) { p.Amount = big.NewInt(-1) }, + "nil amount": func(p *Posting) { p.Amount = nil }, + "source": func(p *Posting) { p.Source = "a b" }, + "destination": func(p *Posting) { p.Destination = "" }, + "source scope": func(p *Posting) { p.SourceScope = "S" }, + "destination scope": func(p *Posting) { p.DestinationScope = "a-b" }, + "asset": func(p *Posting) { p.Asset = "usd" }, + "color": func(p *Posting) { p.Color = "red" }, + } { + t.Run(name, func(t *testing.T) { + p := valid + mutate(&p) + require.False(t, ValidatePosting(p)) + }) + } +} + func TestParseNumber(t *testing.T) { n, ok := ParseNumber("42") require.True(t, ok) diff --git a/internal/gen/ast.go b/internal/gen/ast.go index 64172765..9c43a3a6 100644 --- a/internal/gen/ast.go +++ b/internal/gen/ast.go @@ -68,7 +68,7 @@ type Source struct { // numerator always renders through a runtime-bound number var, so any sign is // expressible without unary minus). Unlike literals and portion vars, the value // is unconstrained: negative and over-one portions are reachable this way, and -// both engines must agree on what they do. numscript-only: the oracle's grammar +// both engines must reject them the same way. numscript-only: the oracle's grammar // has no division expression. type PortionDiv struct { Num *big.Int diff --git a/internal/gen/gen.go b/internal/gen/gen.go index c3d15327..6c0a748d 100644 --- a/internal/gen/gen.go +++ b/internal/gen/gen.go @@ -273,8 +273,8 @@ func srcColor(rng *rand.Rand, numscriptOnly bool) (string, bool) { } // divPortion maybe replaces one allotment clause's portion with a division -// expression `$n/den`, whose value is unconstrained: negative (a negative -// share), zero, ordinary, or above one. Only next to a `remaining` clause — +// expression `$n/den`, whose value is unconstrained: negative (rejected by both +// engines), zero, ordinary, or above one. Only next to a `remaining` clause — // without one the sum must be exactly 1 and nearly every draw would be a // same-on-both-engines rejection that compares nothing. func divPortion(rng *rand.Rand, numscriptOnly, withRemaining bool) *PortionDiv { diff --git a/internal/interpreter/interpreter.go b/internal/interpreter/interpreter.go index 037d7ecf..9d48a80c 100644 --- a/internal/interpreter/interpreter.go +++ b/internal/interpreter/interpreter.go @@ -103,24 +103,10 @@ func evaluateVarOrigin(env *evalEnv, type_ string, expr parser.ValueExpr) (Value return evaluateExpr(env, expr) } -// Check the following invariants: -// - no negative postings -// - no invalid account names -// - no invalid asset names -// - no invalid colors func checkPostingInvariants(posting Posting) InterpreterError { - isAmtNegative := posting.Amount.Cmp(big.NewInt(0)) == -1 - - isInvalidPosting := (isAmtNegative || - !funds.ValidateAsset(posting.Asset) || - !funds.ValidateColor(posting.Color) || - !funds.ValidateAccount(posting.Source) || - !funds.ValidateAccount(posting.Destination)) - - if isInvalidPosting { + if !funds.ValidatePosting(posting) { return InternalError{Posting: posting} } - return nil } @@ -864,6 +850,9 @@ func (s *programState) makeAllotment(monetary *big.Int, items []parser.Allotment if err != nil { return nil, err } + if rat.Sign() < 0 { + return nil, NegativePortion{Range: allotment.Value.GetRange(), Portion: *rat} + } totalAllotment.Add(totalAllotment, rat) allotments = append(allotments, rat) diff --git a/internal/interpreter/interpreter_error.go b/internal/interpreter/interpreter_error.go index 02820f8b..fc301228 100644 --- a/internal/interpreter/interpreter_error.go +++ b/internal/interpreter/interpreter_error.go @@ -217,6 +217,15 @@ func (e InvalidAllotmentSum) Error() string { return fmt.Sprintf("Invalid allotment: portions sum should be 1 (got %s instead)", e.ActualSum.String()) } +type NegativePortion struct { + parser.Range + Portion big.Rat +} + +func (e NegativePortion) Error() string { + return fmt.Sprintf("Invalid allotment: portions cannot be negative (got %s)", e.Portion.String()) +} + type QueryBalanceError struct { parser.Range WrappedError error diff --git a/internal/interpreter/interpreter_test.go b/internal/interpreter/interpreter_test.go index e0e5eeb9..cc40961c 100644 --- a/internal/interpreter/interpreter_test.go +++ b/internal/interpreter/interpreter_test.go @@ -594,6 +594,105 @@ func TestInvalidDestinationAllotmentSumOverOneWithRemaining(t *testing.T) { test(t, tc) } +func TestNegativeSourceAllotmentPortion(t *testing.T) { + tc := NewTestCase() + src := tc.compile(t, `vars { + number $n + } + + send [COIN 90] ( + source = { + $n/3 from @a + remaining from @b + } + destination = @dest + )`) + tc.setVarsFromJSON(t, `{"n": "-1"}`) + tc.setBalance("a", "COIN", 500) + tc.setBalance("b", "COIN", 500) + + tc.expected = CaseResult{ + Error: interpreter.NegativePortion{ + Range: parser.RangeOfIndexed(src, "$n/3", 0), + Portion: *big.NewRat(-1, 3), + }, + } + test(t, tc) +} + +func TestNegativeDestinationAllotmentPortion(t *testing.T) { + tc := NewTestCase() + src := tc.compile(t, `vars { + number $n + } + + send [COIN 90] ( + source = @world + destination = { + $n/3 to @a + remaining to @b + } + )`) + tc.setVarsFromJSON(t, `{"n": "-1"}`) + + tc.expected = CaseResult{ + Error: interpreter.NegativePortion{ + Range: parser.RangeOfIndexed(src, "$n/3", 0), + Portion: *big.NewRat(-1, 3), + }, + } + test(t, tc) +} + +func TestNegativeAllotmentPortionBalancedToOne(t *testing.T) { + // the portions sum to 1, so only the per-clause check rejects it + tc := NewTestCase() + src := tc.compile(t, `vars { + number $n + number $m + } + + send [COIN 90] ( + source = @world + destination = { + $m/3 to @a + $n/3 to @b + } + )`) + tc.setVarsFromJSON(t, `{"n": "-1", "m": "4"}`) + + tc.expected = CaseResult{ + Error: interpreter.NegativePortion{ + Range: parser.RangeOfIndexed(src, "$n/3", 0), + Portion: *big.NewRat(-1, 3), + }, + } + test(t, tc) +} + +func TestZeroAllotmentPortion(t *testing.T) { + tc := NewTestCase() + tc.compile(t, `vars { + number $n + } + + send [COIN 90] ( + source = @world + destination = { + $n/3 to @a + remaining to @b + } + )`) + tc.setVarsFromJSON(t, `{"n": "0"}`) + + tc.expected = CaseResult{ + Postings: []Posting{ + {Asset: "COIN", Amount: big.NewInt(90), Source: "world", Destination: "b"}, + }, + } + test(t, tc) +} + func TestRejectsDuplicateRemainingAllotments(t *testing.T) { tc := NewTestCase() src := tc.compile(t, `send [COIN 100] ( diff --git a/internal/interpreter/testdata/script-tests/dest-allot-negative-portion-expr.num b/internal/interpreter/testdata/script-tests/dest-allot-negative-portion-expr.num deleted file mode 100644 index 64d3e6b8..00000000 --- a/internal/interpreter/testdata/script-tests/dest-allot-negative-portion-expr.num +++ /dev/null @@ -1,11 +0,0 @@ -vars { - number $n -} - -send [COIN 90] ( - source = @world - destination = { - $n/3 to @a - remaining to @b - } -) diff --git a/internal/interpreter/testdata/script-tests/dest-allot-negative-portion-expr.num.specs.json b/internal/interpreter/testdata/script-tests/dest-allot-negative-portion-expr.num.specs.json deleted file mode 100644 index e8dfed6c..00000000 --- a/internal/interpreter/testdata/script-tests/dest-allot-negative-portion-expr.num.specs.json +++ /dev/null @@ -1,19 +0,0 @@ -{ - "$schema": "https://raw.githubusercontent.com/formancehq/numscript/main/v1.specs.schema.json", - "testCases": [ - { - "it": "a negative allotment share receives nothing and the remaining share is capped by the sent amount", - "variables": { - "n": "-1" - }, - "expect.postings": [ - { - "source": "world", - "destination": "b", - "amount": 90, - "asset": "COIN" - } - ] - } - ] -} diff --git a/internal/interpreter/testdata/script-tests/src-allot-negative-portion-expr.num b/internal/interpreter/testdata/script-tests/src-allot-negative-portion-expr.num deleted file mode 100644 index 47aa1d0e..00000000 --- a/internal/interpreter/testdata/script-tests/src-allot-negative-portion-expr.num +++ /dev/null @@ -1,11 +0,0 @@ -vars { - number $n -} - -send [COIN 90] ( - source = { - $n/3 from @a - remaining from @b - } - destination = @d -) diff --git a/internal/interpreter/testdata/script-tests/src-allot-negative-portion-expr.num.specs.json b/internal/interpreter/testdata/script-tests/src-allot-negative-portion-expr.num.specs.json deleted file mode 100644 index ce8666c5..00000000 --- a/internal/interpreter/testdata/script-tests/src-allot-negative-portion-expr.num.specs.json +++ /dev/null @@ -1,16 +0,0 @@ -{ - "$schema": "https://raw.githubusercontent.com/formancehq/numscript/main/v1.specs.schema.json", - "testCases": [ - { - "it": "a negative allotment share fails as missing funds, not silently over-pulling the remaining share", - "variables": { - "n": "-1" - }, - "balances": [ - { "account": "a", "asset": "COIN", "amount": 500 }, - { "account": "b", "asset": "COIN", "amount": 500 } - ], - "expect.error.missingFunds": true - } - ] -} diff --git a/internal/ir/assemble.go b/internal/ir/assemble.go new file mode 100644 index 00000000..755a5f1d --- /dev/null +++ b/internal/ir/assemble.go @@ -0,0 +1,1063 @@ +package ir + +import ( + "errors" + "fmt" + "math" + "math/big" + "slices" + + "github.com/formancehq/numscript/internal/vm" +) + +const maxReg = 0xFF + +// ErrRegisterBankOverflow is the assembler refusing a program that needs more +// than maxReg simultaneously-live registers in one bank — the instruction +// encoding's one-byte operand limit. It is a capacity bound of the compiled +// engine, not a semantic rejection of the script. +var ErrRegisterBankOverflow = errors.New("register bank overflow") + +// ErrProgramTooLarge is the assembler refusing a program whose instruction +// stream outgrows what a jump target can address — the encoding's uint16 +// operand limit. Like ErrRegisterBankOverflow, a capacity bound of the +// compiled engine, not a semantic rejection of the script. +var ErrProgramTooLarge = errors.New("program too large") + +// regPool assigns each virtual Reg a physical bank index (0..maxReg-1), +// reusing an index once its Reg's last reference has been processed. +// +// Correctness of a plain linear scan (rather than full dataflow liveness +// analysis) hinges on IR control flow being forward-jump-only, which +// Assemble's patch-resolution pass already enforces ("backward jump to +// label" is a hard error): program order is then a superset of every +// possible runtime execution order, so a Reg's last *textual* reference is +// also its true last possible use, and freeing its slot right after is +// always safe — nothing later in the stream can reach back to read it. +// +// Assemble runs the instruction list through this twice: once in scanning +// mode (via newScanRegPool) purely to record each Reg's last-referenced +// instruction position, discarding everything else it produces, and once +// for real (via newAllocRegPool, seeded with that position map) to hand out +// and reuse physical slots. Freeing is deferred to instruction boundaries +// (endInstr), not individual references, so that two unrelated Regs which +// both happen to make their last appearance in the same instruction (e.g. +// dest and an operand) never get aliased onto the same physical slot before +// that instruction has finished reading them. +type regPool struct { + scanning bool + curPos int // current instruction position, set by the assembler driver + + // scanning mode: last instruction position, so far, referencing each Reg + lastUse map[Reg]int + + // allocation mode + indexByReg map[Reg]byte + freeList []byte + next int + pendingFree map[Reg]struct{} // regs to free once endInstr() runs for curPos +} + +func newScanRegPool() regPool { + return regPool{scanning: true, lastUse: map[Reg]int{}} +} + +func newAllocRegPool(lastUse map[Reg]int) regPool { + return regPool{ + lastUse: lastUse, + indexByReg: map[Reg]byte{}, + pendingFree: map[Reg]struct{}{}, + } +} + +// endInstr frees every register whose last reference was curPos. Called by +// the assembler driver after an instruction has been fully assembled. +// +// The freed slots are pushed in ascending Reg order, not map order: freeList +// is popped as a stack, so the push order decides which physical slot the next +// allocation reuses, and ranging over the map directly made the same program +// assemble to different (equivalent) bytecode from one run to the next. +func (b *regPool) endInstr() { + if len(b.pendingFree) == 0 { + return + } + + regs := make([]Reg, 0, len(b.pendingFree)) + for r := range b.pendingFree { + regs = append(regs, r) + } + slices.Sort(regs) + + for _, r := range regs { + b.freeList = append(b.freeList, b.indexByReg[r]) + delete(b.indexByReg, r) + delete(b.pendingFree, r) + } +} + +type constPool[T any] struct { + indexByValue map[string]uint16 + items []T + toString func(T) string +} + +func newConstPool[T any](toString func(T) string) constPool[T] { + return constPool[T]{ + indexByValue: map[string]uint16{}, + toString: toString, + } +} + +func (p *constPool[T]) alloc(item T) (uint16, error) { + strValue := p.toString(item) + index, ok := p.indexByValue[strValue] + if !ok { + l := len(p.items) + if l > math.MaxUint16 { + return 0, fmt.Errorf("error: too many consts (overflowed the u16 len)") + } + index = uint16(l) + p.indexByValue[strValue] = index + p.items = append(p.items, item) + } + + return index, nil +} + +func (b *regPool) Index(r Reg) (byte, error) { + if b.scanning { + b.lastUse[r] = b.curPos + return 0, nil + } + + idx, ok := b.indexByReg[r] + if !ok { + if n := len(b.freeList); n > 0 { + idx = b.freeList[n-1] + b.freeList = b.freeList[:n-1] + } else { + if b.next >= maxReg { + return 0, fmt.Errorf("%w: more than %d registers live at once in one bank", ErrRegisterBankOverflow, maxReg) + } + idx = byte(b.next) + b.next++ + } + b.indexByReg[r] = idx + } + + if b.lastUse[r] == b.curPos { + b.pendingFree[r] = struct{}{} + } + + return idx, nil +} + +type patch struct { + Label Label + index int + // delta is the forward jump offset, relative to the instruction following index + getInstruction func(delta uint16) vm.Instruction +} + +// assembler lowers IR instructions into a vm.Program. +type assembler struct { + instructions []vm.Instruction + + patches []patch + labels map[Label]uint16 + + // one register bank per VM register bank + ints regPool + strings regPool + Portions regPool + bools regPool + + intsPool constPool[big.Int] + stringsPool constPool[string] +} + +// newAssembler builds an assembler over fresh int/string/portion/bool +// register pools — scanning ones if lastUse is nil (see regPool's doc +// comment), real allocating ones (seeded with lastUse) otherwise. +func newAssembler(lastUse *regUsage) *assembler { + a := &assembler{ + labels: map[Label]uint16{}, + + intsPool: newConstPool(func(i big.Int) string { + return i.String() + }), + stringsPool: newConstPool(func(s string) string { + return s + }), + } + + if lastUse == nil { + a.ints, a.strings, a.Portions, a.bools = newScanRegPool(), newScanRegPool(), newScanRegPool(), newScanRegPool() + } else { + a.ints = newAllocRegPool(lastUse.ints) + a.strings = newAllocRegPool(lastUse.strings) + a.Portions = newAllocRegPool(lastUse.portions) + a.bools = newAllocRegPool(lastUse.bools) + } + + return a +} + +// setPos moves every bank's current instruction position, used to key +// live-range tracking (see regPool's doc comment). +func (a *assembler) setPos(pos int) { + a.ints.curPos = pos + a.strings.curPos = pos + a.Portions.curPos = pos + a.bools.curPos = pos +} + +// endInstr frees, in every bank, any register whose last reference was the +// instruction just assembled. +func (a *assembler) endInstr() { + a.ints.endInstr() + a.strings.endInstr() + a.Portions.endInstr() + a.bools.endInstr() +} + +type regUsage struct { + ints, strings, portions, bools map[Reg]int +} + +// scanRegUsage dry-runs instrs through the same assemble() dispatch used for +// real emission, to learn each Reg's last-referenced instruction position +// per bank. Everything else the dry run produces (instructions, patches, +// labels, const pools) is discarded — only the position maps survive. +func scanRegUsage(instrs []Instr) (regUsage, error) { + scan := newAssembler(nil) + for i, instr := range instrs { + scan.setPos(i) + if err := instr.assemble(scan); err != nil { + return regUsage{}, err + } + // no endInstr(): scanning mode never frees, it only records + } + return regUsage{ + ints: scan.ints.lastUse, + strings: scan.strings.lastUse, + portions: scan.Portions.lastUse, + bools: scan.bools.lastUse, + }, nil +} + +func Assemble(instrs []Instr) (vm.Program, error) { + lastUse, err := scanRegUsage(instrs) + if err != nil { + return vm.Program{}, err + } + + a := newAssembler(&lastUse) + for i, instr := range instrs { + a.setPos(i) + if err := instr.assemble(a); err != nil { + return vm.Program{}, err + } + a.endInstr() + } + + // now we run the patches + for _, patch := range a.patches { + labelIndex, ok := a.labels[patch.Label] + if !ok { + return vm.Program{}, fmt.Errorf("missing label declaration of `%s`", string(patch.Label)) + } + + next := patch.index + 1 + if int(labelIndex) < next { + return vm.Program{}, fmt.Errorf("backward jump to label `%s`: jumps must go forward", string(patch.Label)) + } + + a.instructions[patch.index] = patch.getInstruction(uint16(int(labelIndex) - next)) + } + + return vm.Program{ + Instructions: a.instructions, + StringsPool: a.stringsPool.items, + IntsPool: a.intsPool.items, + + MaxRegString: byte(a.strings.next), + MaxRegPortion: byte(a.Portions.next), + MaxRegInt: byte(a.ints.next), + MaxRegBool: byte(a.bools.next), + + Version: vm.CurrentBytecodeVersion, + }, nil +} + +func (as *assembler) intReg(r Reg) (byte, error) { return as.ints.Index(r) } +func (as *assembler) strReg(r Reg) (byte, error) { return as.strings.Index(r) } +func (as *assembler) portionReg(r Reg) (byte, error) { return as.Portions.Index(r) } +func (as *assembler) boolReg(r Reg) (byte, error) { return as.bools.Index(r) } + +func (as *assembler) optionalReg( + regPool func(*assembler, Reg) (byte, error), + Reg *Reg, +) (byte, error) { + if Reg == nil { + return maxReg, nil + } else { + reg_, err := regPool(as, *Reg) + if err != nil { + return 0, err + } + return reg_, nil + } + +} + +func (as *assembler) emit(op vm.Opcode, a, b, c byte) { + as.instructions = append(as.instructions, vm.Instruction{ + Opcode: byte(op), + A: a, + B: b, + C: c, + }) +} + +func (as *assembler) emitBC(op vm.Opcode, a byte, bc uint16) { + as.instructions = append(as.instructions, vm.NewBC(op, a, bc)) +} + +// regResolver maps a virtual register to a concrete bank index. op sigs hold +// these as method expressions ((*assembler).intReg, ...) so that a sig is a +// static description of an op, independent of any assembler instance. +type regResolver = func(*assembler, Reg) (byte, error) + +type unaryOpSig struct { + opcode vm.Opcode + dest regResolver + arg regResolver +} + +// copySig is the shape every bank copy shares: dest and src in the same bank. +func copySig(opcode vm.Opcode, bank regResolver) unaryOpSig { + return unaryOpSig{opcode: opcode, dest: bank, arg: bank} +} + +func (OpIntCopy) sig() unaryOpSig { return copySig(vm.Op_IntCopy, (*assembler).intReg) } +func (OpPortionCopy) sig() unaryOpSig { + return copySig(vm.Op_PortionCopy, (*assembler).portionReg) +} +func (OpStrCopy) sig() unaryOpSig { return copySig(vm.Op_StrCopy, (*assembler).strReg) } +func (OpBoolCopy) sig() unaryOpSig { return copySig(vm.Op_BoolCopy, (*assembler).boolReg) } +func (OpNegInt) sig() unaryOpSig { + return unaryOpSig{ + opcode: vm.Op_NegInt, + dest: (*assembler).intReg, + arg: (*assembler).intReg, + } +} +func (OpIntToString) sig() unaryOpSig { + return unaryOpSig{ + opcode: vm.Op_IntToString, + dest: (*assembler).strReg, + arg: (*assembler).intReg, + } +} +func (OpIsZero) sig() unaryOpSig { + return unaryOpSig{ + opcode: vm.Op_IsZero, + dest: (*assembler).boolReg, + arg: (*assembler).intReg, + } +} +func (OpNot) sig() unaryOpSig { + return unaryOpSig{ + opcode: vm.Op_Not, + dest: (*assembler).boolReg, + arg: (*assembler).boolReg, + } +} +func (OpPortionToString) sig() unaryOpSig { + return unaryOpSig{ + opcode: vm.Op_PortionToString, + dest: (*assembler).strReg, + arg: (*assembler).portionReg, + } +} +func (OpIntToPortion) sig() unaryOpSig { + return unaryOpSig{ + opcode: vm.Op_IntToPortion, + dest: (*assembler).portionReg, + arg: (*assembler).intReg, + } +} +func (OpPortionToInt) sig() unaryOpSig { + return unaryOpSig{ + opcode: vm.Op_PortionToInt, + dest: (*assembler).intReg, + arg: (*assembler).portionReg, + } +} +func (i UnaryOp) assemble(a *assembler) error { + sig := i.Op.sig() + + dest, err := sig.dest(a, i.Dest) + if err != nil { + return err + } + arg, err := sig.arg(a, i.Arg) + if err != nil { + return err + } + + a.emit(sig.opcode, dest, arg, maxReg) + return nil +} + +type binaryOpSig struct { + opcode vm.Opcode + dest regResolver + left regResolver + right regResolver +} + +// comparisonSig is the shape every binary comparison shares: two operands of one +// bank, a bool dest. +func comparisonSig(opcode vm.Opcode, operand regResolver) binaryOpSig { + return binaryOpSig{ + opcode: opcode, + dest: (*assembler).boolReg, + left: operand, + right: operand, + } +} + +func (OpLtInt) sig() binaryOpSig { return comparisonSig(vm.Op_LtInt, (*assembler).intReg) } +func (OpEqInt) sig() binaryOpSig { return comparisonSig(vm.Op_EqInt, (*assembler).intReg) } +func (OpLtPortion) sig() binaryOpSig { + return comparisonSig(vm.Op_LtPortion, (*assembler).portionReg) +} +func (OpEqPortion) sig() binaryOpSig { + return comparisonSig(vm.Op_EqPortion, (*assembler).portionReg) +} +func (OpAddInt) sig() binaryOpSig { + return binaryOpSig{ + opcode: vm.Op_AddInt, + dest: (*assembler).intReg, + left: (*assembler).intReg, + right: (*assembler).intReg, + } +} +func (OpSubInt) sig() binaryOpSig { + return binaryOpSig{ + opcode: vm.Op_SubInt, + dest: (*assembler).intReg, + left: (*assembler).intReg, + right: (*assembler).intReg, + } +} +func (OpAddString) sig() binaryOpSig { + return binaryOpSig{ + opcode: vm.Op_AddString, + dest: (*assembler).strReg, + left: (*assembler).strReg, + right: (*assembler).strReg, + } +} +func (OpStrEq) sig() binaryOpSig { + return binaryOpSig{ + opcode: vm.Op_StrEq, + dest: (*assembler).boolReg, + left: (*assembler).strReg, + right: (*assembler).strReg, + } +} + +// portionArithSig is the shape of portion addition and subtraction: three +// portion operands. +func portionArithSig(opcode vm.Opcode) binaryOpSig { + return binaryOpSig{ + opcode: opcode, + dest: (*assembler).portionReg, + left: (*assembler).portionReg, + right: (*assembler).portionReg, + } +} + +func (OpAddPortion) sig() binaryOpSig { return portionArithSig(vm.Op_AddPortion) } +func (OpSubPortion) sig() binaryOpSig { return portionArithSig(vm.Op_SubPortion) } +func (OpMulPortion) sig() binaryOpSig { return portionArithSig(vm.Op_MulPortion) } +func (OpMakePortion) sig() binaryOpSig { + return binaryOpSig{ + opcode: vm.Op_MkPortion, + dest: (*assembler).portionReg, + left: (*assembler).intReg, + right: (*assembler).intReg, + } +} +func (OpMonetaryToString) sig() binaryOpSig { + return binaryOpSig{ + opcode: vm.Op_MonetaryToString, + dest: (*assembler).strReg, + left: (*assembler).strReg, + right: (*assembler).intReg, + } +} + +func (i BinaryOp) assemble(a *assembler) error { + sig := i.Op.sig() + + dest, err := sig.dest(a, i.Dest) + if err != nil { + return err + } + left, err := sig.left(a, i.Left) + if err != nil { + return err + } + right, err := sig.right(a, i.Right) + if err != nil { + return err + } + + a.emit(sig.opcode, dest, left, right) + return nil +} + +func (i LoadInt) assemble(a *assembler) error { + dest, err := a.intReg(i.Dest) + if err != nil { + return err + } + + poolIndex, err := a.intsPool.alloc(i.Value) + if err != nil { + return err + } + + a.emitBC(vm.Op_LoadInt, dest, poolIndex) + return nil +} + +func (i LoadStr) assemble(a *assembler) error { + dest, err := a.strReg(i.Dest) + if err != nil { + return err + } + + poolIndex, err := a.stringsPool.alloc(i.Value) + if err != nil { + return err + } + + a.emitBC(vm.Op_LoadStr, dest, poolIndex) + return nil +} + +func (i ConstBool) assemble(a *assembler) error { + dest, err := a.boolReg(i.Dest) + if err != nil { + return err + } + + opcode := vm.Op_ConstFalse + if i.Value { + opcode = vm.Op_ConstTrue + } + + a.emit(opcode, dest, maxReg, maxReg) + return nil +} + +func (i CheckEnoughFunds) assemble(a *assembler) error { + got, err := a.intReg(i.Got) + if err != nil { + return err + } + + needed, err := a.intReg(i.Needed) + if err != nil { + return err + } + + a.emit(vm.Op_CheckEnoughFunds, got, needed, maxReg) + return nil +} + +// Save needs a fourth operand (scope) but its base word's three slots are all +// taken, so it spills scope into an ext word, like PullAccount. +func (i Save) assemble(a *assembler) error { + account, err := a.strReg(i.Account) + if err != nil { + return err + } + asset, err := a.strReg(i.Asset) + if err != nil { + return err + } + amount, err := a.optionalReg((*assembler).intReg, i.Amount) + if err != nil { + return err + } + scope, err := a.optionalReg((*assembler).strReg, i.Scope) + if err != nil { + return err + } + + a.emit(vm.Op_Save, account, asset, amount) + a.instructions = append(a.instructions, vm.Instruction{ + Opcode: maxReg, + A: scope, + B: maxReg, + C: maxReg, + }) + return nil +} + +func (i AssertLeftover) assemble(a *assembler) error { + portion, err := a.portionReg(i.Portion) + if err != nil { + return err + } + var exact byte + if i.Exact { + exact = 1 + } + a.emit(vm.Op_AssertLeftover, portion, exact, maxReg) + return nil +} + +func (i SetCurrentAsset) assemble(a *assembler) error { + assetReg, err := a.strReg(i.Asset) + if err != nil { + return err + } + + a.emit(vm.Op_SetCurrentAsset, assetReg, maxReg, maxReg) + return nil +} + +func (i PullAccount) assemble(a *assembler) error { + dest, err := a.intReg(i.Dest) + if err != nil { + return err + } + + account, err := a.strReg(i.Account) + if err != nil { + return err + } + + cap, err := a.optionalReg((*assembler).intReg, i.Cap) + if err != nil { + return err + } + + overdraft, err := a.optionalReg((*assembler).intReg, i.Overdraft) + if err != nil { + return err + } + + color, err := a.optionalReg((*assembler).strReg, i.Color) + if err != nil { + return err + } + + scope, err := a.optionalReg((*assembler).strReg, i.Scope) + if err != nil { + return err + } + + a.emit(vm.Op_PullAccount, dest, account, cap) + + a.instructions = append(a.instructions, vm.Instruction{ + Opcode: maxReg, // <- UNUSED + A: overdraft, // overdraft (int) + B: color, // color (str) + C: scope, // scope (str) + }) + + return nil +} + +func (i SendToAccount) assemble(a *assembler) error { + account, err := a.optionalReg((*assembler).strReg, i.Account) + if err != nil { + return err + } + + cap, err := a.optionalReg((*assembler).intReg, i.Cap) + if err != nil { + return err + } + + scope, err := a.optionalReg((*assembler).strReg, i.Scope) + if err != nil { + return err + } + + a.emit(vm.Op_SendToAccount, account, cap, scope) + return nil +} + +func (i AssertSameAsset) assemble(a *assembler) error { + left, err := a.strReg(i.Left) + if err != nil { + return err + } + right, err := a.strReg(i.Right) + if err != nil { + return err + } + + a.emit(vm.Op_AssertSameAsset, left, right, maxReg) + + return nil +} + +func (i AssertValidAccount) assemble(a *assembler) error { + account, err := a.strReg(i.Account) + if err != nil { + return err + } + + a.emit(vm.Op_AssertValidAccount, account, maxReg, maxReg) + + return nil +} + +func (i AssertValidColor) assemble(a *assembler) error { + color, err := a.strReg(i.Color) + if err != nil { + return err + } + + a.emit(vm.Op_AssertValidColor, color, maxReg, maxReg) + + return nil +} + +func (i AssertValidScope) assemble(a *assembler) error { + scope, err := a.strReg(i.Scope) + if err != nil { + return err + } + + a.emit(vm.Op_AssertValidScope, scope, maxReg, maxReg) + + return nil +} + +func (i AssertUnscoped) assemble(a *assembler) error { + scope, err := a.strReg(i.Scope) + if err != nil { + return err + } + account, err := a.strReg(i.Account) + if err != nil { + return err + } + + a.emit(vm.Op_AssertUnscoped, scope, account, maxReg) + + return nil +} + +func (i AssertNonNegativeBalance) assemble(a *assembler) error { + balance, err := a.intReg(i.Balance) + if err != nil { + return err + } + account, err := a.strReg(i.Account) + if err != nil { + return err + } + + a.emit(vm.Op_AssertNonNegativeBalance, balance, account, maxReg) + + return nil +} + +func (i AssertNonNegativeAmount) assemble(a *assembler) error { + amount, err := a.intReg(i.Amount) + if err != nil { + return err + } + + a.emit(vm.Op_AssertNonNegativeAmount, amount, maxReg, maxReg) + + return nil +} + +func (i AssertNonNegativePortion) assemble(a *assembler) error { + portion, err := a.portionReg(i.Portion) + if err != nil { + return err + } + + a.emit(vm.Op_AssertNonNegativePortion, portion, maxReg, maxReg) + + return nil +} + +func (i SetTxMeta) assemble(a *assembler) error { + key, err := a.strReg(i.Key) + if err != nil { + return err + } + value, err := a.strReg(i.Value) + if err != nil { + return err + } + + a.emit(vm.Op_SetTxMeta, key, value, maxReg) + + return nil +} + +// emitMeta emits the shared base word for a meta(account, key) read (dest, +// account, key) plus an ext word carrying scope, shared by MetaStr/Int/Portion. +// MetaMonetary needs its ext word for a second destination too, so it builds its +// own instead of sharing this helper. +func (a *assembler) emitMeta(opcode vm.Opcode, dest byte, account, key Reg, scope *Reg) error { + acc, err := a.strReg(account) + if err != nil { + return err + } + k, err := a.strReg(key) + if err != nil { + return err + } + s, err := a.optionalReg((*assembler).strReg, scope) + if err != nil { + return err + } + + a.emit(opcode, dest, acc, k) + a.instructions = append(a.instructions, vm.Instruction{ + Opcode: maxReg, + A: s, + B: maxReg, + C: maxReg, + }) + return nil +} + +func (MetaStr) assembleMeta(a *assembler, dest, account, key Reg, scope *Reg) error { + d, err := a.strReg(dest) + if err != nil { + return err + } + return a.emitMeta(vm.Op_MetaStr, d, account, key, scope) +} + +func (MetaInt) assembleMeta(a *assembler, dest, account, key Reg, scope *Reg) error { + d, err := a.intReg(dest) + if err != nil { + return err + } + return a.emitMeta(vm.Op_MetaInt, d, account, key, scope) +} + +func (MetaPortion) assembleMeta(a *assembler, dest, account, key Reg, scope *Reg) error { + d, err := a.portionReg(dest) + if err != nil { + return err + } + return a.emitMeta(vm.Op_MetaPortion, d, account, key, scope) +} + +func (i MetaVar) assemble(a *assembler) error { + return i.Typ.assembleMeta(a, i.Dest, i.Account, i.Key, i.Scope) +} + +// MetaMonetary needs four operands plus scope, so it spills the second +// destination and scope into its own ext word rather than sharing emitMeta +// (which would otherwise burn a second, conflicting ext word). +func (i MetaMonetary) assemble(a *assembler) error { + destAsset, err := a.strReg(i.DestAsset) + if err != nil { + return err + } + destAmount, err := a.intReg(i.DestAmount) + if err != nil { + return err + } + account, err := a.strReg(i.Account) + if err != nil { + return err + } + key, err := a.strReg(i.Key) + if err != nil { + return err + } + scope, err := a.optionalReg((*assembler).strReg, i.Scope) + if err != nil { + return err + } + + a.emit(vm.Op_MetaMonetary, destAsset, account, key) + a.instructions = append(a.instructions, vm.Instruction{ + Opcode: maxReg, + A: destAmount, + B: scope, + C: maxReg, + }) + return nil +} + +// SetAccountMeta needs a fourth operand (scope) but its base word's three slots +// are all taken, so it spills scope into an ext word, like Save. +func (i SetAccountMeta) assemble(a *assembler) error { + account, err := a.strReg(i.Account) + if err != nil { + return err + } + key, err := a.strReg(i.Key) + if err != nil { + return err + } + value, err := a.strReg(i.Value) + if err != nil { + return err + } + scope, err := a.optionalReg((*assembler).strReg, i.Scope) + if err != nil { + return err + } + + a.emit(vm.Op_SetAccountMeta, account, key, value) + a.instructions = append(a.instructions, vm.Instruction{ + Opcode: maxReg, + A: scope, + B: maxReg, + C: maxReg, + }) + + return nil +} + +// FetchBalance needs a fourth operand (scope) but its base word's three slots +// are all taken, so it spills scope into an ext word, like Save. +func (i FetchBalance) assemble(a *assembler) error { + dest, err := a.intReg(i.Dest) + if err != nil { + return err + } + account, err := a.strReg(i.Account) + if err != nil { + return err + } + asset, err := a.strReg(i.Asset) + if err != nil { + return err + } + scope, err := a.optionalReg((*assembler).strReg, i.Scope) + if err != nil { + return err + } + + a.emit(vm.Op_Balance, dest, account, asset) + a.instructions = append(a.instructions, vm.Instruction{ + Opcode: maxReg, + A: scope, + B: maxReg, + C: maxReg, + }) + + return nil +} + +// assembleCondJmp emits either conditional jump: they differ only in the opcode. +func (a *assembler) assembleCondJmp(opcode vm.Opcode, cond Reg, target Label) error { + condReg, err := a.boolReg(cond) + if err != nil { + return err + } + + a.patches = append(a.patches, patch{ + Label: target, + index: len(a.instructions), + getInstruction: func(delta uint16) vm.Instruction { + return vm.NewBC(opcode, condReg, delta) + }, + }) + + // Emit dummy instruction + a.emit(0, 0, 0, 0) + + return nil +} + +func (i JmpIfFalse) assemble(a *assembler) error { + return a.assembleCondJmp(vm.Op_JmpIfFalse, i.Cond, i.Target) +} + +func (i JmpIfTrue) assemble(a *assembler) error { + return a.assembleCondJmp(vm.Op_JmpIfTrue, i.Cond, i.Target) +} + +func (i Jmp) assemble(a *assembler) error { + a.patches = append(a.patches, patch{ + Label: i.Target, + index: len(a.instructions), + getInstruction: func(delta uint16) vm.Instruction { + return vm.NewBC(vm.Op_Jmp, 0, delta) + }, + }) + + // Emit dummy instruction + a.emit(0, 0, 0, 0) + + return nil +} + +func (VarInt) assembleLoad(a *assembler, dest Reg, index uint16) error { + d, err := a.intReg(dest) + if err != nil { + return err + } + a.emitBC(vm.Op_LoadVarInt, d, index) + return nil +} + +func (VarStr) assembleLoad(a *assembler, dest Reg, index uint16) error { + d, err := a.strReg(dest) + if err != nil { + return err + } + a.emitBC(vm.Op_LoadVarStr, d, index) + return nil +} + +func (i LoadVar) assemble(a *assembler) error { + return i.Typ.assembleLoad(a, i.Dest, i.Index) +} + +func (i LabelMarker) assemble(a *assembler) error { + l := len(a.instructions) + if l > math.MaxUint16 { + return fmt.Errorf("%w: a label sits past instruction %d, the most a jump target can address", ErrProgramTooLarge, math.MaxUint16) + } + + if _, ok := a.labels[i.Label]; ok { + return fmt.Errorf("duplicate label declaration of `%s`", string(i.Label)) + } + a.labels[i.Label] = uint16(l) + + return nil +} + +// the mark ops take no register, so there is nothing to allocate and no way to +// fail here + +func (i MarkPush) assemble(a *assembler) error { + a.emit(vm.Op_MarkPush, maxReg, maxReg, maxReg) + return nil +} + +func (i MarkEnd) assemble(a *assembler) error { + var rewind byte + if i.Rewind { + rewind = 1 + } + a.emit(vm.Op_MarkEnd, rewind, maxReg, maxReg) + return nil +} diff --git a/internal/ir/assemble_test.go b/internal/ir/assemble_test.go new file mode 100644 index 00000000..ac7c398c --- /dev/null +++ b/internal/ir/assemble_test.go @@ -0,0 +1,252 @@ +package ir + +import ( + "math/big" + "testing" + + "github.com/formancehq/numscript/internal/vm" + "github.com/stretchr/testify/require" +) + +func TestAssemble_AddInt(t *testing.T) { + // Three distinct virtual int registers map to the first three int-bank + // indices in first-use order. + prog, err := Assemble([]Instr{ + BinaryOp{Op: OpAddInt{}, Dest: 10, Left: 20, Right: 30}, + }) + if err != nil { + t.Fatalf("Assemble: %v", err) + } + + instrs := prog.Instructions + if len(instrs) != 1 { + t.Fatalf("got %d instructions, want 1", len(instrs)) + } + want := vm.Instruction{Opcode: byte(vm.Op_AddInt), A: 0, B: 1, C: 2} + if instrs[0] != want { + t.Errorf("got %+v, want %+v", instrs[0], want) + } +} + +func TestAssemble_AddInt_ReusesRegisterIndices(t *testing.T) { + // A virtual register reused across operands/instructions keeps the same + // bank index; new ones get fresh indices in first-use order. + prog, err := Assemble([]Instr{ + // Reg 7 -> 0, Reg 8 -> 1 ; dest==left==7 + BinaryOp{Op: OpAddInt{}, Dest: 7, Left: 7, Right: 8}, + // Reg 9 -> 2 ; reuses 7->0 and 8->1 + BinaryOp{Op: OpAddInt{}, Dest: 9, Left: 7, Right: 8}, + }) + if err != nil { + t.Fatalf("Assemble: %v", err) + } + + got := prog.Instructions + want := []vm.Instruction{ + {Opcode: byte(vm.Op_AddInt), A: 0, B: 0, C: 1}, + {Opcode: byte(vm.Op_AddInt), A: 2, B: 0, C: 1}, + } + if len(got) != len(want) { + t.Fatalf("got %d instructions, want %d", len(got), len(want)) + } + for i := range want { + if got[i] != want[i] { + t.Errorf("instr[%d] = %+v, want %+v", i, got[i], want[i]) + } + } +} + +func TestAssemble_Empty(t *testing.T) { + prog, err := Assemble(nil) + if err != nil { + t.Fatalf("Assemble: %v", err) + } + if len(prog.Instructions) != 0 { + t.Errorf("expected no instructions, got %d", len(prog.Instructions)) + } +} + +func TestAssemble_MaxRegPerBank(t *testing.T) { + prog, err := Assemble([]Instr{ + LoadStr{Dest: 0, Value: "USD/2"}, + LoadInt{Dest: 1, Value: *big.NewInt(10)}, + BinaryOp{Op: OpMakePortion{}, Left: 1, Right: 1, Dest: 3}, + }) + require.NoError(t, err) + + require.Equal(t, byte(1), prog.MaxRegString, "one str reg") + require.Equal(t, byte(1), prog.MaxRegInt, "one int reg") + require.Equal(t, byte(1), prog.MaxRegPortion, "one portion reg") + require.Equal(t, byte(0), prog.MaxRegBool, "no bool reg") +} + +// The bool value is in the opcode, so the two constants differ only there, +// and bool registers are indexed in their own bank. Neither constant is +// read again after being set, so the allocator frees the first one's slot +// right after it and reuses it for the second — they end up sharing bank +// index 0. +func TestAssemble_ConstBool(t *testing.T) { + prog, err := Assemble([]Instr{ + LoadInt{Dest: 0, Value: *big.NewInt(1)}, + ConstBool{Dest: 1, Value: true}, + ConstBool{Dest: 2, Value: false}, + }) + require.NoError(t, err) + + require.Equal(t, []vm.Instruction{ + vm.NewBC(vm.Op_LoadInt, 0, 0), + {Opcode: byte(vm.Op_ConstTrue), A: 0, B: 0xFF, C: 0xFF}, + {Opcode: byte(vm.Op_ConstFalse), A: 0, B: 0xFF, C: 0xFF}, + }, prog.Instructions) + + require.Equal(t, byte(1), prog.MaxRegInt) + require.Equal(t, byte(1), prog.MaxRegBool) +} + +// 255 registers per bank, because 0xFF is the "operand unset" sentinel. The +// allocator reuses a slot once its register's last reference has passed, so +// the cap is only reachable when that many registers are genuinely live at +// once — not merely ever defined (see TestAssemble_SingleUseRegistersDontAccumulate). +func TestAssemble_RegisterBankOverflow(t *testing.T) { + // liveInts defines n distinct int registers, then reads every one of + // them back (via CheckEnoughFunds, which takes two int operands and + // produces no destination register, so the read-back itself never adds + // extra pressure), so all n stay simultaneously live right up until the + // last definition runs — the peak liveness a real allocator has to + // accommodate. + liveInts := func(n int) []Instr { + instrs := make([]Instr, 0, 2*n) + for i := range n { + instrs = append(instrs, LoadInt{Dest: Reg(i), Value: *big.NewInt(int64(i))}) + } + for i := range n { + instrs = append(instrs, CheckEnoughFunds{Got: Reg(i), Needed: Reg(i)}) + } + return instrs + } + + t.Run("255 simultaneously live registers fit", func(t *testing.T) { + prog, err := Assemble(liveInts(255)) + require.NoError(t, err) + require.Equal(t, byte(255), prog.MaxRegInt) + }) + + t.Run("256 simultaneously live registers do not", func(t *testing.T) { + _, err := Assemble(liveInts(256)) + require.ErrorContains(t, err, "register bank overflow") + }) + + t.Run("banks are counted separately", func(t *testing.T) { + instrs := liveInts(255) + for i := range 255 { + instrs = append(instrs, LoadStr{Dest: Reg(2000 + i), Value: "x"}) + } + for i := range 255 { + instrs = append(instrs, AssertSameAsset{Left: Reg(2000 + i), Right: Reg(2000 + i)}) + } + _, err := Assemble(instrs) + require.NoError(t, err) + }) +} + +// Unlike liveInts above, each register here is used exactly once, at +// definition, and never read again — so however many are defined, only one +// physical slot is ever needed at a time. +func TestAssemble_SingleUseRegistersDontAccumulate(t *testing.T) { + instrs := make([]Instr, 300) + for i := range instrs { + instrs[i] = LoadInt{Dest: Reg(i), Value: *big.NewInt(int64(i))} + } + prog, err := Assemble(instrs) + require.NoError(t, err) + require.Equal(t, byte(1), prog.MaxRegInt) +} + +func TestAssemble_JmpDelta(t *testing.T) { + t.Run("delta counts the instructions skipped", func(t *testing.T) { + prog, err := Assemble([]Instr{ + ConstBool{Dest: 0, Value: true}, // 0 + JmpIfFalse{Cond: 0, Target: "end"}, // 1 + LoadInt{Dest: 1, Value: *big.NewInt(1)}, // 2 + LoadInt{Dest: 2, Value: *big.NewInt(2)}, // 3 + LabelMarker{Label: "end"}, // -> 4 + }) + require.NoError(t, err) + + require.Equal(t, vm.NewBC(vm.Op_JmpIfFalse, 0, 2), prog.Instructions[1]) + }) + + // the two conditional jumps differ only in the opcode + t.Run("jmp_if_true emits its own opcode", func(t *testing.T) { + prog, err := Assemble([]Instr{ + ConstBool{Dest: 0, Value: true}, + JmpIfTrue{Cond: 0, Target: "end"}, + LoadInt{Dest: 1, Value: *big.NewInt(1)}, + LabelMarker{Label: "end"}, + }) + require.NoError(t, err) + + require.Equal(t, vm.NewBC(vm.Op_JmpIfTrue, 0, 1), prog.Instructions[1]) + }) + + t.Run("jump to the immediately following instruction has delta 0", func(t *testing.T) { + prog, err := Assemble([]Instr{ + ConstBool{Dest: 0, Value: true}, + JmpIfFalse{Cond: 0, Target: "end"}, + LabelMarker{Label: "end"}, + LoadInt{Dest: 1, Value: *big.NewInt(1)}, + }) + require.NoError(t, err) + + require.Equal(t, vm.NewBC(vm.Op_JmpIfFalse, 0, 0), prog.Instructions[1]) + }) + + t.Run("backward jump is rejected", func(t *testing.T) { + _, err := Assemble([]Instr{ + LabelMarker{Label: "start"}, + ConstBool{Dest: 0, Value: true}, + JmpIfFalse{Cond: 0, Target: "start"}, + }) + require.ErrorContains(t, err, "backward jump") + }) + + t.Run("duplicate label is rejected", func(t *testing.T) { + _, err := Assemble([]Instr{ + ConstBool{Dest: 0, Value: true}, + JmpIfFalse{Cond: 0, Target: "end"}, + LabelMarker{Label: "end"}, + LoadInt{Dest: 1, Value: *big.NewInt(1)}, + LabelMarker{Label: "end"}, + }) + require.ErrorContains(t, err, "duplicate label") + }) + + t.Run("jump to itself is rejected", func(t *testing.T) { + _, err := Assemble([]Instr{ + ConstBool{Dest: 0, Value: true}, + LabelMarker{Label: "self"}, + JmpIfTrue{Cond: 0, Target: "self"}, + }) + require.ErrorContains(t, err, "backward jump") + }) + + t.Run("unconditional jmp patches its delta", func(t *testing.T) { + prog, err := Assemble([]Instr{ + Jmp{Target: "end"}, + LoadInt{Dest: 0, Value: *big.NewInt(0)}, + LabelMarker{Label: "end"}, + }) + require.NoError(t, err) + + // one instruction (the load) sits between the jump and the label + require.Equal(t, vm.NewBC(vm.Op_Jmp, 0, 1), prog.Instructions[0]) + }) + + t.Run("backward unconditional jmp is rejected", func(t *testing.T) { + _, err := Assemble([]Instr{ + LabelMarker{Label: "start"}, + Jmp{Target: "start"}, + }) + require.ErrorContains(t, err, "backward jump") + }) +} diff --git a/internal/ir/builder.go b/internal/ir/builder.go new file mode 100644 index 00000000..0955d69b --- /dev/null +++ b/internal/ir/builder.go @@ -0,0 +1,41 @@ +package ir + +import "fmt" + +// Builder accumulates an instruction stream, handing out the registers and +// labels it needs. +type Builder struct { + instrs []Instr + nextReg Reg + nextLabelID int +} + +func (b *Builder) FreshReg() Reg { + r := b.nextReg + b.nextReg++ + return r +} + +// FreshLabel suffixes prefix with a counter: "inorder_end" -> #inorder_end_0. +func (b *Builder) FreshLabel(prefix string) Label { + l := Label(fmt.Sprintf("%s_%d", prefix, b.nextLabelID)) + b.nextLabelID++ + return l +} + +func (b *Builder) Push(instr Instr) { + b.instrs = append(b.instrs, instr) +} + +// PushWithDest allocates the register the instruction writes to. Allocating here +// rather than up front keeps registers numbered in emission order, which is what +// makes a Dump of the result parse back to the same program. +func (b *Builder) PushWithDest(getInstr func(dest Reg) Instr) Reg { + dest := b.FreshReg() + b.Push(getInstr(dest)) + return dest +} + +func (b *Builder) Instrs() []Instr { + return b.instrs +} diff --git a/internal/ir/dump.go b/internal/ir/dump.go new file mode 100644 index 00000000..26316df7 --- /dev/null +++ b/internal/ir/dump.go @@ -0,0 +1,252 @@ +package ir + +import ( + "fmt" + "strings" +) + +func (r Reg) String() string { return fmt.Sprintf("$r%d", uint(r)) } +func (l Label) String() string { return fmt.Sprintf("#%s", string(l)) } + +func (OpAddInt) String() string { return "add_int" } +func (OpSubInt) String() string { return "sub_int" } +func (OpAddString) String() string { return "add_string" } +func (OpStrEq) String() string { return "str_eq" } +func (OpLtInt) String() string { return "lt_int" } +func (OpEqInt) String() string { return "eq_int" } +func (OpLtPortion) String() string { return "lt_portion" } +func (OpEqPortion) String() string { return "eq_portion" } +func (OpAddPortion) String() string { return "add_portion" } +func (OpSubPortion) String() string { return "sub_portion" } +func (OpMulPortion) String() string { return "mul_portion" } +func (OpMakePortion) String() string { return "mk_portion" } +func (OpMonetaryToString) String() string { return "monetary_to_string" } + +func (OpIntCopy) String() string { return "int_copy" } +func (OpPortionCopy) String() string { return "portion_copy" } +func (OpStrCopy) String() string { return "str_copy" } +func (OpBoolCopy) String() string { return "bool_copy" } +func (OpNegInt) String() string { return "neg_int" } +func (OpIntToString) String() string { return "int_to_string" } +func (OpIsZero) String() string { return "is_zero" } +func (OpNot) String() string { return "not" } +func (OpPortionToString) String() string { return "portion_to_string" } +func (OpIntToPortion) String() string { return "int_to_portion" } +func (OpPortionToInt) String() string { return "portion_to_int" } + +func (i PullAccount) String() string { + opts := joinOpts( + optLabel("cap", i.Cap), + optLabel("overdraft", i.Overdraft), + optLabel("color", i.Color), + optLabel("scope", i.Scope), + ) + s := fmt.Sprintf("%s = pull_account(account: %s", i.Dest, i.Account) + if opts != "" { + s += ", " + opts + } + return s + ")" +} + +func (i SendToAccount) String() string { + opts := joinOpts(optLabel("account", i.Account), optLabel("cap", i.Cap), optLabel("scope", i.Scope)) + return fmt.Sprintf("send_to_account(%s)", opts) +} + +func (i CheckEnoughFunds) String() string { + return fmt.Sprintf("check_enough_funds(%s, %s)", i.Got, i.Needed) +} + +func (i Save) String() string { + opts := joinOpts(optLabel("amount", i.Amount), optLabel("scope", i.Scope)) + s := fmt.Sprintf("save(account: %s, asset: %s", i.Account, i.Asset) + if opts != "" { + s += ", " + opts + } + return s + ")" +} + +func (i AssertLeftover) String() string { + if i.Exact { + return fmt.Sprintf("assert_leftover_exact(%s)", i.Portion) + } + return fmt.Sprintf("assert_leftover(%s)", i.Portion) +} + +func (i SetCurrentAsset) String() string { + return fmt.Sprintf("set_current_asset(%s)", i.Asset) +} + +func (i AssertSameAsset) String() string { + return fmt.Sprintf("assert_same_asset(%s, %s)", i.Left, i.Right) +} + +func (i AssertValidAccount) String() string { + return fmt.Sprintf("assert_valid_account(%s)", i.Account) +} + +func (i AssertValidColor) String() string { + return fmt.Sprintf("assert_valid_color(%s)", i.Color) +} + +func (i AssertValidScope) String() string { + return fmt.Sprintf("assert_valid_scope(%s)", i.Scope) +} + +func (i AssertUnscoped) String() string { + return fmt.Sprintf("assert_unscoped(%s, %s)", i.Scope, i.Account) +} + +func (i AssertNonNegativeBalance) String() string { + return fmt.Sprintf("assert_non_negative_balance(%s, %s)", i.Balance, i.Account) +} + +func (i AssertNonNegativeAmount) String() string { + return fmt.Sprintf("assert_non_negative_amount(%s)", i.Amount) +} + +func (i AssertNonNegativePortion) String() string { + return fmt.Sprintf("assert_non_negative_portion(%s)", i.Portion) +} + +func (i SetTxMeta) String() string { + return fmt.Sprintf("set_tx_meta(%s, %s)", i.Key, i.Value) +} + +func (i SetAccountMeta) String() string { + s := fmt.Sprintf("set_account_meta(%s, %s, %s)", i.Account, i.Key, i.Value) + if i.Scope != nil { + s = fmt.Sprintf("set_account_meta(%s, %s, %s, %s)", i.Account, i.Key, i.Value, optLabel("scope", i.Scope)) + } + return s +} + +func (i MetaVar) String() string { + s := fmt.Sprintf("%s = meta<%s>(%s, %s)", i.Dest, i.Typ, i.Account, i.Key) + if i.Scope != nil { + s = fmt.Sprintf("%s = meta<%s>(%s, %s, %s)", i.Dest, i.Typ, i.Account, i.Key, optLabel("scope", i.Scope)) + } + return s +} + +func (i MetaMonetary) String() string { + s := fmt.Sprintf("[%s, %s] = meta_monetary(%s, %s)", i.DestAsset, i.DestAmount, i.Account, i.Key) + if i.Scope != nil { + s = fmt.Sprintf("[%s, %s] = meta_monetary(%s, %s, %s)", i.DestAsset, i.DestAmount, i.Account, i.Key, optLabel("scope", i.Scope)) + } + return s +} + +func (MetaStr) String() string { return "str" } +func (MetaInt) String() string { return "int" } +func (MetaPortion) String() string { return "portion" } + +func (i FetchBalance) String() string { + s := fmt.Sprintf("%s = balance(%s, %s)", i.Dest, i.Account, i.Asset) + if i.Scope != nil { + s = fmt.Sprintf("%s = balance(%s, %s, %s)", i.Dest, i.Account, i.Asset, optLabel("scope", i.Scope)) + } + return s +} + +func (i LoadVar) String() string { + return fmt.Sprintf("%s = load_var<%s>(%d)", i.Dest, i.Typ, i.Index) +} + +func (VarInt) String() string { return "int" } +func (VarStr) String() string { return "str" } + +func (i JmpIfFalse) String() string { + return fmt.Sprintf("jmp_if_false(%s, %s)", i.Cond, i.Target) +} + +func (i JmpIfTrue) String() string { + return fmt.Sprintf("jmp_if_true(%s, %s)", i.Cond, i.Target) +} + +func (i Jmp) String() string { + return fmt.Sprintf("jmp(%s)", i.Target) +} + +func (i LoadInt) String() string { + return fmt.Sprintf("%s = %s", i.Dest, &i.Value) +} + +func (i LoadStr) String() string { + return fmt.Sprintf("%s = %q", i.Dest, i.Value) +} + +func (i ConstBool) String() string { + return fmt.Sprintf("%s = %t", i.Dest, i.Value) +} + +// infixAlias returns the infix spelling of an op, for the two that have one. +func infixAlias(k BinKind) (string, bool) { + switch k.(type) { + case OpAddInt: + return "+", true + case OpSubInt: + return "-", true + default: + return "", false + } +} + +func (i BinaryOp) String() string { + alias, hasAlias := infixAlias(i.Op) + switch { + case !hasAlias: + return fmt.Sprintf("%s = %s(%s, %s)", i.Dest, i.Op, i.Left, i.Right) + case i.Dest == i.Left: + // e.g. $acc += $Reg + return fmt.Sprintf("%s %s= %s", i.Dest, alias, i.Right) + default: + // e.g. $tot = $l + $r + return fmt.Sprintf("%s = %s %s %s", i.Dest, i.Left, alias, i.Right) + } +} + +func (i UnaryOp) String() string { + return fmt.Sprintf("%s = %s(%s)", i.Dest, i.Op, i.Arg) +} + +func (i LabelMarker) String() string { return i.Label.String() } + +func (i MarkPush) String() string { return "mark_push()" } + +func (i MarkEnd) String() string { + if i.Rewind { + return "mark_rewind()" + } + return "mark_commit()" +} + +// Dump renders a program: labels flush-left, instructions indented. +func Dump(code []Instr) string { + var b strings.Builder + for _, in := range code { + if _, ok := in.(LabelMarker); ok { + fmt.Fprintf(&b, "%s\n", in) + } else { + fmt.Fprintf(&b, " %s\n", in) + } + } + return b.String() +} + +func optLabel(name string, r *Reg) string { + if r == nil { + return "" + } + return fmt.Sprintf("%s: %s", name, *r) +} + +func joinOpts(parts ...string) string { + kept := make([]string, 0, len(parts)) + for _, p := range parts { + if p != "" { + kept = append(kept, p) + } + } + return strings.Join(kept, ", ") +} diff --git a/internal/ir/instr.go b/internal/ir/instr.go new file mode 100644 index 00000000..fd87f464 --- /dev/null +++ b/internal/ir/instr.go @@ -0,0 +1,326 @@ +// Package ir is the compiler's intermediate representation, and everything that +// operates on it: Parse and Dump convert to and from the textual format (see +// ir-textual-format.md), Typecheck checks register types, Assemble lowers to a +// vm.Program. The grammar's AST is internal, so callers only see instructions. +package ir + +import ( + "fmt" + "math/big" +) + +type Reg uint + +type Label string + +type BinKind interface { + fmt.Stringer + sig() binaryOpSig +} + +type ( + OpAddInt struct{} + OpSubInt struct{} + OpAddString struct{} + // The comparisons: only `<` and `==`, per type. `>`, `<=`, `>=` and `!=` are + // front-end normalisations over these plus OpNot — see the table in + // internal/vm/instruction.go. + OpLtInt struct{} + OpEqInt struct{} + OpLtPortion struct{} + OpEqPortion struct{} + // OpStrEq yields a bool: it is the one comparison that produces a value + // rather than trapping, and what the jumps branch on. + OpStrEq struct{} + OpAddPortion struct{} + OpSubPortion struct{} + OpMulPortion struct{} + OpMakePortion struct{} + // OpMonetaryToString takes the asset (str) and the amount (int) of a monetary + // and produces its "ASSET AMOUNT" form, the inverse of funds.ParseMonetary. + OpMonetaryToString struct{} +) + +type UnKind interface { + fmt.Stringer + sig() unaryOpSig +} + +type ( + // One copy per register bank. A monetary has none: it is a (str, int) pair, so + // copying one is a str_copy plus an int_copy. + OpIntCopy struct{} + OpPortionCopy struct{} + OpStrCopy struct{} + OpBoolCopy struct{} + + OpNegInt struct{} + OpIntToString struct{} + // OpIsZero projects an int onto a bool, which is how a quantity reaches a + // jump: the jumps take a bool, so the projection has to be explicit. + OpIsZero struct{} + // OpNot is the only bool -> bool operation. + OpNot struct{} + OpPortionToString struct{} + // The two directions across the int/portion boundary. OpIntToPortion is + // exact; OpPortionToInt floors (big.Rat's denominator is always positive, so + // big.Int.Div is the floor). Together with OpMulPortion they are what an + // allotment share is made of. + OpIntToPortion struct{} + OpPortionToInt struct{} +) + +type VarType interface { + fmt.Stringer + assembleLoad(a *assembler, Dest Reg, Index uint16) error +} + +type ( + VarInt struct{} + VarStr struct{} +) + +type MetaType interface { + fmt.Stringer + assembleMeta(a *assembler, Dest, Account, Key Reg, Scope *Reg) error +} + +type ( + MetaStr struct{} + MetaInt struct{} + MetaPortion struct{} +) + +type ( + PullAccount struct { + Dest Reg // int: amount pulled + Account Reg // str + Cap, Overdraft, Color, Scope *Reg // int, int, str, str + } + SendToAccount struct { + Account, Cap, Scope *Reg // str, int, str + } + Save struct { + Account Reg // str + Asset Reg // str + Amount *Reg // int; nil = save all + Scope *Reg // str + } + CheckEnoughFunds struct{ Got, Needed Reg } // int + AssertLeftover struct { + Portion Reg // the allotment leftover (1 - sum of the given Portions) + Exact bool // no `remaining` clause: leftover must be exactly 0, else >= 0 + } + SetCurrentAsset struct{ Asset Reg } // str + AssertSameAsset struct{ Left, Right Reg } // str, str + AssertValidAccount struct{ Account Reg } // str + AssertValidColor struct{ Color Reg } // str + AssertValidScope struct{ Scope Reg } // str + AssertUnscoped struct{ Scope, Account Reg } // str, str + AssertNonNegativeBalance struct{ Balance, Account Reg } // int (the amount), str + AssertNonNegativeAmount struct{ Amount Reg } // int (a sent/saved amount, not tied to an account) + AssertNonNegativePortion struct{ Portion Reg } // portion (an allotment clause portion) + SetTxMeta struct{ Key, Value Reg } // str, str + SetAccountMeta struct { + Account, Key, Value Reg // str, str, str + Scope *Reg // str + } + MetaVar struct { + Dest Reg + Account, Key Reg // str, str + Scope *Reg // str + Typ MetaType + } + // MetaMonetary is meta: one store read yields both halves, so it is + // the only two-destination read and is not a MetaType. + MetaMonetary struct { + DestAsset Reg // str + DestAmount Reg // int + Account Reg // str + Key Reg // str + Scope *Reg + } + FetchBalance struct { + Dest Reg // int (the amount; the asset is the Asset operand) + Account, Asset Reg // str, str + Scope *Reg + } // reads the run-state (impure) + LoadVar struct { + Dest Reg + Typ VarType + Index uint16 + } + // The two conditional jumps differ only in which edge of the bool jumps, so + // either branch of a condition is one instruction and no negation is needed. + JmpIfFalse struct { + Cond Reg // bool + Target Label + } + JmpIfTrue struct { + Cond Reg // bool + Target Label + } + Jmp struct { + Target Label + } + LoadInt struct { + Dest Reg + Value big.Int + } + LoadStr struct { + Dest Reg + Value string + } + // ConstBool assembles to Op_ConstTrue or Op_ConstFalse: the value is in the + // opcode, so there is no pool entry. + ConstBool struct { + Dest Reg + Value bool + } + BinaryOp struct { + Op BinKind + Dest, Left, Right Reg + } + UnaryOp struct { + Op UnKind + Dest, Arg Reg + } + LabelMarker struct{ Label Label } + + // The mark ops, used for oneof backtracking. Neither takes a register: the mark + // is a source-queue depth on a LIFO the run-state owns, so no operand can name a + // depth it never marked. + // + // There is no "rewind but keep the mark" instruction: a retry is + // MarkEnd{Rewind: true} followed by a fresh MarkPush, which after the rollback + // yields a mark identical to the closed one. Mark depth is therefore a function + // of position in the instruction stream, so a verifier could decide statically + // that pushes and ends balance on every path, that no MarkEnd runs at depth 0, + // and that no SendToAccount, SetCurrentAsset or Save sits at depth > 0. No such + // pass exists yet — the VM enforces it at execution time — but keep emission + // verifiable: never emit a mark op on only one side of a branch. + MarkPush struct{} // opens a region at the current source-queue depth + // MarkEnd closes the innermost region. Rewind undoes what it did; otherwise the + // region's pulls and postings are committed. Dumps as mark_rewind / mark_commit. + MarkEnd struct{ Rewind bool } +) + +type Instr interface { + dests() []Reg // registers written + sources() []Reg // registers read + assemble(a *assembler) error +} + +func (i PullAccount) dests() []Reg { return []Reg{i.Dest} } +func (i PullAccount) sources() []Reg { + return present(&i.Account, i.Cap, i.Overdraft, i.Color, i.Scope) +} + +func (i SendToAccount) dests() []Reg { return nil } +func (i SendToAccount) sources() []Reg { return present(i.Account, i.Cap, i.Scope) } + +func (i CheckEnoughFunds) dests() []Reg { return nil } +func (i CheckEnoughFunds) sources() []Reg { return []Reg{i.Got, i.Needed} } + +func (i Save) dests() []Reg { return nil } +func (i Save) sources() []Reg { + regs := []Reg{i.Account, i.Asset} + if i.Amount != nil { + regs = append(regs, *i.Amount) + } + if i.Scope != nil { + regs = append(regs, *i.Scope) + } + return regs +} + +func (i AssertLeftover) dests() []Reg { return nil } +func (i AssertLeftover) sources() []Reg { return []Reg{i.Portion} } + +func (i SetCurrentAsset) dests() []Reg { return nil } +func (i SetCurrentAsset) sources() []Reg { return []Reg{i.Asset} } + +func (i AssertSameAsset) dests() []Reg { return nil } +func (i AssertSameAsset) sources() []Reg { return []Reg{i.Left, i.Right} } + +func (i AssertValidAccount) dests() []Reg { return nil } +func (i AssertValidAccount) sources() []Reg { return []Reg{i.Account} } + +func (i AssertValidColor) dests() []Reg { return nil } +func (i AssertValidColor) sources() []Reg { return []Reg{i.Color} } + +func (i AssertValidScope) dests() []Reg { return nil } +func (i AssertValidScope) sources() []Reg { return []Reg{i.Scope} } + +func (i AssertUnscoped) dests() []Reg { return nil } +func (i AssertUnscoped) sources() []Reg { return []Reg{i.Scope, i.Account} } + +func (i AssertNonNegativeBalance) dests() []Reg { return nil } +func (i AssertNonNegativeBalance) sources() []Reg { return []Reg{i.Balance, i.Account} } + +func (i AssertNonNegativeAmount) dests() []Reg { return nil } +func (i AssertNonNegativeAmount) sources() []Reg { return []Reg{i.Amount} } + +func (i AssertNonNegativePortion) dests() []Reg { return nil } +func (i AssertNonNegativePortion) sources() []Reg { return []Reg{i.Portion} } + +func (i SetTxMeta) dests() []Reg { return nil } +func (i SetTxMeta) sources() []Reg { return []Reg{i.Key, i.Value} } + +func (i SetAccountMeta) dests() []Reg { return nil } +func (i SetAccountMeta) sources() []Reg { return present(&i.Account, &i.Key, &i.Value, i.Scope) } + +func (i MetaVar) dests() []Reg { return []Reg{i.Dest} } +func (i MetaVar) sources() []Reg { return present(&i.Account, &i.Key, i.Scope) } + +func (i MetaMonetary) dests() []Reg { return []Reg{i.DestAsset, i.DestAmount} } +func (i MetaMonetary) sources() []Reg { return present(&i.Account, &i.Key, i.Scope) } + +func (i FetchBalance) dests() []Reg { return []Reg{i.Dest} } +func (i FetchBalance) sources() []Reg { return present(&i.Account, &i.Asset, i.Scope) } + +func (i LoadVar) dests() []Reg { return []Reg{i.Dest} } +func (i LoadVar) sources() []Reg { return nil } + +func (i JmpIfFalse) dests() []Reg { return nil } +func (i JmpIfFalse) sources() []Reg { return []Reg{i.Cond} } + +func (i JmpIfTrue) dests() []Reg { return nil } +func (i JmpIfTrue) sources() []Reg { return []Reg{i.Cond} } + +func (i Jmp) dests() []Reg { return nil } +func (i Jmp) sources() []Reg { return nil } + +func (i LoadInt) dests() []Reg { return []Reg{i.Dest} } +func (i LoadInt) sources() []Reg { return nil } + +func (i LoadStr) dests() []Reg { return []Reg{i.Dest} } +func (i LoadStr) sources() []Reg { return nil } + +func (i ConstBool) dests() []Reg { return []Reg{i.Dest} } +func (i ConstBool) sources() []Reg { return nil } + +func (i BinaryOp) dests() []Reg { return []Reg{i.Dest} } +func (i BinaryOp) sources() []Reg { return []Reg{i.Left, i.Right} } + +func (i UnaryOp) dests() []Reg { return []Reg{i.Dest} } +func (i UnaryOp) sources() []Reg { return []Reg{i.Arg} } + +func (i LabelMarker) dests() []Reg { return nil } +func (i LabelMarker) sources() []Reg { return nil } + +func (i MarkPush) dests() []Reg { return nil } +func (i MarkPush) sources() []Reg { return nil } + +func (i MarkEnd) dests() []Reg { return nil } +func (i MarkEnd) sources() []Reg { return nil } + +func present(regs ...*Reg) []Reg { + out := make([]Reg, 0, len(regs)) + for _, r := range regs { + if r != nil { + out = append(out, *r) + } + } + return out +} diff --git a/internal/ir/internal/syntax/antlrParser/IR.interp b/internal/ir/internal/syntax/antlrParser/IR.interp new file mode 100644 index 00000000..aff84b76 --- /dev/null +++ b/internal/ir/internal/syntax/antlrParser/IR.interp @@ -0,0 +1,71 @@ +token literal names: +null +':' +null +null +null +null +null +null +null +null +null +'(' +')' +'[' +']' +',' +'=' +'+' +'-' +'+=' +'-=' +'<' +'>' +'_' + +token symbolic names: +null +null +WS +NEWLINE +TYPE_KEYWORD +BOOL +REG +LABEL +INT +STRING +IDENTIFIER +LPAREN +RPAREN +LBRACKET +RBRACKET +COMMA +EQ +PLUS +MINUS +PLUS_EQ +MINUS_EQ +LT +GT +UNDERSCORE + +rule names: +program +line +labelMarker +instruction +dest +regList +instrCall +instrName +typeName +args +arg +value +const_ +reg + + +atn: +[4, 1, 23, 129, 2, 0, 7, 0, 2, 1, 7, 1, 2, 2, 7, 2, 2, 3, 7, 3, 2, 4, 7, 4, 2, 5, 7, 5, 2, 6, 7, 6, 2, 7, 7, 7, 2, 8, 7, 8, 2, 9, 7, 9, 2, 10, 7, 10, 2, 11, 7, 11, 2, 12, 7, 12, 2, 13, 7, 13, 1, 0, 5, 0, 30, 8, 0, 10, 0, 12, 0, 33, 9, 0, 1, 0, 1, 0, 1, 1, 1, 1, 3, 1, 39, 8, 1, 1, 2, 1, 2, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 3, 3, 62, 8, 3, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 3, 4, 70, 8, 4, 1, 5, 1, 5, 1, 5, 5, 5, 75, 8, 5, 10, 5, 12, 5, 78, 9, 5, 1, 6, 1, 6, 1, 6, 1, 6, 1, 6, 1, 7, 1, 7, 1, 7, 1, 7, 1, 7, 3, 7, 90, 8, 7, 1, 8, 1, 8, 1, 9, 1, 9, 1, 9, 5, 9, 97, 8, 9, 10, 9, 12, 9, 100, 9, 9, 3, 9, 102, 8, 9, 1, 10, 1, 10, 1, 10, 1, 10, 3, 10, 108, 8, 10, 1, 11, 1, 11, 1, 11, 1, 11, 1, 11, 1, 11, 1, 11, 3, 11, 117, 8, 11, 1, 12, 1, 12, 3, 12, 121, 8, 12, 1, 12, 1, 12, 3, 12, 125, 8, 12, 1, 13, 1, 13, 1, 13, 0, 0, 14, 0, 2, 4, 6, 8, 10, 12, 14, 16, 18, 20, 22, 24, 26, 0, 2, 1, 0, 17, 18, 1, 0, 19, 20, 133, 0, 31, 1, 0, 0, 0, 2, 38, 1, 0, 0, 0, 4, 40, 1, 0, 0, 0, 6, 61, 1, 0, 0, 0, 8, 69, 1, 0, 0, 0, 10, 71, 1, 0, 0, 0, 12, 79, 1, 0, 0, 0, 14, 84, 1, 0, 0, 0, 16, 91, 1, 0, 0, 0, 18, 101, 1, 0, 0, 0, 20, 107, 1, 0, 0, 0, 22, 116, 1, 0, 0, 0, 24, 124, 1, 0, 0, 0, 26, 126, 1, 0, 0, 0, 28, 30, 3, 2, 1, 0, 29, 28, 1, 0, 0, 0, 30, 33, 1, 0, 0, 0, 31, 29, 1, 0, 0, 0, 31, 32, 1, 0, 0, 0, 32, 34, 1, 0, 0, 0, 33, 31, 1, 0, 0, 0, 34, 35, 5, 0, 0, 1, 35, 1, 1, 0, 0, 0, 36, 39, 3, 4, 2, 0, 37, 39, 3, 6, 3, 0, 38, 36, 1, 0, 0, 0, 38, 37, 1, 0, 0, 0, 39, 3, 1, 0, 0, 0, 40, 41, 5, 7, 0, 0, 41, 5, 1, 0, 0, 0, 42, 43, 3, 8, 4, 0, 43, 44, 5, 16, 0, 0, 44, 45, 3, 12, 6, 0, 45, 62, 1, 0, 0, 0, 46, 62, 3, 12, 6, 0, 47, 48, 3, 8, 4, 0, 48, 49, 5, 16, 0, 0, 49, 50, 3, 24, 12, 0, 50, 62, 1, 0, 0, 0, 51, 52, 3, 8, 4, 0, 52, 53, 5, 16, 0, 0, 53, 54, 3, 26, 13, 0, 54, 55, 7, 0, 0, 0, 55, 56, 3, 26, 13, 0, 56, 62, 1, 0, 0, 0, 57, 58, 3, 26, 13, 0, 58, 59, 7, 1, 0, 0, 59, 60, 3, 26, 13, 0, 60, 62, 1, 0, 0, 0, 61, 42, 1, 0, 0, 0, 61, 46, 1, 0, 0, 0, 61, 47, 1, 0, 0, 0, 61, 51, 1, 0, 0, 0, 61, 57, 1, 0, 0, 0, 62, 7, 1, 0, 0, 0, 63, 70, 3, 26, 13, 0, 64, 70, 5, 23, 0, 0, 65, 66, 5, 13, 0, 0, 66, 67, 3, 10, 5, 0, 67, 68, 5, 14, 0, 0, 68, 70, 1, 0, 0, 0, 69, 63, 1, 0, 0, 0, 69, 64, 1, 0, 0, 0, 69, 65, 1, 0, 0, 0, 70, 9, 1, 0, 0, 0, 71, 76, 3, 26, 13, 0, 72, 73, 5, 15, 0, 0, 73, 75, 3, 26, 13, 0, 74, 72, 1, 0, 0, 0, 75, 78, 1, 0, 0, 0, 76, 74, 1, 0, 0, 0, 76, 77, 1, 0, 0, 0, 77, 11, 1, 0, 0, 0, 78, 76, 1, 0, 0, 0, 79, 80, 3, 14, 7, 0, 80, 81, 5, 11, 0, 0, 81, 82, 3, 18, 9, 0, 82, 83, 5, 12, 0, 0, 83, 13, 1, 0, 0, 0, 84, 89, 5, 10, 0, 0, 85, 86, 5, 21, 0, 0, 86, 87, 3, 16, 8, 0, 87, 88, 5, 22, 0, 0, 88, 90, 1, 0, 0, 0, 89, 85, 1, 0, 0, 0, 89, 90, 1, 0, 0, 0, 90, 15, 1, 0, 0, 0, 91, 92, 5, 4, 0, 0, 92, 17, 1, 0, 0, 0, 93, 98, 3, 20, 10, 0, 94, 95, 5, 15, 0, 0, 95, 97, 3, 20, 10, 0, 96, 94, 1, 0, 0, 0, 97, 100, 1, 0, 0, 0, 98, 96, 1, 0, 0, 0, 98, 99, 1, 0, 0, 0, 99, 102, 1, 0, 0, 0, 100, 98, 1, 0, 0, 0, 101, 93, 1, 0, 0, 0, 101, 102, 1, 0, 0, 0, 102, 19, 1, 0, 0, 0, 103, 108, 3, 22, 11, 0, 104, 105, 5, 10, 0, 0, 105, 106, 5, 1, 0, 0, 106, 108, 3, 22, 11, 0, 107, 103, 1, 0, 0, 0, 107, 104, 1, 0, 0, 0, 108, 21, 1, 0, 0, 0, 109, 117, 3, 26, 13, 0, 110, 117, 5, 7, 0, 0, 111, 117, 5, 8, 0, 0, 112, 113, 5, 13, 0, 0, 113, 114, 3, 10, 5, 0, 114, 115, 5, 14, 0, 0, 115, 117, 1, 0, 0, 0, 116, 109, 1, 0, 0, 0, 116, 110, 1, 0, 0, 0, 116, 111, 1, 0, 0, 0, 116, 112, 1, 0, 0, 0, 117, 23, 1, 0, 0, 0, 118, 125, 5, 9, 0, 0, 119, 121, 5, 18, 0, 0, 120, 119, 1, 0, 0, 0, 120, 121, 1, 0, 0, 0, 121, 122, 1, 0, 0, 0, 122, 125, 5, 8, 0, 0, 123, 125, 5, 5, 0, 0, 124, 118, 1, 0, 0, 0, 124, 120, 1, 0, 0, 0, 124, 123, 1, 0, 0, 0, 125, 25, 1, 0, 0, 0, 126, 127, 5, 6, 0, 0, 127, 27, 1, 0, 0, 0, 12, 31, 38, 61, 69, 76, 89, 98, 101, 107, 116, 120, 124] \ No newline at end of file diff --git a/internal/ir/internal/syntax/antlrParser/IR.tokens b/internal/ir/internal/syntax/antlrParser/IR.tokens new file mode 100644 index 00000000..4da26b21 --- /dev/null +++ b/internal/ir/internal/syntax/antlrParser/IR.tokens @@ -0,0 +1,37 @@ +T__0=1 +WS=2 +NEWLINE=3 +TYPE_KEYWORD=4 +BOOL=5 +REG=6 +LABEL=7 +INT=8 +STRING=9 +IDENTIFIER=10 +LPAREN=11 +RPAREN=12 +LBRACKET=13 +RBRACKET=14 +COMMA=15 +EQ=16 +PLUS=17 +MINUS=18 +PLUS_EQ=19 +MINUS_EQ=20 +LT=21 +GT=22 +UNDERSCORE=23 +':'=1 +'('=11 +')'=12 +'['=13 +']'=14 +','=15 +'='=16 +'+'=17 +'-'=18 +'+='=19 +'-='=20 +'<'=21 +'>'=22 +'_'=23 diff --git a/internal/ir/internal/syntax/antlrParser/IRLexer.interp b/internal/ir/internal/syntax/antlrParser/IRLexer.interp new file mode 100644 index 00000000..165146b5 --- /dev/null +++ b/internal/ir/internal/syntax/antlrParser/IRLexer.interp @@ -0,0 +1,86 @@ +token literal names: +null +':' +null +null +null +null +null +null +null +null +null +'(' +')' +'[' +']' +',' +'=' +'+' +'-' +'+=' +'-=' +'<' +'>' +'_' + +token symbolic names: +null +null +WS +NEWLINE +TYPE_KEYWORD +BOOL +REG +LABEL +INT +STRING +IDENTIFIER +LPAREN +RPAREN +LBRACKET +RBRACKET +COMMA +EQ +PLUS +MINUS +PLUS_EQ +MINUS_EQ +LT +GT +UNDERSCORE + +rule names: +T__0 +WS +NEWLINE +TYPE_KEYWORD +BOOL +REG +LABEL +INT +STRING +IDENTIFIER +LPAREN +RPAREN +LBRACKET +RBRACKET +COMMA +EQ +PLUS +MINUS +PLUS_EQ +MINUS_EQ +LT +GT +UNDERSCORE + +channel names: +DEFAULT_TOKEN_CHANNEL +HIDDEN + +mode names: +DEFAULT_MODE + +atn: +[4, 0, 23, 164, 6, -1, 2, 0, 7, 0, 2, 1, 7, 1, 2, 2, 7, 2, 2, 3, 7, 3, 2, 4, 7, 4, 2, 5, 7, 5, 2, 6, 7, 6, 2, 7, 7, 7, 2, 8, 7, 8, 2, 9, 7, 9, 2, 10, 7, 10, 2, 11, 7, 11, 2, 12, 7, 12, 2, 13, 7, 13, 2, 14, 7, 14, 2, 15, 7, 15, 2, 16, 7, 16, 2, 17, 7, 17, 2, 18, 7, 18, 2, 19, 7, 19, 2, 20, 7, 20, 2, 21, 7, 21, 2, 22, 7, 22, 1, 0, 1, 0, 1, 1, 4, 1, 51, 8, 1, 11, 1, 12, 1, 52, 1, 1, 1, 1, 1, 2, 4, 2, 58, 8, 2, 11, 2, 12, 2, 59, 1, 2, 1, 2, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 3, 3, 85, 8, 3, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 3, 4, 96, 8, 4, 1, 5, 1, 5, 1, 5, 5, 5, 101, 8, 5, 10, 5, 12, 5, 104, 9, 5, 1, 6, 1, 6, 1, 6, 5, 6, 109, 8, 6, 10, 6, 12, 6, 112, 9, 6, 1, 7, 4, 7, 115, 8, 7, 11, 7, 12, 7, 116, 1, 8, 1, 8, 1, 8, 1, 8, 5, 8, 123, 8, 8, 10, 8, 12, 8, 126, 9, 8, 1, 8, 1, 8, 1, 9, 1, 9, 5, 9, 132, 8, 9, 10, 9, 12, 9, 135, 9, 9, 1, 10, 1, 10, 1, 11, 1, 11, 1, 12, 1, 12, 1, 13, 1, 13, 1, 14, 1, 14, 1, 15, 1, 15, 1, 16, 1, 16, 1, 17, 1, 17, 1, 18, 1, 18, 1, 18, 1, 19, 1, 19, 1, 19, 1, 20, 1, 20, 1, 21, 1, 21, 1, 22, 1, 22, 0, 0, 23, 1, 1, 3, 2, 5, 3, 7, 4, 9, 5, 11, 6, 13, 7, 15, 8, 17, 9, 19, 10, 21, 11, 23, 12, 25, 13, 27, 14, 29, 15, 31, 16, 33, 17, 35, 18, 37, 19, 39, 20, 41, 21, 43, 22, 45, 23, 1, 0, 8, 2, 0, 9, 9, 32, 32, 2, 0, 10, 10, 13, 13, 3, 0, 65, 90, 95, 95, 97, 122, 4, 0, 48, 57, 65, 90, 95, 95, 97, 122, 1, 0, 48, 57, 3, 0, 10, 10, 13, 13, 34, 34, 1, 0, 97, 122, 3, 0, 48, 57, 95, 95, 97, 122, 175, 0, 1, 1, 0, 0, 0, 0, 3, 1, 0, 0, 0, 0, 5, 1, 0, 0, 0, 0, 7, 1, 0, 0, 0, 0, 9, 1, 0, 0, 0, 0, 11, 1, 0, 0, 0, 0, 13, 1, 0, 0, 0, 0, 15, 1, 0, 0, 0, 0, 17, 1, 0, 0, 0, 0, 19, 1, 0, 0, 0, 0, 21, 1, 0, 0, 0, 0, 23, 1, 0, 0, 0, 0, 25, 1, 0, 0, 0, 0, 27, 1, 0, 0, 0, 0, 29, 1, 0, 0, 0, 0, 31, 1, 0, 0, 0, 0, 33, 1, 0, 0, 0, 0, 35, 1, 0, 0, 0, 0, 37, 1, 0, 0, 0, 0, 39, 1, 0, 0, 0, 0, 41, 1, 0, 0, 0, 0, 43, 1, 0, 0, 0, 0, 45, 1, 0, 0, 0, 1, 47, 1, 0, 0, 0, 3, 50, 1, 0, 0, 0, 5, 57, 1, 0, 0, 0, 7, 84, 1, 0, 0, 0, 9, 95, 1, 0, 0, 0, 11, 97, 1, 0, 0, 0, 13, 105, 1, 0, 0, 0, 15, 114, 1, 0, 0, 0, 17, 118, 1, 0, 0, 0, 19, 129, 1, 0, 0, 0, 21, 136, 1, 0, 0, 0, 23, 138, 1, 0, 0, 0, 25, 140, 1, 0, 0, 0, 27, 142, 1, 0, 0, 0, 29, 144, 1, 0, 0, 0, 31, 146, 1, 0, 0, 0, 33, 148, 1, 0, 0, 0, 35, 150, 1, 0, 0, 0, 37, 152, 1, 0, 0, 0, 39, 155, 1, 0, 0, 0, 41, 158, 1, 0, 0, 0, 43, 160, 1, 0, 0, 0, 45, 162, 1, 0, 0, 0, 47, 48, 5, 58, 0, 0, 48, 2, 1, 0, 0, 0, 49, 51, 7, 0, 0, 0, 50, 49, 1, 0, 0, 0, 51, 52, 1, 0, 0, 0, 52, 50, 1, 0, 0, 0, 52, 53, 1, 0, 0, 0, 53, 54, 1, 0, 0, 0, 54, 55, 6, 1, 0, 0, 55, 4, 1, 0, 0, 0, 56, 58, 7, 1, 0, 0, 57, 56, 1, 0, 0, 0, 58, 59, 1, 0, 0, 0, 59, 57, 1, 0, 0, 0, 59, 60, 1, 0, 0, 0, 60, 61, 1, 0, 0, 0, 61, 62, 6, 2, 0, 0, 62, 6, 1, 0, 0, 0, 63, 64, 5, 105, 0, 0, 64, 65, 5, 110, 0, 0, 65, 85, 5, 116, 0, 0, 66, 67, 5, 115, 0, 0, 67, 68, 5, 116, 0, 0, 68, 85, 5, 114, 0, 0, 69, 70, 5, 112, 0, 0, 70, 71, 5, 111, 0, 0, 71, 72, 5, 114, 0, 0, 72, 73, 5, 116, 0, 0, 73, 74, 5, 105, 0, 0, 74, 75, 5, 111, 0, 0, 75, 85, 5, 110, 0, 0, 76, 77, 5, 109, 0, 0, 77, 78, 5, 111, 0, 0, 78, 79, 5, 110, 0, 0, 79, 80, 5, 101, 0, 0, 80, 81, 5, 116, 0, 0, 81, 82, 5, 97, 0, 0, 82, 83, 5, 114, 0, 0, 83, 85, 5, 121, 0, 0, 84, 63, 1, 0, 0, 0, 84, 66, 1, 0, 0, 0, 84, 69, 1, 0, 0, 0, 84, 76, 1, 0, 0, 0, 85, 8, 1, 0, 0, 0, 86, 87, 5, 116, 0, 0, 87, 88, 5, 114, 0, 0, 88, 89, 5, 117, 0, 0, 89, 96, 5, 101, 0, 0, 90, 91, 5, 102, 0, 0, 91, 92, 5, 97, 0, 0, 92, 93, 5, 108, 0, 0, 93, 94, 5, 115, 0, 0, 94, 96, 5, 101, 0, 0, 95, 86, 1, 0, 0, 0, 95, 90, 1, 0, 0, 0, 96, 10, 1, 0, 0, 0, 97, 98, 5, 36, 0, 0, 98, 102, 7, 2, 0, 0, 99, 101, 7, 3, 0, 0, 100, 99, 1, 0, 0, 0, 101, 104, 1, 0, 0, 0, 102, 100, 1, 0, 0, 0, 102, 103, 1, 0, 0, 0, 103, 12, 1, 0, 0, 0, 104, 102, 1, 0, 0, 0, 105, 106, 5, 35, 0, 0, 106, 110, 7, 2, 0, 0, 107, 109, 7, 3, 0, 0, 108, 107, 1, 0, 0, 0, 109, 112, 1, 0, 0, 0, 110, 108, 1, 0, 0, 0, 110, 111, 1, 0, 0, 0, 111, 14, 1, 0, 0, 0, 112, 110, 1, 0, 0, 0, 113, 115, 7, 4, 0, 0, 114, 113, 1, 0, 0, 0, 115, 116, 1, 0, 0, 0, 116, 114, 1, 0, 0, 0, 116, 117, 1, 0, 0, 0, 117, 16, 1, 0, 0, 0, 118, 124, 5, 34, 0, 0, 119, 120, 5, 92, 0, 0, 120, 123, 5, 34, 0, 0, 121, 123, 8, 5, 0, 0, 122, 119, 1, 0, 0, 0, 122, 121, 1, 0, 0, 0, 123, 126, 1, 0, 0, 0, 124, 122, 1, 0, 0, 0, 124, 125, 1, 0, 0, 0, 125, 127, 1, 0, 0, 0, 126, 124, 1, 0, 0, 0, 127, 128, 5, 34, 0, 0, 128, 18, 1, 0, 0, 0, 129, 133, 7, 6, 0, 0, 130, 132, 7, 7, 0, 0, 131, 130, 1, 0, 0, 0, 132, 135, 1, 0, 0, 0, 133, 131, 1, 0, 0, 0, 133, 134, 1, 0, 0, 0, 134, 20, 1, 0, 0, 0, 135, 133, 1, 0, 0, 0, 136, 137, 5, 40, 0, 0, 137, 22, 1, 0, 0, 0, 138, 139, 5, 41, 0, 0, 139, 24, 1, 0, 0, 0, 140, 141, 5, 91, 0, 0, 141, 26, 1, 0, 0, 0, 142, 143, 5, 93, 0, 0, 143, 28, 1, 0, 0, 0, 144, 145, 5, 44, 0, 0, 145, 30, 1, 0, 0, 0, 146, 147, 5, 61, 0, 0, 147, 32, 1, 0, 0, 0, 148, 149, 5, 43, 0, 0, 149, 34, 1, 0, 0, 0, 150, 151, 5, 45, 0, 0, 151, 36, 1, 0, 0, 0, 152, 153, 5, 43, 0, 0, 153, 154, 5, 61, 0, 0, 154, 38, 1, 0, 0, 0, 155, 156, 5, 45, 0, 0, 156, 157, 5, 61, 0, 0, 157, 40, 1, 0, 0, 0, 158, 159, 5, 60, 0, 0, 159, 42, 1, 0, 0, 0, 160, 161, 5, 62, 0, 0, 161, 44, 1, 0, 0, 0, 162, 163, 5, 95, 0, 0, 163, 46, 1, 0, 0, 0, 11, 0, 52, 59, 84, 95, 102, 110, 116, 122, 124, 133, 1, 6, 0, 0] \ No newline at end of file diff --git a/internal/ir/internal/syntax/antlrParser/IRLexer.tokens b/internal/ir/internal/syntax/antlrParser/IRLexer.tokens new file mode 100644 index 00000000..4da26b21 --- /dev/null +++ b/internal/ir/internal/syntax/antlrParser/IRLexer.tokens @@ -0,0 +1,37 @@ +T__0=1 +WS=2 +NEWLINE=3 +TYPE_KEYWORD=4 +BOOL=5 +REG=6 +LABEL=7 +INT=8 +STRING=9 +IDENTIFIER=10 +LPAREN=11 +RPAREN=12 +LBRACKET=13 +RBRACKET=14 +COMMA=15 +EQ=16 +PLUS=17 +MINUS=18 +PLUS_EQ=19 +MINUS_EQ=20 +LT=21 +GT=22 +UNDERSCORE=23 +':'=1 +'('=11 +')'=12 +'['=13 +']'=14 +','=15 +'='=16 +'+'=17 +'-'=18 +'+='=19 +'-='=20 +'<'=21 +'>'=22 +'_'=23 diff --git a/internal/ir/internal/syntax/antlrParser/ir_base_listener.go b/internal/ir/internal/syntax/antlrParser/ir_base_listener.go new file mode 100644 index 00000000..7860932b --- /dev/null +++ b/internal/ir/internal/syntax/antlrParser/ir_base_listener.go @@ -0,0 +1,177 @@ +// Code generated from IR.g4 by ANTLR 4.13.2. DO NOT EDIT. + +package antlrParser // IR +import "github.com/antlr4-go/antlr/v4" + +// BaseIRListener is a complete listener for a parse tree produced by IRParser. +type BaseIRListener struct{} + +var _ IRListener = &BaseIRListener{} + +// VisitTerminal is called when a terminal node is visited. +func (s *BaseIRListener) VisitTerminal(node antlr.TerminalNode) {} + +// VisitErrorNode is called when an error node is visited. +func (s *BaseIRListener) VisitErrorNode(node antlr.ErrorNode) {} + +// EnterEveryRule is called when any rule is entered. +func (s *BaseIRListener) EnterEveryRule(ctx antlr.ParserRuleContext) {} + +// ExitEveryRule is called when any rule is exited. +func (s *BaseIRListener) ExitEveryRule(ctx antlr.ParserRuleContext) {} + +// EnterProgram is called when production program is entered. +func (s *BaseIRListener) EnterProgram(ctx *ProgramContext) {} + +// ExitProgram is called when production program is exited. +func (s *BaseIRListener) ExitProgram(ctx *ProgramContext) {} + +// EnterLine is called when production line is entered. +func (s *BaseIRListener) EnterLine(ctx *LineContext) {} + +// ExitLine is called when production line is exited. +func (s *BaseIRListener) ExitLine(ctx *LineContext) {} + +// EnterLabelMarker is called when production labelMarker is entered. +func (s *BaseIRListener) EnterLabelMarker(ctx *LabelMarkerContext) {} + +// ExitLabelMarker is called when production labelMarker is exited. +func (s *BaseIRListener) ExitLabelMarker(ctx *LabelMarkerContext) {} + +// EnterInstrWithDest is called when production instrWithDest is entered. +func (s *BaseIRListener) EnterInstrWithDest(ctx *InstrWithDestContext) {} + +// ExitInstrWithDest is called when production instrWithDest is exited. +func (s *BaseIRListener) ExitInstrWithDest(ctx *InstrWithDestContext) {} + +// EnterInstrNoDest is called when production instrNoDest is entered. +func (s *BaseIRListener) EnterInstrNoDest(ctx *InstrNoDestContext) {} + +// ExitInstrNoDest is called when production instrNoDest is exited. +func (s *BaseIRListener) ExitInstrNoDest(ctx *InstrNoDestContext) {} + +// EnterConstAssign is called when production constAssign is entered. +func (s *BaseIRListener) EnterConstAssign(ctx *ConstAssignContext) {} + +// ExitConstAssign is called when production constAssign is exited. +func (s *BaseIRListener) ExitConstAssign(ctx *ConstAssignContext) {} + +// EnterInfixInstr is called when production infixInstr is entered. +func (s *BaseIRListener) EnterInfixInstr(ctx *InfixInstrContext) {} + +// ExitInfixInstr is called when production infixInstr is exited. +func (s *BaseIRListener) ExitInfixInstr(ctx *InfixInstrContext) {} + +// EnterCompoundAssignInstr is called when production compoundAssignInstr is entered. +func (s *BaseIRListener) EnterCompoundAssignInstr(ctx *CompoundAssignInstrContext) {} + +// ExitCompoundAssignInstr is called when production compoundAssignInstr is exited. +func (s *BaseIRListener) ExitCompoundAssignInstr(ctx *CompoundAssignInstrContext) {} + +// EnterDestReg is called when production destReg is entered. +func (s *BaseIRListener) EnterDestReg(ctx *DestRegContext) {} + +// ExitDestReg is called when production destReg is exited. +func (s *BaseIRListener) ExitDestReg(ctx *DestRegContext) {} + +// EnterDestDiscard is called when production destDiscard is entered. +func (s *BaseIRListener) EnterDestDiscard(ctx *DestDiscardContext) {} + +// ExitDestDiscard is called when production destDiscard is exited. +func (s *BaseIRListener) ExitDestDiscard(ctx *DestDiscardContext) {} + +// EnterDestList is called when production destList is entered. +func (s *BaseIRListener) EnterDestList(ctx *DestListContext) {} + +// ExitDestList is called when production destList is exited. +func (s *BaseIRListener) ExitDestList(ctx *DestListContext) {} + +// EnterRegList is called when production regList is entered. +func (s *BaseIRListener) EnterRegList(ctx *RegListContext) {} + +// ExitRegList is called when production regList is exited. +func (s *BaseIRListener) ExitRegList(ctx *RegListContext) {} + +// EnterInstrCall is called when production instrCall is entered. +func (s *BaseIRListener) EnterInstrCall(ctx *InstrCallContext) {} + +// ExitInstrCall is called when production instrCall is exited. +func (s *BaseIRListener) ExitInstrCall(ctx *InstrCallContext) {} + +// EnterInstrName is called when production instrName is entered. +func (s *BaseIRListener) EnterInstrName(ctx *InstrNameContext) {} + +// ExitInstrName is called when production instrName is exited. +func (s *BaseIRListener) ExitInstrName(ctx *InstrNameContext) {} + +// EnterTypeName is called when production typeName is entered. +func (s *BaseIRListener) EnterTypeName(ctx *TypeNameContext) {} + +// ExitTypeName is called when production typeName is exited. +func (s *BaseIRListener) ExitTypeName(ctx *TypeNameContext) {} + +// EnterArgs is called when production args is entered. +func (s *BaseIRListener) EnterArgs(ctx *ArgsContext) {} + +// ExitArgs is called when production args is exited. +func (s *BaseIRListener) ExitArgs(ctx *ArgsContext) {} + +// EnterPositionalArg is called when production positionalArg is entered. +func (s *BaseIRListener) EnterPositionalArg(ctx *PositionalArgContext) {} + +// ExitPositionalArg is called when production positionalArg is exited. +func (s *BaseIRListener) ExitPositionalArg(ctx *PositionalArgContext) {} + +// EnterLabeledArg is called when production labeledArg is entered. +func (s *BaseIRListener) EnterLabeledArg(ctx *LabeledArgContext) {} + +// ExitLabeledArg is called when production labeledArg is exited. +func (s *BaseIRListener) ExitLabeledArg(ctx *LabeledArgContext) {} + +// EnterValReg is called when production valReg is entered. +func (s *BaseIRListener) EnterValReg(ctx *ValRegContext) {} + +// ExitValReg is called when production valReg is exited. +func (s *BaseIRListener) ExitValReg(ctx *ValRegContext) {} + +// EnterValLabel is called when production valLabel is entered. +func (s *BaseIRListener) EnterValLabel(ctx *ValLabelContext) {} + +// ExitValLabel is called when production valLabel is exited. +func (s *BaseIRListener) ExitValLabel(ctx *ValLabelContext) {} + +// EnterValInt is called when production valInt is entered. +func (s *BaseIRListener) EnterValInt(ctx *ValIntContext) {} + +// ExitValInt is called when production valInt is exited. +func (s *BaseIRListener) ExitValInt(ctx *ValIntContext) {} + +// EnterValRegList is called when production valRegList is entered. +func (s *BaseIRListener) EnterValRegList(ctx *ValRegListContext) {} + +// ExitValRegList is called when production valRegList is exited. +func (s *BaseIRListener) ExitValRegList(ctx *ValRegListContext) {} + +// EnterConstString is called when production constString is entered. +func (s *BaseIRListener) EnterConstString(ctx *ConstStringContext) {} + +// ExitConstString is called when production constString is exited. +func (s *BaseIRListener) ExitConstString(ctx *ConstStringContext) {} + +// EnterConstInt is called when production constInt is entered. +func (s *BaseIRListener) EnterConstInt(ctx *ConstIntContext) {} + +// ExitConstInt is called when production constInt is exited. +func (s *BaseIRListener) ExitConstInt(ctx *ConstIntContext) {} + +// EnterConstBool is called when production constBool is entered. +func (s *BaseIRListener) EnterConstBool(ctx *ConstBoolContext) {} + +// ExitConstBool is called when production constBool is exited. +func (s *BaseIRListener) ExitConstBool(ctx *ConstBoolContext) {} + +// EnterReg is called when production reg is entered. +func (s *BaseIRListener) EnterReg(ctx *RegContext) {} + +// ExitReg is called when production reg is exited. +func (s *BaseIRListener) ExitReg(ctx *RegContext) {} diff --git a/internal/ir/internal/syntax/antlrParser/ir_lexer.go b/internal/ir/internal/syntax/antlrParser/ir_lexer.go new file mode 100644 index 00000000..3a2f0f76 --- /dev/null +++ b/internal/ir/internal/syntax/antlrParser/ir_lexer.go @@ -0,0 +1,197 @@ +// Code generated from IR.g4 by ANTLR 4.13.2. DO NOT EDIT. + +package antlrParser + +import ( + "fmt" + "github.com/antlr4-go/antlr/v4" + "sync" + "unicode" +) + +// Suppress unused import error +var _ = fmt.Printf +var _ = sync.Once{} +var _ = unicode.IsLetter + +type IRLexer struct { + *antlr.BaseLexer + channelNames []string + modeNames []string + // TODO: EOF string +} + +var IRLexerLexerStaticData struct { + once sync.Once + serializedATN []int32 + ChannelNames []string + ModeNames []string + LiteralNames []string + SymbolicNames []string + RuleNames []string + PredictionContextCache *antlr.PredictionContextCache + atn *antlr.ATN + decisionToDFA []*antlr.DFA +} + +func irlexerLexerInit() { + staticData := &IRLexerLexerStaticData + staticData.ChannelNames = []string{ + "DEFAULT_TOKEN_CHANNEL", "HIDDEN", + } + staticData.ModeNames = []string{ + "DEFAULT_MODE", + } + staticData.LiteralNames = []string{ + "", "':'", "", "", "", "", "", "", "", "", "", "'('", "')'", "'['", + "']'", "','", "'='", "'+'", "'-'", "'+='", "'-='", "'<'", "'>'", "'_'", + } + staticData.SymbolicNames = []string{ + "", "", "WS", "NEWLINE", "TYPE_KEYWORD", "BOOL", "REG", "LABEL", "INT", + "STRING", "IDENTIFIER", "LPAREN", "RPAREN", "LBRACKET", "RBRACKET", + "COMMA", "EQ", "PLUS", "MINUS", "PLUS_EQ", "MINUS_EQ", "LT", "GT", "UNDERSCORE", + } + staticData.RuleNames = []string{ + "T__0", "WS", "NEWLINE", "TYPE_KEYWORD", "BOOL", "REG", "LABEL", "INT", + "STRING", "IDENTIFIER", "LPAREN", "RPAREN", "LBRACKET", "RBRACKET", + "COMMA", "EQ", "PLUS", "MINUS", "PLUS_EQ", "MINUS_EQ", "LT", "GT", "UNDERSCORE", + } + staticData.PredictionContextCache = antlr.NewPredictionContextCache() + staticData.serializedATN = []int32{ + 4, 0, 23, 164, 6, -1, 2, 0, 7, 0, 2, 1, 7, 1, 2, 2, 7, 2, 2, 3, 7, 3, 2, + 4, 7, 4, 2, 5, 7, 5, 2, 6, 7, 6, 2, 7, 7, 7, 2, 8, 7, 8, 2, 9, 7, 9, 2, + 10, 7, 10, 2, 11, 7, 11, 2, 12, 7, 12, 2, 13, 7, 13, 2, 14, 7, 14, 2, 15, + 7, 15, 2, 16, 7, 16, 2, 17, 7, 17, 2, 18, 7, 18, 2, 19, 7, 19, 2, 20, 7, + 20, 2, 21, 7, 21, 2, 22, 7, 22, 1, 0, 1, 0, 1, 1, 4, 1, 51, 8, 1, 11, 1, + 12, 1, 52, 1, 1, 1, 1, 1, 2, 4, 2, 58, 8, 2, 11, 2, 12, 2, 59, 1, 2, 1, + 2, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, + 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 3, 3, 85, 8, 3, + 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 1, 4, 3, 4, 96, 8, 4, 1, + 5, 1, 5, 1, 5, 5, 5, 101, 8, 5, 10, 5, 12, 5, 104, 9, 5, 1, 6, 1, 6, 1, + 6, 5, 6, 109, 8, 6, 10, 6, 12, 6, 112, 9, 6, 1, 7, 4, 7, 115, 8, 7, 11, + 7, 12, 7, 116, 1, 8, 1, 8, 1, 8, 1, 8, 5, 8, 123, 8, 8, 10, 8, 12, 8, 126, + 9, 8, 1, 8, 1, 8, 1, 9, 1, 9, 5, 9, 132, 8, 9, 10, 9, 12, 9, 135, 9, 9, + 1, 10, 1, 10, 1, 11, 1, 11, 1, 12, 1, 12, 1, 13, 1, 13, 1, 14, 1, 14, 1, + 15, 1, 15, 1, 16, 1, 16, 1, 17, 1, 17, 1, 18, 1, 18, 1, 18, 1, 19, 1, 19, + 1, 19, 1, 20, 1, 20, 1, 21, 1, 21, 1, 22, 1, 22, 0, 0, 23, 1, 1, 3, 2, + 5, 3, 7, 4, 9, 5, 11, 6, 13, 7, 15, 8, 17, 9, 19, 10, 21, 11, 23, 12, 25, + 13, 27, 14, 29, 15, 31, 16, 33, 17, 35, 18, 37, 19, 39, 20, 41, 21, 43, + 22, 45, 23, 1, 0, 8, 2, 0, 9, 9, 32, 32, 2, 0, 10, 10, 13, 13, 3, 0, 65, + 90, 95, 95, 97, 122, 4, 0, 48, 57, 65, 90, 95, 95, 97, 122, 1, 0, 48, 57, + 3, 0, 10, 10, 13, 13, 34, 34, 1, 0, 97, 122, 3, 0, 48, 57, 95, 95, 97, + 122, 175, 0, 1, 1, 0, 0, 0, 0, 3, 1, 0, 0, 0, 0, 5, 1, 0, 0, 0, 0, 7, 1, + 0, 0, 0, 0, 9, 1, 0, 0, 0, 0, 11, 1, 0, 0, 0, 0, 13, 1, 0, 0, 0, 0, 15, + 1, 0, 0, 0, 0, 17, 1, 0, 0, 0, 0, 19, 1, 0, 0, 0, 0, 21, 1, 0, 0, 0, 0, + 23, 1, 0, 0, 0, 0, 25, 1, 0, 0, 0, 0, 27, 1, 0, 0, 0, 0, 29, 1, 0, 0, 0, + 0, 31, 1, 0, 0, 0, 0, 33, 1, 0, 0, 0, 0, 35, 1, 0, 0, 0, 0, 37, 1, 0, 0, + 0, 0, 39, 1, 0, 0, 0, 0, 41, 1, 0, 0, 0, 0, 43, 1, 0, 0, 0, 0, 45, 1, 0, + 0, 0, 1, 47, 1, 0, 0, 0, 3, 50, 1, 0, 0, 0, 5, 57, 1, 0, 0, 0, 7, 84, 1, + 0, 0, 0, 9, 95, 1, 0, 0, 0, 11, 97, 1, 0, 0, 0, 13, 105, 1, 0, 0, 0, 15, + 114, 1, 0, 0, 0, 17, 118, 1, 0, 0, 0, 19, 129, 1, 0, 0, 0, 21, 136, 1, + 0, 0, 0, 23, 138, 1, 0, 0, 0, 25, 140, 1, 0, 0, 0, 27, 142, 1, 0, 0, 0, + 29, 144, 1, 0, 0, 0, 31, 146, 1, 0, 0, 0, 33, 148, 1, 0, 0, 0, 35, 150, + 1, 0, 0, 0, 37, 152, 1, 0, 0, 0, 39, 155, 1, 0, 0, 0, 41, 158, 1, 0, 0, + 0, 43, 160, 1, 0, 0, 0, 45, 162, 1, 0, 0, 0, 47, 48, 5, 58, 0, 0, 48, 2, + 1, 0, 0, 0, 49, 51, 7, 0, 0, 0, 50, 49, 1, 0, 0, 0, 51, 52, 1, 0, 0, 0, + 52, 50, 1, 0, 0, 0, 52, 53, 1, 0, 0, 0, 53, 54, 1, 0, 0, 0, 54, 55, 6, + 1, 0, 0, 55, 4, 1, 0, 0, 0, 56, 58, 7, 1, 0, 0, 57, 56, 1, 0, 0, 0, 58, + 59, 1, 0, 0, 0, 59, 57, 1, 0, 0, 0, 59, 60, 1, 0, 0, 0, 60, 61, 1, 0, 0, + 0, 61, 62, 6, 2, 0, 0, 62, 6, 1, 0, 0, 0, 63, 64, 5, 105, 0, 0, 64, 65, + 5, 110, 0, 0, 65, 85, 5, 116, 0, 0, 66, 67, 5, 115, 0, 0, 67, 68, 5, 116, + 0, 0, 68, 85, 5, 114, 0, 0, 69, 70, 5, 112, 0, 0, 70, 71, 5, 111, 0, 0, + 71, 72, 5, 114, 0, 0, 72, 73, 5, 116, 0, 0, 73, 74, 5, 105, 0, 0, 74, 75, + 5, 111, 0, 0, 75, 85, 5, 110, 0, 0, 76, 77, 5, 109, 0, 0, 77, 78, 5, 111, + 0, 0, 78, 79, 5, 110, 0, 0, 79, 80, 5, 101, 0, 0, 80, 81, 5, 116, 0, 0, + 81, 82, 5, 97, 0, 0, 82, 83, 5, 114, 0, 0, 83, 85, 5, 121, 0, 0, 84, 63, + 1, 0, 0, 0, 84, 66, 1, 0, 0, 0, 84, 69, 1, 0, 0, 0, 84, 76, 1, 0, 0, 0, + 85, 8, 1, 0, 0, 0, 86, 87, 5, 116, 0, 0, 87, 88, 5, 114, 0, 0, 88, 89, + 5, 117, 0, 0, 89, 96, 5, 101, 0, 0, 90, 91, 5, 102, 0, 0, 91, 92, 5, 97, + 0, 0, 92, 93, 5, 108, 0, 0, 93, 94, 5, 115, 0, 0, 94, 96, 5, 101, 0, 0, + 95, 86, 1, 0, 0, 0, 95, 90, 1, 0, 0, 0, 96, 10, 1, 0, 0, 0, 97, 98, 5, + 36, 0, 0, 98, 102, 7, 2, 0, 0, 99, 101, 7, 3, 0, 0, 100, 99, 1, 0, 0, 0, + 101, 104, 1, 0, 0, 0, 102, 100, 1, 0, 0, 0, 102, 103, 1, 0, 0, 0, 103, + 12, 1, 0, 0, 0, 104, 102, 1, 0, 0, 0, 105, 106, 5, 35, 0, 0, 106, 110, + 7, 2, 0, 0, 107, 109, 7, 3, 0, 0, 108, 107, 1, 0, 0, 0, 109, 112, 1, 0, + 0, 0, 110, 108, 1, 0, 0, 0, 110, 111, 1, 0, 0, 0, 111, 14, 1, 0, 0, 0, + 112, 110, 1, 0, 0, 0, 113, 115, 7, 4, 0, 0, 114, 113, 1, 0, 0, 0, 115, + 116, 1, 0, 0, 0, 116, 114, 1, 0, 0, 0, 116, 117, 1, 0, 0, 0, 117, 16, 1, + 0, 0, 0, 118, 124, 5, 34, 0, 0, 119, 120, 5, 92, 0, 0, 120, 123, 5, 34, + 0, 0, 121, 123, 8, 5, 0, 0, 122, 119, 1, 0, 0, 0, 122, 121, 1, 0, 0, 0, + 123, 126, 1, 0, 0, 0, 124, 122, 1, 0, 0, 0, 124, 125, 1, 0, 0, 0, 125, + 127, 1, 0, 0, 0, 126, 124, 1, 0, 0, 0, 127, 128, 5, 34, 0, 0, 128, 18, + 1, 0, 0, 0, 129, 133, 7, 6, 0, 0, 130, 132, 7, 7, 0, 0, 131, 130, 1, 0, + 0, 0, 132, 135, 1, 0, 0, 0, 133, 131, 1, 0, 0, 0, 133, 134, 1, 0, 0, 0, + 134, 20, 1, 0, 0, 0, 135, 133, 1, 0, 0, 0, 136, 137, 5, 40, 0, 0, 137, + 22, 1, 0, 0, 0, 138, 139, 5, 41, 0, 0, 139, 24, 1, 0, 0, 0, 140, 141, 5, + 91, 0, 0, 141, 26, 1, 0, 0, 0, 142, 143, 5, 93, 0, 0, 143, 28, 1, 0, 0, + 0, 144, 145, 5, 44, 0, 0, 145, 30, 1, 0, 0, 0, 146, 147, 5, 61, 0, 0, 147, + 32, 1, 0, 0, 0, 148, 149, 5, 43, 0, 0, 149, 34, 1, 0, 0, 0, 150, 151, 5, + 45, 0, 0, 151, 36, 1, 0, 0, 0, 152, 153, 5, 43, 0, 0, 153, 154, 5, 61, + 0, 0, 154, 38, 1, 0, 0, 0, 155, 156, 5, 45, 0, 0, 156, 157, 5, 61, 0, 0, + 157, 40, 1, 0, 0, 0, 158, 159, 5, 60, 0, 0, 159, 42, 1, 0, 0, 0, 160, 161, + 5, 62, 0, 0, 161, 44, 1, 0, 0, 0, 162, 163, 5, 95, 0, 0, 163, 46, 1, 0, + 0, 0, 11, 0, 52, 59, 84, 95, 102, 110, 116, 122, 124, 133, 1, 6, 0, 0, + } + deserializer := antlr.NewATNDeserializer(nil) + staticData.atn = deserializer.Deserialize(staticData.serializedATN) + atn := staticData.atn + staticData.decisionToDFA = make([]*antlr.DFA, len(atn.DecisionToState)) + decisionToDFA := staticData.decisionToDFA + for index, state := range atn.DecisionToState { + decisionToDFA[index] = antlr.NewDFA(state, index) + } +} + +// IRLexerInit initializes any static state used to implement IRLexer. By default the +// static state used to implement the lexer is lazily initialized during the first call to +// NewIRLexer(). You can call this function if you wish to initialize the static state ahead +// of time. +func IRLexerInit() { + staticData := &IRLexerLexerStaticData + staticData.once.Do(irlexerLexerInit) +} + +// NewIRLexer produces a new lexer instance for the optional input antlr.CharStream. +func NewIRLexer(input antlr.CharStream) *IRLexer { + IRLexerInit() + l := new(IRLexer) + l.BaseLexer = antlr.NewBaseLexer(input) + staticData := &IRLexerLexerStaticData + l.Interpreter = antlr.NewLexerATNSimulator(l, staticData.atn, staticData.decisionToDFA, staticData.PredictionContextCache) + l.channelNames = staticData.ChannelNames + l.modeNames = staticData.ModeNames + l.RuleNames = staticData.RuleNames + l.LiteralNames = staticData.LiteralNames + l.SymbolicNames = staticData.SymbolicNames + l.GrammarFileName = "IR.g4" + // TODO: l.EOF = antlr.TokenEOF + + return l +} + +// IRLexer tokens. +const ( + IRLexerT__0 = 1 + IRLexerWS = 2 + IRLexerNEWLINE = 3 + IRLexerTYPE_KEYWORD = 4 + IRLexerBOOL = 5 + IRLexerREG = 6 + IRLexerLABEL = 7 + IRLexerINT = 8 + IRLexerSTRING = 9 + IRLexerIDENTIFIER = 10 + IRLexerLPAREN = 11 + IRLexerRPAREN = 12 + IRLexerLBRACKET = 13 + IRLexerRBRACKET = 14 + IRLexerCOMMA = 15 + IRLexerEQ = 16 + IRLexerPLUS = 17 + IRLexerMINUS = 18 + IRLexerPLUS_EQ = 19 + IRLexerMINUS_EQ = 20 + IRLexerLT = 21 + IRLexerGT = 22 + IRLexerUNDERSCORE = 23 +) diff --git a/internal/ir/internal/syntax/antlrParser/ir_listener.go b/internal/ir/internal/syntax/antlrParser/ir_listener.go new file mode 100644 index 00000000..ce84a4f0 --- /dev/null +++ b/internal/ir/internal/syntax/antlrParser/ir_listener.go @@ -0,0 +1,165 @@ +// Code generated from IR.g4 by ANTLR 4.13.2. DO NOT EDIT. + +package antlrParser // IR +import "github.com/antlr4-go/antlr/v4" + +// IRListener is a complete listener for a parse tree produced by IRParser. +type IRListener interface { + antlr.ParseTreeListener + + // EnterProgram is called when entering the program production. + EnterProgram(c *ProgramContext) + + // EnterLine is called when entering the line production. + EnterLine(c *LineContext) + + // EnterLabelMarker is called when entering the labelMarker production. + EnterLabelMarker(c *LabelMarkerContext) + + // EnterInstrWithDest is called when entering the instrWithDest production. + EnterInstrWithDest(c *InstrWithDestContext) + + // EnterInstrNoDest is called when entering the instrNoDest production. + EnterInstrNoDest(c *InstrNoDestContext) + + // EnterConstAssign is called when entering the constAssign production. + EnterConstAssign(c *ConstAssignContext) + + // EnterInfixInstr is called when entering the infixInstr production. + EnterInfixInstr(c *InfixInstrContext) + + // EnterCompoundAssignInstr is called when entering the compoundAssignInstr production. + EnterCompoundAssignInstr(c *CompoundAssignInstrContext) + + // EnterDestReg is called when entering the destReg production. + EnterDestReg(c *DestRegContext) + + // EnterDestDiscard is called when entering the destDiscard production. + EnterDestDiscard(c *DestDiscardContext) + + // EnterDestList is called when entering the destList production. + EnterDestList(c *DestListContext) + + // EnterRegList is called when entering the regList production. + EnterRegList(c *RegListContext) + + // EnterInstrCall is called when entering the instrCall production. + EnterInstrCall(c *InstrCallContext) + + // EnterInstrName is called when entering the instrName production. + EnterInstrName(c *InstrNameContext) + + // EnterTypeName is called when entering the typeName production. + EnterTypeName(c *TypeNameContext) + + // EnterArgs is called when entering the args production. + EnterArgs(c *ArgsContext) + + // EnterPositionalArg is called when entering the positionalArg production. + EnterPositionalArg(c *PositionalArgContext) + + // EnterLabeledArg is called when entering the labeledArg production. + EnterLabeledArg(c *LabeledArgContext) + + // EnterValReg is called when entering the valReg production. + EnterValReg(c *ValRegContext) + + // EnterValLabel is called when entering the valLabel production. + EnterValLabel(c *ValLabelContext) + + // EnterValInt is called when entering the valInt production. + EnterValInt(c *ValIntContext) + + // EnterValRegList is called when entering the valRegList production. + EnterValRegList(c *ValRegListContext) + + // EnterConstString is called when entering the constString production. + EnterConstString(c *ConstStringContext) + + // EnterConstInt is called when entering the constInt production. + EnterConstInt(c *ConstIntContext) + + // EnterConstBool is called when entering the constBool production. + EnterConstBool(c *ConstBoolContext) + + // EnterReg is called when entering the reg production. + EnterReg(c *RegContext) + + // ExitProgram is called when exiting the program production. + ExitProgram(c *ProgramContext) + + // ExitLine is called when exiting the line production. + ExitLine(c *LineContext) + + // ExitLabelMarker is called when exiting the labelMarker production. + ExitLabelMarker(c *LabelMarkerContext) + + // ExitInstrWithDest is called when exiting the instrWithDest production. + ExitInstrWithDest(c *InstrWithDestContext) + + // ExitInstrNoDest is called when exiting the instrNoDest production. + ExitInstrNoDest(c *InstrNoDestContext) + + // ExitConstAssign is called when exiting the constAssign production. + ExitConstAssign(c *ConstAssignContext) + + // ExitInfixInstr is called when exiting the infixInstr production. + ExitInfixInstr(c *InfixInstrContext) + + // ExitCompoundAssignInstr is called when exiting the compoundAssignInstr production. + ExitCompoundAssignInstr(c *CompoundAssignInstrContext) + + // ExitDestReg is called when exiting the destReg production. + ExitDestReg(c *DestRegContext) + + // ExitDestDiscard is called when exiting the destDiscard production. + ExitDestDiscard(c *DestDiscardContext) + + // ExitDestList is called when exiting the destList production. + ExitDestList(c *DestListContext) + + // ExitRegList is called when exiting the regList production. + ExitRegList(c *RegListContext) + + // ExitInstrCall is called when exiting the instrCall production. + ExitInstrCall(c *InstrCallContext) + + // ExitInstrName is called when exiting the instrName production. + ExitInstrName(c *InstrNameContext) + + // ExitTypeName is called when exiting the typeName production. + ExitTypeName(c *TypeNameContext) + + // ExitArgs is called when exiting the args production. + ExitArgs(c *ArgsContext) + + // ExitPositionalArg is called when exiting the positionalArg production. + ExitPositionalArg(c *PositionalArgContext) + + // ExitLabeledArg is called when exiting the labeledArg production. + ExitLabeledArg(c *LabeledArgContext) + + // ExitValReg is called when exiting the valReg production. + ExitValReg(c *ValRegContext) + + // ExitValLabel is called when exiting the valLabel production. + ExitValLabel(c *ValLabelContext) + + // ExitValInt is called when exiting the valInt production. + ExitValInt(c *ValIntContext) + + // ExitValRegList is called when exiting the valRegList production. + ExitValRegList(c *ValRegListContext) + + // ExitConstString is called when exiting the constString production. + ExitConstString(c *ConstStringContext) + + // ExitConstInt is called when exiting the constInt production. + ExitConstInt(c *ConstIntContext) + + // ExitConstBool is called when exiting the constBool production. + ExitConstBool(c *ConstBoolContext) + + // ExitReg is called when exiting the reg production. + ExitReg(c *RegContext) +} diff --git a/internal/ir/internal/syntax/antlrParser/ir_parser.go b/internal/ir/internal/syntax/antlrParser/ir_parser.go new file mode 100644 index 00000000..07a25f27 --- /dev/null +++ b/internal/ir/internal/syntax/antlrParser/ir_parser.go @@ -0,0 +1,3049 @@ +// Code generated from IR.g4 by ANTLR 4.13.2. DO NOT EDIT. + +package antlrParser // IR +import ( + "fmt" + "strconv" + "sync" + + "github.com/antlr4-go/antlr/v4" +) + +// Suppress unused import errors +var _ = fmt.Printf +var _ = strconv.Itoa +var _ = sync.Once{} + +type IRParser struct { + *antlr.BaseParser +} + +var IRParserStaticData struct { + once sync.Once + serializedATN []int32 + LiteralNames []string + SymbolicNames []string + RuleNames []string + PredictionContextCache *antlr.PredictionContextCache + atn *antlr.ATN + decisionToDFA []*antlr.DFA +} + +func irParserInit() { + staticData := &IRParserStaticData + staticData.LiteralNames = []string{ + "", "':'", "", "", "", "", "", "", "", "", "", "'('", "')'", "'['", + "']'", "','", "'='", "'+'", "'-'", "'+='", "'-='", "'<'", "'>'", "'_'", + } + staticData.SymbolicNames = []string{ + "", "", "WS", "NEWLINE", "TYPE_KEYWORD", "BOOL", "REG", "LABEL", "INT", + "STRING", "IDENTIFIER", "LPAREN", "RPAREN", "LBRACKET", "RBRACKET", + "COMMA", "EQ", "PLUS", "MINUS", "PLUS_EQ", "MINUS_EQ", "LT", "GT", "UNDERSCORE", + } + staticData.RuleNames = []string{ + "program", "line", "labelMarker", "instruction", "dest", "regList", + "instrCall", "instrName", "typeName", "args", "arg", "value", "const_", + "reg", + } + staticData.PredictionContextCache = antlr.NewPredictionContextCache() + staticData.serializedATN = []int32{ + 4, 1, 23, 129, 2, 0, 7, 0, 2, 1, 7, 1, 2, 2, 7, 2, 2, 3, 7, 3, 2, 4, 7, + 4, 2, 5, 7, 5, 2, 6, 7, 6, 2, 7, 7, 7, 2, 8, 7, 8, 2, 9, 7, 9, 2, 10, 7, + 10, 2, 11, 7, 11, 2, 12, 7, 12, 2, 13, 7, 13, 1, 0, 5, 0, 30, 8, 0, 10, + 0, 12, 0, 33, 9, 0, 1, 0, 1, 0, 1, 1, 1, 1, 3, 1, 39, 8, 1, 1, 2, 1, 2, + 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, + 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 1, 3, 3, 3, 62, 8, 3, 1, 4, 1, 4, 1, + 4, 1, 4, 1, 4, 1, 4, 3, 4, 70, 8, 4, 1, 5, 1, 5, 1, 5, 5, 5, 75, 8, 5, + 10, 5, 12, 5, 78, 9, 5, 1, 6, 1, 6, 1, 6, 1, 6, 1, 6, 1, 7, 1, 7, 1, 7, + 1, 7, 1, 7, 3, 7, 90, 8, 7, 1, 8, 1, 8, 1, 9, 1, 9, 1, 9, 5, 9, 97, 8, + 9, 10, 9, 12, 9, 100, 9, 9, 3, 9, 102, 8, 9, 1, 10, 1, 10, 1, 10, 1, 10, + 3, 10, 108, 8, 10, 1, 11, 1, 11, 1, 11, 1, 11, 1, 11, 1, 11, 1, 11, 3, + 11, 117, 8, 11, 1, 12, 1, 12, 3, 12, 121, 8, 12, 1, 12, 1, 12, 3, 12, 125, + 8, 12, 1, 13, 1, 13, 1, 13, 0, 0, 14, 0, 2, 4, 6, 8, 10, 12, 14, 16, 18, + 20, 22, 24, 26, 0, 2, 1, 0, 17, 18, 1, 0, 19, 20, 133, 0, 31, 1, 0, 0, + 0, 2, 38, 1, 0, 0, 0, 4, 40, 1, 0, 0, 0, 6, 61, 1, 0, 0, 0, 8, 69, 1, 0, + 0, 0, 10, 71, 1, 0, 0, 0, 12, 79, 1, 0, 0, 0, 14, 84, 1, 0, 0, 0, 16, 91, + 1, 0, 0, 0, 18, 101, 1, 0, 0, 0, 20, 107, 1, 0, 0, 0, 22, 116, 1, 0, 0, + 0, 24, 124, 1, 0, 0, 0, 26, 126, 1, 0, 0, 0, 28, 30, 3, 2, 1, 0, 29, 28, + 1, 0, 0, 0, 30, 33, 1, 0, 0, 0, 31, 29, 1, 0, 0, 0, 31, 32, 1, 0, 0, 0, + 32, 34, 1, 0, 0, 0, 33, 31, 1, 0, 0, 0, 34, 35, 5, 0, 0, 1, 35, 1, 1, 0, + 0, 0, 36, 39, 3, 4, 2, 0, 37, 39, 3, 6, 3, 0, 38, 36, 1, 0, 0, 0, 38, 37, + 1, 0, 0, 0, 39, 3, 1, 0, 0, 0, 40, 41, 5, 7, 0, 0, 41, 5, 1, 0, 0, 0, 42, + 43, 3, 8, 4, 0, 43, 44, 5, 16, 0, 0, 44, 45, 3, 12, 6, 0, 45, 62, 1, 0, + 0, 0, 46, 62, 3, 12, 6, 0, 47, 48, 3, 8, 4, 0, 48, 49, 5, 16, 0, 0, 49, + 50, 3, 24, 12, 0, 50, 62, 1, 0, 0, 0, 51, 52, 3, 8, 4, 0, 52, 53, 5, 16, + 0, 0, 53, 54, 3, 26, 13, 0, 54, 55, 7, 0, 0, 0, 55, 56, 3, 26, 13, 0, 56, + 62, 1, 0, 0, 0, 57, 58, 3, 26, 13, 0, 58, 59, 7, 1, 0, 0, 59, 60, 3, 26, + 13, 0, 60, 62, 1, 0, 0, 0, 61, 42, 1, 0, 0, 0, 61, 46, 1, 0, 0, 0, 61, + 47, 1, 0, 0, 0, 61, 51, 1, 0, 0, 0, 61, 57, 1, 0, 0, 0, 62, 7, 1, 0, 0, + 0, 63, 70, 3, 26, 13, 0, 64, 70, 5, 23, 0, 0, 65, 66, 5, 13, 0, 0, 66, + 67, 3, 10, 5, 0, 67, 68, 5, 14, 0, 0, 68, 70, 1, 0, 0, 0, 69, 63, 1, 0, + 0, 0, 69, 64, 1, 0, 0, 0, 69, 65, 1, 0, 0, 0, 70, 9, 1, 0, 0, 0, 71, 76, + 3, 26, 13, 0, 72, 73, 5, 15, 0, 0, 73, 75, 3, 26, 13, 0, 74, 72, 1, 0, + 0, 0, 75, 78, 1, 0, 0, 0, 76, 74, 1, 0, 0, 0, 76, 77, 1, 0, 0, 0, 77, 11, + 1, 0, 0, 0, 78, 76, 1, 0, 0, 0, 79, 80, 3, 14, 7, 0, 80, 81, 5, 11, 0, + 0, 81, 82, 3, 18, 9, 0, 82, 83, 5, 12, 0, 0, 83, 13, 1, 0, 0, 0, 84, 89, + 5, 10, 0, 0, 85, 86, 5, 21, 0, 0, 86, 87, 3, 16, 8, 0, 87, 88, 5, 22, 0, + 0, 88, 90, 1, 0, 0, 0, 89, 85, 1, 0, 0, 0, 89, 90, 1, 0, 0, 0, 90, 15, + 1, 0, 0, 0, 91, 92, 5, 4, 0, 0, 92, 17, 1, 0, 0, 0, 93, 98, 3, 20, 10, + 0, 94, 95, 5, 15, 0, 0, 95, 97, 3, 20, 10, 0, 96, 94, 1, 0, 0, 0, 97, 100, + 1, 0, 0, 0, 98, 96, 1, 0, 0, 0, 98, 99, 1, 0, 0, 0, 99, 102, 1, 0, 0, 0, + 100, 98, 1, 0, 0, 0, 101, 93, 1, 0, 0, 0, 101, 102, 1, 0, 0, 0, 102, 19, + 1, 0, 0, 0, 103, 108, 3, 22, 11, 0, 104, 105, 5, 10, 0, 0, 105, 106, 5, + 1, 0, 0, 106, 108, 3, 22, 11, 0, 107, 103, 1, 0, 0, 0, 107, 104, 1, 0, + 0, 0, 108, 21, 1, 0, 0, 0, 109, 117, 3, 26, 13, 0, 110, 117, 5, 7, 0, 0, + 111, 117, 5, 8, 0, 0, 112, 113, 5, 13, 0, 0, 113, 114, 3, 10, 5, 0, 114, + 115, 5, 14, 0, 0, 115, 117, 1, 0, 0, 0, 116, 109, 1, 0, 0, 0, 116, 110, + 1, 0, 0, 0, 116, 111, 1, 0, 0, 0, 116, 112, 1, 0, 0, 0, 117, 23, 1, 0, + 0, 0, 118, 125, 5, 9, 0, 0, 119, 121, 5, 18, 0, 0, 120, 119, 1, 0, 0, 0, + 120, 121, 1, 0, 0, 0, 121, 122, 1, 0, 0, 0, 122, 125, 5, 8, 0, 0, 123, + 125, 5, 5, 0, 0, 124, 118, 1, 0, 0, 0, 124, 120, 1, 0, 0, 0, 124, 123, + 1, 0, 0, 0, 125, 25, 1, 0, 0, 0, 126, 127, 5, 6, 0, 0, 127, 27, 1, 0, 0, + 0, 12, 31, 38, 61, 69, 76, 89, 98, 101, 107, 116, 120, 124, + } + deserializer := antlr.NewATNDeserializer(nil) + staticData.atn = deserializer.Deserialize(staticData.serializedATN) + atn := staticData.atn + staticData.decisionToDFA = make([]*antlr.DFA, len(atn.DecisionToState)) + decisionToDFA := staticData.decisionToDFA + for index, state := range atn.DecisionToState { + decisionToDFA[index] = antlr.NewDFA(state, index) + } +} + +// IRParserInit initializes any static state used to implement IRParser. By default the +// static state used to implement the parser is lazily initialized during the first call to +// NewIRParser(). You can call this function if you wish to initialize the static state ahead +// of time. +func IRParserInit() { + staticData := &IRParserStaticData + staticData.once.Do(irParserInit) +} + +// NewIRParser produces a new parser instance for the optional input antlr.TokenStream. +func NewIRParser(input antlr.TokenStream) *IRParser { + IRParserInit() + this := new(IRParser) + this.BaseParser = antlr.NewBaseParser(input) + staticData := &IRParserStaticData + this.Interpreter = antlr.NewParserATNSimulator(this, staticData.atn, staticData.decisionToDFA, staticData.PredictionContextCache) + this.RuleNames = staticData.RuleNames + this.LiteralNames = staticData.LiteralNames + this.SymbolicNames = staticData.SymbolicNames + this.GrammarFileName = "IR.g4" + + return this +} + +// IRParser tokens. +const ( + IRParserEOF = antlr.TokenEOF + IRParserT__0 = 1 + IRParserWS = 2 + IRParserNEWLINE = 3 + IRParserTYPE_KEYWORD = 4 + IRParserBOOL = 5 + IRParserREG = 6 + IRParserLABEL = 7 + IRParserINT = 8 + IRParserSTRING = 9 + IRParserIDENTIFIER = 10 + IRParserLPAREN = 11 + IRParserRPAREN = 12 + IRParserLBRACKET = 13 + IRParserRBRACKET = 14 + IRParserCOMMA = 15 + IRParserEQ = 16 + IRParserPLUS = 17 + IRParserMINUS = 18 + IRParserPLUS_EQ = 19 + IRParserMINUS_EQ = 20 + IRParserLT = 21 + IRParserGT = 22 + IRParserUNDERSCORE = 23 +) + +// IRParser rules. +const ( + IRParserRULE_program = 0 + IRParserRULE_line = 1 + IRParserRULE_labelMarker = 2 + IRParserRULE_instruction = 3 + IRParserRULE_dest = 4 + IRParserRULE_regList = 5 + IRParserRULE_instrCall = 6 + IRParserRULE_instrName = 7 + IRParserRULE_typeName = 8 + IRParserRULE_args = 9 + IRParserRULE_arg = 10 + IRParserRULE_value = 11 + IRParserRULE_const_ = 12 + IRParserRULE_reg = 13 +) + +// IProgramContext is an interface to support dynamic dispatch. +type IProgramContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + + // Getter signatures + EOF() antlr.TerminalNode + AllLine() []ILineContext + Line(i int) ILineContext + + // IsProgramContext differentiates from other interfaces. + IsProgramContext() +} + +type ProgramContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyProgramContext() *ProgramContext { + var p = new(ProgramContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_program + return p +} + +func InitEmptyProgramContext(p *ProgramContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_program +} + +func (*ProgramContext) IsProgramContext() {} + +func NewProgramContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *ProgramContext { + var p = new(ProgramContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_program + + return p +} + +func (s *ProgramContext) GetParser() antlr.Parser { return s.parser } + +func (s *ProgramContext) EOF() antlr.TerminalNode { + return s.GetToken(IRParserEOF, 0) +} + +func (s *ProgramContext) AllLine() []ILineContext { + children := s.GetChildren() + len := 0 + for _, ctx := range children { + if _, ok := ctx.(ILineContext); ok { + len++ + } + } + + tst := make([]ILineContext, len) + i := 0 + for _, ctx := range children { + if t, ok := ctx.(ILineContext); ok { + tst[i] = t.(ILineContext) + i++ + } + } + + return tst +} + +func (s *ProgramContext) Line(i int) ILineContext { + var t antlr.RuleContext + j := 0 + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(ILineContext); ok { + if j == i { + t = ctx.(antlr.RuleContext) + break + } + j++ + } + } + + if t == nil { + return nil + } + + return t.(ILineContext) +} + +func (s *ProgramContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ProgramContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +func (s *ProgramContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterProgram(s) + } +} + +func (s *ProgramContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitProgram(s) + } +} + +func (p *IRParser) Program() (localctx IProgramContext) { + localctx = NewProgramContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 0, IRParserRULE_program) + var _la int + + p.EnterOuterAlt(localctx, 1) + p.SetState(31) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + _la = p.GetTokenStream().LA(1) + + for (int64(_la) & ^0x3f) == 0 && ((int64(1)<<_la)&8398016) != 0 { + { + p.SetState(28) + p.Line() + } + + p.SetState(33) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + _la = p.GetTokenStream().LA(1) + } + { + p.SetState(34) + p.Match(IRParserEOF) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// ILineContext is an interface to support dynamic dispatch. +type ILineContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + + // Getter signatures + LabelMarker() ILabelMarkerContext + Instruction() IInstructionContext + + // IsLineContext differentiates from other interfaces. + IsLineContext() +} + +type LineContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyLineContext() *LineContext { + var p = new(LineContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_line + return p +} + +func InitEmptyLineContext(p *LineContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_line +} + +func (*LineContext) IsLineContext() {} + +func NewLineContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *LineContext { + var p = new(LineContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_line + + return p +} + +func (s *LineContext) GetParser() antlr.Parser { return s.parser } + +func (s *LineContext) LabelMarker() ILabelMarkerContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(ILabelMarkerContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(ILabelMarkerContext) +} + +func (s *LineContext) Instruction() IInstructionContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IInstructionContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IInstructionContext) +} + +func (s *LineContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *LineContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +func (s *LineContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterLine(s) + } +} + +func (s *LineContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitLine(s) + } +} + +func (p *IRParser) Line() (localctx ILineContext) { + localctx = NewLineContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 2, IRParserRULE_line) + p.SetState(38) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + + switch p.GetTokenStream().LA(1) { + case IRParserLABEL: + p.EnterOuterAlt(localctx, 1) + { + p.SetState(36) + p.LabelMarker() + } + + case IRParserREG, IRParserIDENTIFIER, IRParserLBRACKET, IRParserUNDERSCORE: + p.EnterOuterAlt(localctx, 2) + { + p.SetState(37) + p.Instruction() + } + + default: + p.SetError(antlr.NewNoViableAltException(p, nil, nil, nil, nil, nil)) + goto errorExit + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// ILabelMarkerContext is an interface to support dynamic dispatch. +type ILabelMarkerContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + + // Getter signatures + LABEL() antlr.TerminalNode + + // IsLabelMarkerContext differentiates from other interfaces. + IsLabelMarkerContext() +} + +type LabelMarkerContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyLabelMarkerContext() *LabelMarkerContext { + var p = new(LabelMarkerContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_labelMarker + return p +} + +func InitEmptyLabelMarkerContext(p *LabelMarkerContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_labelMarker +} + +func (*LabelMarkerContext) IsLabelMarkerContext() {} + +func NewLabelMarkerContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *LabelMarkerContext { + var p = new(LabelMarkerContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_labelMarker + + return p +} + +func (s *LabelMarkerContext) GetParser() antlr.Parser { return s.parser } + +func (s *LabelMarkerContext) LABEL() antlr.TerminalNode { + return s.GetToken(IRParserLABEL, 0) +} + +func (s *LabelMarkerContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *LabelMarkerContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +func (s *LabelMarkerContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterLabelMarker(s) + } +} + +func (s *LabelMarkerContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitLabelMarker(s) + } +} + +func (p *IRParser) LabelMarker() (localctx ILabelMarkerContext) { + localctx = NewLabelMarkerContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 4, IRParserRULE_labelMarker) + p.EnterOuterAlt(localctx, 1) + { + p.SetState(40) + p.Match(IRParserLABEL) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IInstructionContext is an interface to support dynamic dispatch. +type IInstructionContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + // IsInstructionContext differentiates from other interfaces. + IsInstructionContext() +} + +type InstructionContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyInstructionContext() *InstructionContext { + var p = new(InstructionContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_instruction + return p +} + +func InitEmptyInstructionContext(p *InstructionContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_instruction +} + +func (*InstructionContext) IsInstructionContext() {} + +func NewInstructionContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *InstructionContext { + var p = new(InstructionContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_instruction + + return p +} + +func (s *InstructionContext) GetParser() antlr.Parser { return s.parser } + +func (s *InstructionContext) CopyAll(ctx *InstructionContext) { + s.CopyFrom(&ctx.BaseParserRuleContext) +} + +func (s *InstructionContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *InstructionContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +type InstrNoDestContext struct { + InstructionContext +} + +func NewInstrNoDestContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *InstrNoDestContext { + var p = new(InstrNoDestContext) + + InitEmptyInstructionContext(&p.InstructionContext) + p.parser = parser + p.CopyAll(ctx.(*InstructionContext)) + + return p +} + +func (s *InstrNoDestContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *InstrNoDestContext) InstrCall() IInstrCallContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IInstrCallContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IInstrCallContext) +} + +func (s *InstrNoDestContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterInstrNoDest(s) + } +} + +func (s *InstrNoDestContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitInstrNoDest(s) + } +} + +type CompoundAssignInstrContext struct { + InstructionContext + left IRegContext + op antlr.Token + right IRegContext +} + +func NewCompoundAssignInstrContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *CompoundAssignInstrContext { + var p = new(CompoundAssignInstrContext) + + InitEmptyInstructionContext(&p.InstructionContext) + p.parser = parser + p.CopyAll(ctx.(*InstructionContext)) + + return p +} + +func (s *CompoundAssignInstrContext) GetOp() antlr.Token { return s.op } + +func (s *CompoundAssignInstrContext) SetOp(v antlr.Token) { s.op = v } + +func (s *CompoundAssignInstrContext) GetLeft() IRegContext { return s.left } + +func (s *CompoundAssignInstrContext) GetRight() IRegContext { return s.right } + +func (s *CompoundAssignInstrContext) SetLeft(v IRegContext) { s.left = v } + +func (s *CompoundAssignInstrContext) SetRight(v IRegContext) { s.right = v } + +func (s *CompoundAssignInstrContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *CompoundAssignInstrContext) AllReg() []IRegContext { + children := s.GetChildren() + len := 0 + for _, ctx := range children { + if _, ok := ctx.(IRegContext); ok { + len++ + } + } + + tst := make([]IRegContext, len) + i := 0 + for _, ctx := range children { + if t, ok := ctx.(IRegContext); ok { + tst[i] = t.(IRegContext) + i++ + } + } + + return tst +} + +func (s *CompoundAssignInstrContext) Reg(i int) IRegContext { + var t antlr.RuleContext + j := 0 + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IRegContext); ok { + if j == i { + t = ctx.(antlr.RuleContext) + break + } + j++ + } + } + + if t == nil { + return nil + } + + return t.(IRegContext) +} + +func (s *CompoundAssignInstrContext) PLUS_EQ() antlr.TerminalNode { + return s.GetToken(IRParserPLUS_EQ, 0) +} + +func (s *CompoundAssignInstrContext) MINUS_EQ() antlr.TerminalNode { + return s.GetToken(IRParserMINUS_EQ, 0) +} + +func (s *CompoundAssignInstrContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterCompoundAssignInstr(s) + } +} + +func (s *CompoundAssignInstrContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitCompoundAssignInstr(s) + } +} + +type ConstAssignContext struct { + InstructionContext +} + +func NewConstAssignContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *ConstAssignContext { + var p = new(ConstAssignContext) + + InitEmptyInstructionContext(&p.InstructionContext) + p.parser = parser + p.CopyAll(ctx.(*InstructionContext)) + + return p +} + +func (s *ConstAssignContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ConstAssignContext) Dest() IDestContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IDestContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IDestContext) +} + +func (s *ConstAssignContext) EQ() antlr.TerminalNode { + return s.GetToken(IRParserEQ, 0) +} + +func (s *ConstAssignContext) Const_() IConst_Context { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IConst_Context); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IConst_Context) +} + +func (s *ConstAssignContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterConstAssign(s) + } +} + +func (s *ConstAssignContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitConstAssign(s) + } +} + +type InstrWithDestContext struct { + InstructionContext +} + +func NewInstrWithDestContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *InstrWithDestContext { + var p = new(InstrWithDestContext) + + InitEmptyInstructionContext(&p.InstructionContext) + p.parser = parser + p.CopyAll(ctx.(*InstructionContext)) + + return p +} + +func (s *InstrWithDestContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *InstrWithDestContext) Dest() IDestContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IDestContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IDestContext) +} + +func (s *InstrWithDestContext) EQ() antlr.TerminalNode { + return s.GetToken(IRParserEQ, 0) +} + +func (s *InstrWithDestContext) InstrCall() IInstrCallContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IInstrCallContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IInstrCallContext) +} + +func (s *InstrWithDestContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterInstrWithDest(s) + } +} + +func (s *InstrWithDestContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitInstrWithDest(s) + } +} + +type InfixInstrContext struct { + InstructionContext + left IRegContext + op antlr.Token + right IRegContext +} + +func NewInfixInstrContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *InfixInstrContext { + var p = new(InfixInstrContext) + + InitEmptyInstructionContext(&p.InstructionContext) + p.parser = parser + p.CopyAll(ctx.(*InstructionContext)) + + return p +} + +func (s *InfixInstrContext) GetOp() antlr.Token { return s.op } + +func (s *InfixInstrContext) SetOp(v antlr.Token) { s.op = v } + +func (s *InfixInstrContext) GetLeft() IRegContext { return s.left } + +func (s *InfixInstrContext) GetRight() IRegContext { return s.right } + +func (s *InfixInstrContext) SetLeft(v IRegContext) { s.left = v } + +func (s *InfixInstrContext) SetRight(v IRegContext) { s.right = v } + +func (s *InfixInstrContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *InfixInstrContext) Dest() IDestContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IDestContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IDestContext) +} + +func (s *InfixInstrContext) EQ() antlr.TerminalNode { + return s.GetToken(IRParserEQ, 0) +} + +func (s *InfixInstrContext) AllReg() []IRegContext { + children := s.GetChildren() + len := 0 + for _, ctx := range children { + if _, ok := ctx.(IRegContext); ok { + len++ + } + } + + tst := make([]IRegContext, len) + i := 0 + for _, ctx := range children { + if t, ok := ctx.(IRegContext); ok { + tst[i] = t.(IRegContext) + i++ + } + } + + return tst +} + +func (s *InfixInstrContext) Reg(i int) IRegContext { + var t antlr.RuleContext + j := 0 + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IRegContext); ok { + if j == i { + t = ctx.(antlr.RuleContext) + break + } + j++ + } + } + + if t == nil { + return nil + } + + return t.(IRegContext) +} + +func (s *InfixInstrContext) PLUS() antlr.TerminalNode { + return s.GetToken(IRParserPLUS, 0) +} + +func (s *InfixInstrContext) MINUS() antlr.TerminalNode { + return s.GetToken(IRParserMINUS, 0) +} + +func (s *InfixInstrContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterInfixInstr(s) + } +} + +func (s *InfixInstrContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitInfixInstr(s) + } +} + +func (p *IRParser) Instruction() (localctx IInstructionContext) { + localctx = NewInstructionContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 6, IRParserRULE_instruction) + var _la int + + p.SetState(61) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + + switch p.GetInterpreter().AdaptivePredict(p.BaseParser, p.GetTokenStream(), 2, p.GetParserRuleContext()) { + case 1: + localctx = NewInstrWithDestContext(p, localctx) + p.EnterOuterAlt(localctx, 1) + { + p.SetState(42) + p.Dest() + } + { + p.SetState(43) + p.Match(IRParserEQ) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(44) + p.InstrCall() + } + + case 2: + localctx = NewInstrNoDestContext(p, localctx) + p.EnterOuterAlt(localctx, 2) + { + p.SetState(46) + p.InstrCall() + } + + case 3: + localctx = NewConstAssignContext(p, localctx) + p.EnterOuterAlt(localctx, 3) + { + p.SetState(47) + p.Dest() + } + { + p.SetState(48) + p.Match(IRParserEQ) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(49) + p.Const_() + } + + case 4: + localctx = NewInfixInstrContext(p, localctx) + p.EnterOuterAlt(localctx, 4) + { + p.SetState(51) + p.Dest() + } + { + p.SetState(52) + p.Match(IRParserEQ) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(53) + + var _x = p.Reg() + + localctx.(*InfixInstrContext).left = _x + } + { + p.SetState(54) + + var _lt = p.GetTokenStream().LT(1) + + localctx.(*InfixInstrContext).op = _lt + + _la = p.GetTokenStream().LA(1) + + if !(_la == IRParserPLUS || _la == IRParserMINUS) { + var _ri = p.GetErrorHandler().RecoverInline(p) + + localctx.(*InfixInstrContext).op = _ri + } else { + p.GetErrorHandler().ReportMatch(p) + p.Consume() + } + } + { + p.SetState(55) + + var _x = p.Reg() + + localctx.(*InfixInstrContext).right = _x + } + + case 5: + localctx = NewCompoundAssignInstrContext(p, localctx) + p.EnterOuterAlt(localctx, 5) + { + p.SetState(57) + + var _x = p.Reg() + + localctx.(*CompoundAssignInstrContext).left = _x + } + { + p.SetState(58) + + var _lt = p.GetTokenStream().LT(1) + + localctx.(*CompoundAssignInstrContext).op = _lt + + _la = p.GetTokenStream().LA(1) + + if !(_la == IRParserPLUS_EQ || _la == IRParserMINUS_EQ) { + var _ri = p.GetErrorHandler().RecoverInline(p) + + localctx.(*CompoundAssignInstrContext).op = _ri + } else { + p.GetErrorHandler().ReportMatch(p) + p.Consume() + } + } + { + p.SetState(59) + + var _x = p.Reg() + + localctx.(*CompoundAssignInstrContext).right = _x + } + + case antlr.ATNInvalidAltNumber: + goto errorExit + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IDestContext is an interface to support dynamic dispatch. +type IDestContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + // IsDestContext differentiates from other interfaces. + IsDestContext() +} + +type DestContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyDestContext() *DestContext { + var p = new(DestContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_dest + return p +} + +func InitEmptyDestContext(p *DestContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_dest +} + +func (*DestContext) IsDestContext() {} + +func NewDestContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *DestContext { + var p = new(DestContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_dest + + return p +} + +func (s *DestContext) GetParser() antlr.Parser { return s.parser } + +func (s *DestContext) CopyAll(ctx *DestContext) { + s.CopyFrom(&ctx.BaseParserRuleContext) +} + +func (s *DestContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *DestContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +type DestRegContext struct { + DestContext +} + +func NewDestRegContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *DestRegContext { + var p = new(DestRegContext) + + InitEmptyDestContext(&p.DestContext) + p.parser = parser + p.CopyAll(ctx.(*DestContext)) + + return p +} + +func (s *DestRegContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *DestRegContext) Reg() IRegContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IRegContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IRegContext) +} + +func (s *DestRegContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterDestReg(s) + } +} + +func (s *DestRegContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitDestReg(s) + } +} + +type DestDiscardContext struct { + DestContext +} + +func NewDestDiscardContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *DestDiscardContext { + var p = new(DestDiscardContext) + + InitEmptyDestContext(&p.DestContext) + p.parser = parser + p.CopyAll(ctx.(*DestContext)) + + return p +} + +func (s *DestDiscardContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *DestDiscardContext) UNDERSCORE() antlr.TerminalNode { + return s.GetToken(IRParserUNDERSCORE, 0) +} + +func (s *DestDiscardContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterDestDiscard(s) + } +} + +func (s *DestDiscardContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitDestDiscard(s) + } +} + +type DestListContext struct { + DestContext +} + +func NewDestListContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *DestListContext { + var p = new(DestListContext) + + InitEmptyDestContext(&p.DestContext) + p.parser = parser + p.CopyAll(ctx.(*DestContext)) + + return p +} + +func (s *DestListContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *DestListContext) LBRACKET() antlr.TerminalNode { + return s.GetToken(IRParserLBRACKET, 0) +} + +func (s *DestListContext) RegList() IRegListContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IRegListContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IRegListContext) +} + +func (s *DestListContext) RBRACKET() antlr.TerminalNode { + return s.GetToken(IRParserRBRACKET, 0) +} + +func (s *DestListContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterDestList(s) + } +} + +func (s *DestListContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitDestList(s) + } +} + +func (p *IRParser) Dest() (localctx IDestContext) { + localctx = NewDestContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 8, IRParserRULE_dest) + p.SetState(69) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + + switch p.GetTokenStream().LA(1) { + case IRParserREG: + localctx = NewDestRegContext(p, localctx) + p.EnterOuterAlt(localctx, 1) + { + p.SetState(63) + p.Reg() + } + + case IRParserUNDERSCORE: + localctx = NewDestDiscardContext(p, localctx) + p.EnterOuterAlt(localctx, 2) + { + p.SetState(64) + p.Match(IRParserUNDERSCORE) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + case IRParserLBRACKET: + localctx = NewDestListContext(p, localctx) + p.EnterOuterAlt(localctx, 3) + { + p.SetState(65) + p.Match(IRParserLBRACKET) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(66) + p.RegList() + } + { + p.SetState(67) + p.Match(IRParserRBRACKET) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + default: + p.SetError(antlr.NewNoViableAltException(p, nil, nil, nil, nil, nil)) + goto errorExit + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IRegListContext is an interface to support dynamic dispatch. +type IRegListContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + + // Getter signatures + AllReg() []IRegContext + Reg(i int) IRegContext + AllCOMMA() []antlr.TerminalNode + COMMA(i int) antlr.TerminalNode + + // IsRegListContext differentiates from other interfaces. + IsRegListContext() +} + +type RegListContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyRegListContext() *RegListContext { + var p = new(RegListContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_regList + return p +} + +func InitEmptyRegListContext(p *RegListContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_regList +} + +func (*RegListContext) IsRegListContext() {} + +func NewRegListContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *RegListContext { + var p = new(RegListContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_regList + + return p +} + +func (s *RegListContext) GetParser() antlr.Parser { return s.parser } + +func (s *RegListContext) AllReg() []IRegContext { + children := s.GetChildren() + len := 0 + for _, ctx := range children { + if _, ok := ctx.(IRegContext); ok { + len++ + } + } + + tst := make([]IRegContext, len) + i := 0 + for _, ctx := range children { + if t, ok := ctx.(IRegContext); ok { + tst[i] = t.(IRegContext) + i++ + } + } + + return tst +} + +func (s *RegListContext) Reg(i int) IRegContext { + var t antlr.RuleContext + j := 0 + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IRegContext); ok { + if j == i { + t = ctx.(antlr.RuleContext) + break + } + j++ + } + } + + if t == nil { + return nil + } + + return t.(IRegContext) +} + +func (s *RegListContext) AllCOMMA() []antlr.TerminalNode { + return s.GetTokens(IRParserCOMMA) +} + +func (s *RegListContext) COMMA(i int) antlr.TerminalNode { + return s.GetToken(IRParserCOMMA, i) +} + +func (s *RegListContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *RegListContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +func (s *RegListContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterRegList(s) + } +} + +func (s *RegListContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitRegList(s) + } +} + +func (p *IRParser) RegList() (localctx IRegListContext) { + localctx = NewRegListContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 10, IRParserRULE_regList) + var _la int + + p.EnterOuterAlt(localctx, 1) + { + p.SetState(71) + p.Reg() + } + p.SetState(76) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + _la = p.GetTokenStream().LA(1) + + for _la == IRParserCOMMA { + { + p.SetState(72) + p.Match(IRParserCOMMA) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(73) + p.Reg() + } + + p.SetState(78) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + _la = p.GetTokenStream().LA(1) + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IInstrCallContext is an interface to support dynamic dispatch. +type IInstrCallContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + + // Getter signatures + InstrName() IInstrNameContext + LPAREN() antlr.TerminalNode + Args() IArgsContext + RPAREN() antlr.TerminalNode + + // IsInstrCallContext differentiates from other interfaces. + IsInstrCallContext() +} + +type InstrCallContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyInstrCallContext() *InstrCallContext { + var p = new(InstrCallContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_instrCall + return p +} + +func InitEmptyInstrCallContext(p *InstrCallContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_instrCall +} + +func (*InstrCallContext) IsInstrCallContext() {} + +func NewInstrCallContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *InstrCallContext { + var p = new(InstrCallContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_instrCall + + return p +} + +func (s *InstrCallContext) GetParser() antlr.Parser { return s.parser } + +func (s *InstrCallContext) InstrName() IInstrNameContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IInstrNameContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IInstrNameContext) +} + +func (s *InstrCallContext) LPAREN() antlr.TerminalNode { + return s.GetToken(IRParserLPAREN, 0) +} + +func (s *InstrCallContext) Args() IArgsContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IArgsContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IArgsContext) +} + +func (s *InstrCallContext) RPAREN() antlr.TerminalNode { + return s.GetToken(IRParserRPAREN, 0) +} + +func (s *InstrCallContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *InstrCallContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +func (s *InstrCallContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterInstrCall(s) + } +} + +func (s *InstrCallContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitInstrCall(s) + } +} + +func (p *IRParser) InstrCall() (localctx IInstrCallContext) { + localctx = NewInstrCallContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 12, IRParserRULE_instrCall) + p.EnterOuterAlt(localctx, 1) + { + p.SetState(79) + p.InstrName() + } + { + p.SetState(80) + p.Match(IRParserLPAREN) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(81) + p.Args() + } + { + p.SetState(82) + p.Match(IRParserRPAREN) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IInstrNameContext is an interface to support dynamic dispatch. +type IInstrNameContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + + // Getter signatures + IDENTIFIER() antlr.TerminalNode + LT() antlr.TerminalNode + TypeName() ITypeNameContext + GT() antlr.TerminalNode + + // IsInstrNameContext differentiates from other interfaces. + IsInstrNameContext() +} + +type InstrNameContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyInstrNameContext() *InstrNameContext { + var p = new(InstrNameContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_instrName + return p +} + +func InitEmptyInstrNameContext(p *InstrNameContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_instrName +} + +func (*InstrNameContext) IsInstrNameContext() {} + +func NewInstrNameContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *InstrNameContext { + var p = new(InstrNameContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_instrName + + return p +} + +func (s *InstrNameContext) GetParser() antlr.Parser { return s.parser } + +func (s *InstrNameContext) IDENTIFIER() antlr.TerminalNode { + return s.GetToken(IRParserIDENTIFIER, 0) +} + +func (s *InstrNameContext) LT() antlr.TerminalNode { + return s.GetToken(IRParserLT, 0) +} + +func (s *InstrNameContext) TypeName() ITypeNameContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(ITypeNameContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(ITypeNameContext) +} + +func (s *InstrNameContext) GT() antlr.TerminalNode { + return s.GetToken(IRParserGT, 0) +} + +func (s *InstrNameContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *InstrNameContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +func (s *InstrNameContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterInstrName(s) + } +} + +func (s *InstrNameContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitInstrName(s) + } +} + +func (p *IRParser) InstrName() (localctx IInstrNameContext) { + localctx = NewInstrNameContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 14, IRParserRULE_instrName) + var _la int + + p.EnterOuterAlt(localctx, 1) + { + p.SetState(84) + p.Match(IRParserIDENTIFIER) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + p.SetState(89) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + _la = p.GetTokenStream().LA(1) + + if _la == IRParserLT { + { + p.SetState(85) + p.Match(IRParserLT) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(86) + p.TypeName() + } + { + p.SetState(87) + p.Match(IRParserGT) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// ITypeNameContext is an interface to support dynamic dispatch. +type ITypeNameContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + + // Getter signatures + TYPE_KEYWORD() antlr.TerminalNode + + // IsTypeNameContext differentiates from other interfaces. + IsTypeNameContext() +} + +type TypeNameContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyTypeNameContext() *TypeNameContext { + var p = new(TypeNameContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_typeName + return p +} + +func InitEmptyTypeNameContext(p *TypeNameContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_typeName +} + +func (*TypeNameContext) IsTypeNameContext() {} + +func NewTypeNameContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *TypeNameContext { + var p = new(TypeNameContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_typeName + + return p +} + +func (s *TypeNameContext) GetParser() antlr.Parser { return s.parser } + +func (s *TypeNameContext) TYPE_KEYWORD() antlr.TerminalNode { + return s.GetToken(IRParserTYPE_KEYWORD, 0) +} + +func (s *TypeNameContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *TypeNameContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +func (s *TypeNameContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterTypeName(s) + } +} + +func (s *TypeNameContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitTypeName(s) + } +} + +func (p *IRParser) TypeName() (localctx ITypeNameContext) { + localctx = NewTypeNameContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 16, IRParserRULE_typeName) + p.EnterOuterAlt(localctx, 1) + { + p.SetState(91) + p.Match(IRParserTYPE_KEYWORD) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IArgsContext is an interface to support dynamic dispatch. +type IArgsContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + + // Getter signatures + AllArg() []IArgContext + Arg(i int) IArgContext + AllCOMMA() []antlr.TerminalNode + COMMA(i int) antlr.TerminalNode + + // IsArgsContext differentiates from other interfaces. + IsArgsContext() +} + +type ArgsContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyArgsContext() *ArgsContext { + var p = new(ArgsContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_args + return p +} + +func InitEmptyArgsContext(p *ArgsContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_args +} + +func (*ArgsContext) IsArgsContext() {} + +func NewArgsContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *ArgsContext { + var p = new(ArgsContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_args + + return p +} + +func (s *ArgsContext) GetParser() antlr.Parser { return s.parser } + +func (s *ArgsContext) AllArg() []IArgContext { + children := s.GetChildren() + len := 0 + for _, ctx := range children { + if _, ok := ctx.(IArgContext); ok { + len++ + } + } + + tst := make([]IArgContext, len) + i := 0 + for _, ctx := range children { + if t, ok := ctx.(IArgContext); ok { + tst[i] = t.(IArgContext) + i++ + } + } + + return tst +} + +func (s *ArgsContext) Arg(i int) IArgContext { + var t antlr.RuleContext + j := 0 + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IArgContext); ok { + if j == i { + t = ctx.(antlr.RuleContext) + break + } + j++ + } + } + + if t == nil { + return nil + } + + return t.(IArgContext) +} + +func (s *ArgsContext) AllCOMMA() []antlr.TerminalNode { + return s.GetTokens(IRParserCOMMA) +} + +func (s *ArgsContext) COMMA(i int) antlr.TerminalNode { + return s.GetToken(IRParserCOMMA, i) +} + +func (s *ArgsContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ArgsContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +func (s *ArgsContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterArgs(s) + } +} + +func (s *ArgsContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitArgs(s) + } +} + +func (p *IRParser) Args() (localctx IArgsContext) { + localctx = NewArgsContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 18, IRParserRULE_args) + var _la int + + p.EnterOuterAlt(localctx, 1) + p.SetState(101) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + _la = p.GetTokenStream().LA(1) + + if (int64(_la) & ^0x3f) == 0 && ((int64(1)<<_la)&9664) != 0 { + { + p.SetState(93) + p.Arg() + } + p.SetState(98) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + _la = p.GetTokenStream().LA(1) + + for _la == IRParserCOMMA { + { + p.SetState(94) + p.Match(IRParserCOMMA) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(95) + p.Arg() + } + + p.SetState(100) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + _la = p.GetTokenStream().LA(1) + } + + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IArgContext is an interface to support dynamic dispatch. +type IArgContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + // IsArgContext differentiates from other interfaces. + IsArgContext() +} + +type ArgContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyArgContext() *ArgContext { + var p = new(ArgContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_arg + return p +} + +func InitEmptyArgContext(p *ArgContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_arg +} + +func (*ArgContext) IsArgContext() {} + +func NewArgContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *ArgContext { + var p = new(ArgContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_arg + + return p +} + +func (s *ArgContext) GetParser() antlr.Parser { return s.parser } + +func (s *ArgContext) CopyAll(ctx *ArgContext) { + s.CopyFrom(&ctx.BaseParserRuleContext) +} + +func (s *ArgContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ArgContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +type PositionalArgContext struct { + ArgContext +} + +func NewPositionalArgContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *PositionalArgContext { + var p = new(PositionalArgContext) + + InitEmptyArgContext(&p.ArgContext) + p.parser = parser + p.CopyAll(ctx.(*ArgContext)) + + return p +} + +func (s *PositionalArgContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *PositionalArgContext) Value() IValueContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IValueContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IValueContext) +} + +func (s *PositionalArgContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterPositionalArg(s) + } +} + +func (s *PositionalArgContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitPositionalArg(s) + } +} + +type LabeledArgContext struct { + ArgContext +} + +func NewLabeledArgContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *LabeledArgContext { + var p = new(LabeledArgContext) + + InitEmptyArgContext(&p.ArgContext) + p.parser = parser + p.CopyAll(ctx.(*ArgContext)) + + return p +} + +func (s *LabeledArgContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *LabeledArgContext) IDENTIFIER() antlr.TerminalNode { + return s.GetToken(IRParserIDENTIFIER, 0) +} + +func (s *LabeledArgContext) Value() IValueContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IValueContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IValueContext) +} + +func (s *LabeledArgContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterLabeledArg(s) + } +} + +func (s *LabeledArgContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitLabeledArg(s) + } +} + +func (p *IRParser) Arg() (localctx IArgContext) { + localctx = NewArgContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 20, IRParserRULE_arg) + p.SetState(107) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + + switch p.GetTokenStream().LA(1) { + case IRParserREG, IRParserLABEL, IRParserINT, IRParserLBRACKET: + localctx = NewPositionalArgContext(p, localctx) + p.EnterOuterAlt(localctx, 1) + { + p.SetState(103) + p.Value() + } + + case IRParserIDENTIFIER: + localctx = NewLabeledArgContext(p, localctx) + p.EnterOuterAlt(localctx, 2) + { + p.SetState(104) + p.Match(IRParserIDENTIFIER) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(105) + p.Match(IRParserT__0) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(106) + p.Value() + } + + default: + p.SetError(antlr.NewNoViableAltException(p, nil, nil, nil, nil, nil)) + goto errorExit + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IValueContext is an interface to support dynamic dispatch. +type IValueContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + // IsValueContext differentiates from other interfaces. + IsValueContext() +} + +type ValueContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyValueContext() *ValueContext { + var p = new(ValueContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_value + return p +} + +func InitEmptyValueContext(p *ValueContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_value +} + +func (*ValueContext) IsValueContext() {} + +func NewValueContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *ValueContext { + var p = new(ValueContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_value + + return p +} + +func (s *ValueContext) GetParser() antlr.Parser { return s.parser } + +func (s *ValueContext) CopyAll(ctx *ValueContext) { + s.CopyFrom(&ctx.BaseParserRuleContext) +} + +func (s *ValueContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ValueContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +type ValRegContext struct { + ValueContext +} + +func NewValRegContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *ValRegContext { + var p = new(ValRegContext) + + InitEmptyValueContext(&p.ValueContext) + p.parser = parser + p.CopyAll(ctx.(*ValueContext)) + + return p +} + +func (s *ValRegContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ValRegContext) Reg() IRegContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IRegContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IRegContext) +} + +func (s *ValRegContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterValReg(s) + } +} + +func (s *ValRegContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitValReg(s) + } +} + +type ValIntContext struct { + ValueContext +} + +func NewValIntContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *ValIntContext { + var p = new(ValIntContext) + + InitEmptyValueContext(&p.ValueContext) + p.parser = parser + p.CopyAll(ctx.(*ValueContext)) + + return p +} + +func (s *ValIntContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ValIntContext) INT() antlr.TerminalNode { + return s.GetToken(IRParserINT, 0) +} + +func (s *ValIntContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterValInt(s) + } +} + +func (s *ValIntContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitValInt(s) + } +} + +type ValLabelContext struct { + ValueContext +} + +func NewValLabelContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *ValLabelContext { + var p = new(ValLabelContext) + + InitEmptyValueContext(&p.ValueContext) + p.parser = parser + p.CopyAll(ctx.(*ValueContext)) + + return p +} + +func (s *ValLabelContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ValLabelContext) LABEL() antlr.TerminalNode { + return s.GetToken(IRParserLABEL, 0) +} + +func (s *ValLabelContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterValLabel(s) + } +} + +func (s *ValLabelContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitValLabel(s) + } +} + +type ValRegListContext struct { + ValueContext +} + +func NewValRegListContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *ValRegListContext { + var p = new(ValRegListContext) + + InitEmptyValueContext(&p.ValueContext) + p.parser = parser + p.CopyAll(ctx.(*ValueContext)) + + return p +} + +func (s *ValRegListContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ValRegListContext) LBRACKET() antlr.TerminalNode { + return s.GetToken(IRParserLBRACKET, 0) +} + +func (s *ValRegListContext) RegList() IRegListContext { + var t antlr.RuleContext + for _, ctx := range s.GetChildren() { + if _, ok := ctx.(IRegListContext); ok { + t = ctx.(antlr.RuleContext) + break + } + } + + if t == nil { + return nil + } + + return t.(IRegListContext) +} + +func (s *ValRegListContext) RBRACKET() antlr.TerminalNode { + return s.GetToken(IRParserRBRACKET, 0) +} + +func (s *ValRegListContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterValRegList(s) + } +} + +func (s *ValRegListContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitValRegList(s) + } +} + +func (p *IRParser) Value() (localctx IValueContext) { + localctx = NewValueContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 22, IRParserRULE_value) + p.SetState(116) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + + switch p.GetTokenStream().LA(1) { + case IRParserREG: + localctx = NewValRegContext(p, localctx) + p.EnterOuterAlt(localctx, 1) + { + p.SetState(109) + p.Reg() + } + + case IRParserLABEL: + localctx = NewValLabelContext(p, localctx) + p.EnterOuterAlt(localctx, 2) + { + p.SetState(110) + p.Match(IRParserLABEL) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + case IRParserINT: + localctx = NewValIntContext(p, localctx) + p.EnterOuterAlt(localctx, 3) + { + p.SetState(111) + p.Match(IRParserINT) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + case IRParserLBRACKET: + localctx = NewValRegListContext(p, localctx) + p.EnterOuterAlt(localctx, 4) + { + p.SetState(112) + p.Match(IRParserLBRACKET) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + { + p.SetState(113) + p.RegList() + } + { + p.SetState(114) + p.Match(IRParserRBRACKET) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + default: + p.SetError(antlr.NewNoViableAltException(p, nil, nil, nil, nil, nil)) + goto errorExit + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IConst_Context is an interface to support dynamic dispatch. +type IConst_Context interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + // IsConst_Context differentiates from other interfaces. + IsConst_Context() +} + +type Const_Context struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyConst_Context() *Const_Context { + var p = new(Const_Context) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_const_ + return p +} + +func InitEmptyConst_Context(p *Const_Context) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_const_ +} + +func (*Const_Context) IsConst_Context() {} + +func NewConst_Context(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *Const_Context { + var p = new(Const_Context) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_const_ + + return p +} + +func (s *Const_Context) GetParser() antlr.Parser { return s.parser } + +func (s *Const_Context) CopyAll(ctx *Const_Context) { + s.CopyFrom(&ctx.BaseParserRuleContext) +} + +func (s *Const_Context) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *Const_Context) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +type ConstStringContext struct { + Const_Context +} + +func NewConstStringContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *ConstStringContext { + var p = new(ConstStringContext) + + InitEmptyConst_Context(&p.Const_Context) + p.parser = parser + p.CopyAll(ctx.(*Const_Context)) + + return p +} + +func (s *ConstStringContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ConstStringContext) STRING() antlr.TerminalNode { + return s.GetToken(IRParserSTRING, 0) +} + +func (s *ConstStringContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterConstString(s) + } +} + +func (s *ConstStringContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitConstString(s) + } +} + +type ConstIntContext struct { + Const_Context +} + +func NewConstIntContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *ConstIntContext { + var p = new(ConstIntContext) + + InitEmptyConst_Context(&p.Const_Context) + p.parser = parser + p.CopyAll(ctx.(*Const_Context)) + + return p +} + +func (s *ConstIntContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ConstIntContext) INT() antlr.TerminalNode { + return s.GetToken(IRParserINT, 0) +} + +func (s *ConstIntContext) MINUS() antlr.TerminalNode { + return s.GetToken(IRParserMINUS, 0) +} + +func (s *ConstIntContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterConstInt(s) + } +} + +func (s *ConstIntContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitConstInt(s) + } +} + +type ConstBoolContext struct { + Const_Context +} + +func NewConstBoolContext(parser antlr.Parser, ctx antlr.ParserRuleContext) *ConstBoolContext { + var p = new(ConstBoolContext) + + InitEmptyConst_Context(&p.Const_Context) + p.parser = parser + p.CopyAll(ctx.(*Const_Context)) + + return p +} + +func (s *ConstBoolContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *ConstBoolContext) BOOL() antlr.TerminalNode { + return s.GetToken(IRParserBOOL, 0) +} + +func (s *ConstBoolContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterConstBool(s) + } +} + +func (s *ConstBoolContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitConstBool(s) + } +} + +func (p *IRParser) Const_() (localctx IConst_Context) { + localctx = NewConst_Context(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 24, IRParserRULE_const_) + var _la int + + p.SetState(124) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + + switch p.GetTokenStream().LA(1) { + case IRParserSTRING: + localctx = NewConstStringContext(p, localctx) + p.EnterOuterAlt(localctx, 1) + { + p.SetState(118) + p.Match(IRParserSTRING) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + case IRParserINT, IRParserMINUS: + localctx = NewConstIntContext(p, localctx) + p.EnterOuterAlt(localctx, 2) + p.SetState(120) + p.GetErrorHandler().Sync(p) + if p.HasError() { + goto errorExit + } + _la = p.GetTokenStream().LA(1) + + if _la == IRParserMINUS { + { + p.SetState(119) + p.Match(IRParserMINUS) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + } + { + p.SetState(122) + p.Match(IRParserINT) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + case IRParserBOOL: + localctx = NewConstBoolContext(p, localctx) + p.EnterOuterAlt(localctx, 3) + { + p.SetState(123) + p.Match(IRParserBOOL) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + + default: + p.SetError(antlr.NewNoViableAltException(p, nil, nil, nil, nil, nil)) + goto errorExit + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} + +// IRegContext is an interface to support dynamic dispatch. +type IRegContext interface { + antlr.ParserRuleContext + + // GetParser returns the parser. + GetParser() antlr.Parser + + // Getter signatures + REG() antlr.TerminalNode + + // IsRegContext differentiates from other interfaces. + IsRegContext() +} + +type RegContext struct { + antlr.BaseParserRuleContext + parser antlr.Parser +} + +func NewEmptyRegContext() *RegContext { + var p = new(RegContext) + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_reg + return p +} + +func InitEmptyRegContext(p *RegContext) { + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, nil, -1) + p.RuleIndex = IRParserRULE_reg +} + +func (*RegContext) IsRegContext() {} + +func NewRegContext(parser antlr.Parser, parent antlr.ParserRuleContext, invokingState int) *RegContext { + var p = new(RegContext) + + antlr.InitBaseParserRuleContext(&p.BaseParserRuleContext, parent, invokingState) + + p.parser = parser + p.RuleIndex = IRParserRULE_reg + + return p +} + +func (s *RegContext) GetParser() antlr.Parser { return s.parser } + +func (s *RegContext) REG() antlr.TerminalNode { + return s.GetToken(IRParserREG, 0) +} + +func (s *RegContext) GetRuleContext() antlr.RuleContext { + return s +} + +func (s *RegContext) ToStringTree(ruleNames []string, recog antlr.Recognizer) string { + return antlr.TreesStringTree(s, ruleNames, recog) +} + +func (s *RegContext) EnterRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.EnterReg(s) + } +} + +func (s *RegContext) ExitRule(listener antlr.ParseTreeListener) { + if listenerT, ok := listener.(IRListener); ok { + listenerT.ExitReg(s) + } +} + +func (p *IRParser) Reg() (localctx IRegContext) { + localctx = NewRegContext(p, p.GetParserRuleContext(), p.GetState()) + p.EnterRule(localctx, 26, IRParserRULE_reg) + p.EnterOuterAlt(localctx, 1) + { + p.SetState(126) + p.Match(IRParserREG) + if p.HasError() { + // Recognition error - abort rule + goto errorExit + } + } + +errorExit: + if p.HasError() { + v := p.GetError() + localctx.SetException(v) + p.GetErrorHandler().ReportError(p, v) + p.GetErrorHandler().Recover(p, v) + p.SetError(nil) + } + p.ExitRule() + return localctx + goto errorExit // Trick to prevent compiler error if the label is not used +} diff --git a/internal/ir/internal/syntax/ast.go b/internal/ir/internal/syntax/ast.go new file mode 100644 index 00000000..55606612 --- /dev/null +++ b/internal/ir/internal/syntax/ast.go @@ -0,0 +1,125 @@ +package syntax + +import ( + "github.com/formancehq/numscript/internal/parser" +) + +// ---- AST types for the IR textual format ---- + +// Program is the root of the parsed IR text. +type Program struct { + Stmts []Stmt +} + +// Stmt is either a LabelStmt or an InstrStmt. +type Stmt interface { + stmt() +} + +// LabelStmt represents a label marker line, e.g. "#inorder_end_0". +type LabelStmt struct { + Range parser.Range + Name string +} + +func (*LabelStmt) stmt() {} + +// InstrStmt represents one instruction line. +type InstrStmt struct { + Range parser.Range + + // Dest is the destination; nil when the instruction has no dest (e.g. "set_current_asset($r3)"). + Dest *Dest + + // One of these is set: + Call *InstrCall // e.g. "mk_monetary($r0, $r1)" + Const *Const // e.g. "$r0 = \"USD/2\"" or "$r0 = 42" + Infix *Infix // e.g. "$r3 = $r1 + $r2" + CompoundAssign *Infix // e.g. "$r5 += $r9" (Left is implicit from Dest) +} + +func (*InstrStmt) stmt() {} + +// Dest is the LHS of an assignment. +type Dest struct { + Range parser.Range + Kind DestKind + Regs []RegRef // non-empty for DestReg and DestList +} + +type DestKind int + +const ( + DestReg DestKind = iota // single register + DestDiscard // _ + DestList // [$r0, $r1] +) + +// RegRef is a parsed register reference, e.g. "$r0". +type RegRef struct { + Range parser.Range + Name string +} + +// InstrCall is a function-call instruction: name(args). +type InstrCall struct { + Range parser.Range + Name string // e.g. "mk_monetary", "pull_account" + TypeParam string // "" if none, else "int", "str", etc. + Args []Arg // may be empty +} + +// Arg is a single argument to an instruction. +type Arg struct { + Range parser.Range + Label string // empty for positional args; "account", "cap", etc. for labeled + Value Value +} + +// Value is a value that can appear as an argument. +type Value struct { + Range parser.Range + Kind ValueKind + + // Exactly one of these is set, depending on Kind: + Reg *RegRef // ValReg + Label *string // ValLabel: the label name without '#' + Int *string // ValInt: raw numeric string + Regs *[]RegRef // ValRegList +} + +type ValueKind int + +const ( + ValReg ValueKind = iota // $r0 + ValLabel // #my_label + ValInt // 42 + ValRegList // [$r0, $r1] +) + +// Const is a constant literal: a string, an integer or a bool. +type Const struct { + Range parser.Range + Kind ConstKind + + // Exactly one is set: + StrVal *string // raw quoted literal + IntVal *string // raw numeric string + BoolVal *bool +} + +type ConstKind int + +const ( + ConstString ConstKind = iota + ConstInt + ConstBool +) + +// Infix is a binary operation with infix syntax. +type Infix struct { + Range parser.Range + Op string // "+" or "-" + Left RegRef + Right RegRef +} diff --git a/internal/ir/internal/syntax/parser.go b/internal/ir/internal/syntax/parser.go new file mode 100644 index 00000000..d95a9b38 --- /dev/null +++ b/internal/ir/internal/syntax/parser.go @@ -0,0 +1,352 @@ +package syntax + +import ( + "github.com/formancehq/numscript/internal/parser" + + "github.com/antlr4-go/antlr/v4" + antlrParser "github.com/formancehq/numscript/internal/ir/internal/syntax/antlrParser" +) + +// ParserError is a parse error with range information. +type ParserError struct { + Range parser.Range + Msg string +} + +func (e ParserError) Error() string { + return e.Msg +} + +type ParseResult struct { + Value Program + Errors []ParserError +} + +type errorListener struct { + antlr.DefaultErrorListener + Errors []ParserError +} + +func (l *errorListener) SyntaxError(_ antlr.Recognizer, offendingSymbol any, startL, startC int, msg string, _ antlr.RecognitionException) { + length := 1 + if token, ok := offendingSymbol.(antlr.Token); ok { + length = len(token.GetText()) + } + endL := startL + endC := startC + length - 1 + l.Errors = append(l.Errors, ParserError{ + Msg: msg, + Range: parser.Range{ + Start: parser.Position{Character: startC, Line: startL - 1}, + End: parser.Position{Character: endC, Line: endL - 1}, + }, + }) +} + +// Parse parses an IR textual program and returns the AST. +func Parse(input string) ParseResult { + listener := &errorListener{} + + is := antlr.NewInputStream(input) + lexer := antlrParser.NewIRLexer(is) + lexer.RemoveErrorListeners() + lexer.AddErrorListener(listener) + + stream := antlr.NewCommonTokenStream(lexer, antlr.TokenDefaultChannel) + + p := antlrParser.NewIRParser(stream) + p.RemoveErrorListeners() + p.AddErrorListener(listener) + + tree := p.Program() + + // On a syntax error ANTLR's recovery leaves partial nodes behind (a call + // without its parens, an assignment without its rhs). Walking those means + // dereferencing tokens that were never matched, so don't: the errors are + // what the caller needs anyway. + if len(listener.Errors) > 0 { + return ParseResult{ + Errors: listener.Errors, + } + } + + return ParseResult{ + Value: buildAST(tree), + } +} + +func tokenToRange(tok antlr.Token) parser.Range { + startL := tok.GetLine() - 1 + startC := tok.GetColumn() + endC := startC + len(tok.GetText()) - 1 + return parser.Range{ + Start: parser.Position{Character: startC, Line: startL}, + End: parser.Position{Character: endC, Line: startL}, + } +} + +// ---- AST builder (walks the ANTLR parse tree) ---- + +func buildAST(tree antlrParser.IProgramContext) Program { + if tree == nil { + return Program{} + } + lines := tree.AllLine() + stmts := make([]Stmt, 0, len(lines)) + for _, l := range lines { + if s := buildStmt(l); s != nil { + stmts = append(stmts, s) + } + } + return Program{Stmts: stmts} +} + +func buildStmt(ctx antlrParser.ILineContext) Stmt { + if ctx == nil { + return nil + } + + // label marker + if lm := ctx.LabelMarker(); lm != nil { + return buildLabelMarker(lm) + } + + // instruction + if instr := ctx.Instruction(); instr != nil { + return buildInstruction(instr) + } + + return nil +} + +func buildLabelMarker(ctx antlrParser.ILabelMarkerContext) *LabelStmt { + tok := ctx.LABEL().GetSymbol() + name := tok.GetText()[1:] // strip '#' + return &LabelStmt{ + Range: tokenToRange(tok), + Name: name, + } +} + +func buildInstruction(ctx antlrParser.IInstructionContext) *InstrStmt { + switch c := ctx.(type) { + case *antlrParser.InstrWithDestContext: + return buildInstrWithDest(c) + case *antlrParser.InstrNoDestContext: + return buildInstrNoDest(c) + case *antlrParser.ConstAssignContext: + return buildConstAssign(c) + case *antlrParser.InfixInstrContext: + return buildInfixInstr(c) + case *antlrParser.CompoundAssignInstrContext: + return buildCompoundAssignInstr(c) + } + return nil +} + +func buildInstrWithDest(ctx *antlrParser.InstrWithDestContext) *InstrStmt { + dest := buildDest(ctx.Dest()) + call := buildInstrCall(ctx.InstrCall()) + rng := mergeRanges(dest.Range, call.Range) + return &InstrStmt{ + Range: rng, + Dest: &dest, + Call: &call, + } +} + +func buildInstrNoDest(ctx *antlrParser.InstrNoDestContext) *InstrStmt { + call := buildInstrCall(ctx.InstrCall()) + return &InstrStmt{ + Range: call.Range, + Call: &call, + } +} + +func buildConstAssign(ctx *antlrParser.ConstAssignContext) *InstrStmt { + dest := buildDest(ctx.Dest()) + c := buildConst(ctx.Const_()) + rng := mergeRanges(dest.Range, c.Range) + return &InstrStmt{ + Range: rng, + Dest: &dest, + Const: &c, + } +} + +func buildInfixInstr(ctx *antlrParser.InfixInstrContext) *InstrStmt { + dest := buildDest(ctx.Dest()) + left := buildRegRef(ctx.GetLeft()) + right := buildRegRef(ctx.GetRight()) + op := ctx.GetOp().GetText() + rng := mergeRanges(dest.Range, right.Range) + return &InstrStmt{ + Range: rng, + Dest: &dest, + Infix: &Infix{ + Range: rng, + Op: op, + Left: left, + Right: right, + }, + } +} + +func buildCompoundAssignInstr(ctx *antlrParser.CompoundAssignInstrContext) *InstrStmt { + left := buildRegRef(ctx.GetLeft()) + right := buildRegRef(ctx.GetRight()) + op := ctx.GetOp().GetText() + // strip the trailing '=' + infixOp := op[:len(op)-1] + rng := mergeRanges(left.Range, right.Range) + return &InstrStmt{ + Range: rng, + Dest: &Dest{Kind: DestReg, Regs: []RegRef{left}, Range: left.Range}, + CompoundAssign: &Infix{ + Range: rng, + Op: infixOp, + Left: left, + Right: right, + }, + } +} + +func buildDest(ctx antlrParser.IDestContext) Dest { + switch d := ctx.(type) { + case *antlrParser.DestRegContext: + reg := buildRegRef(d.Reg()) + return Dest{Kind: DestReg, Regs: []RegRef{reg}, Range: reg.Range} + case *antlrParser.DestDiscardContext: + tok := d.UNDERSCORE().GetSymbol() + return Dest{Kind: DestDiscard, Range: tokenToRange(tok)} + case *antlrParser.DestListContext: + regs := buildRegList(d.RegList()) + if len(regs) == 0 { + return Dest{Kind: DestList, Range: tokenToRange(d.LBRACKET().GetSymbol())} + } + rng := mergeRanges( + tokenToRange(d.LBRACKET().GetSymbol()), + tokenToRange(d.RBRACKET().GetSymbol()), + ) + return Dest{Kind: DestList, Regs: regs, Range: rng} + } + return Dest{} +} + +func buildRegList(ctx antlrParser.IRegListContext) []RegRef { + if ctx == nil { + return nil + } + allRegs := ctx.AllReg() + regs := make([]RegRef, len(allRegs)) + for i, r := range allRegs { + regs[i] = buildRegRef(r) + } + return regs +} + +func buildInstrCall(ctx antlrParser.IInstrCallContext) InstrCall { + nameCtx := ctx.InstrName() + name, typeParam := buildInstrName(nameCtx) + args := buildArgs(ctx.Args()) + + rngStart := tokenToRange(ctx.LPAREN().GetSymbol()) + rngEnd := tokenToRange(ctx.RPAREN().GetSymbol()) + + return InstrCall{ + Range: mergeRanges(rngStart, rngEnd), + Name: name, + TypeParam: typeParam, + Args: args, + } +} + +func buildInstrName(ctx antlrParser.IInstrNameContext) (name string, typeParam string) { + name = ctx.IDENTIFIER().GetText() + if tn := ctx.TypeName(); tn != nil { + typeParam = tn.(*antlrParser.TypeNameContext).TYPE_KEYWORD().GetText() + } + return +} + +func buildArgs(ctx antlrParser.IArgsContext) []Arg { + if ctx == nil { + return nil + } + allArgs := ctx.AllArg() + args := make([]Arg, len(allArgs)) + for i, a := range allArgs { + args[i] = buildArg(a) + } + return args +} + +func buildArg(ctx antlrParser.IArgContext) Arg { + switch a := ctx.(type) { + case *antlrParser.PositionalArgContext: + val := buildValue(a.Value()) + return Arg{Range: val.Range, Value: val} + case *antlrParser.LabeledArgContext: + label := a.IDENTIFIER().GetText() + val := buildValue(a.Value()) + rng := mergeRanges(tokenToRange(a.IDENTIFIER().GetSymbol()), val.Range) + return Arg{Range: rng, Label: label, Value: val} + } + return Arg{} +} + +func buildValue(ctx antlrParser.IValueContext) Value { + switch v := ctx.(type) { + case *antlrParser.ValRegContext: + reg := buildRegRef(v.Reg()) + return Value{Range: reg.Range, Kind: ValReg, Reg: ®} + case *antlrParser.ValLabelContext: + tok := v.LABEL().GetSymbol() + name := tok.GetText()[1:] // strip '#' + return Value{Range: tokenToRange(tok), Kind: ValLabel, Label: &name} + case *antlrParser.ValIntContext: + tok := v.INT().GetSymbol() + s := tok.GetText() + return Value{Range: tokenToRange(tok), Kind: ValInt, Int: &s} + case *antlrParser.ValRegListContext: + regs := buildRegList(v.RegList()) + lTok := v.LBRACKET().GetSymbol() + rTok := v.RBRACKET().GetSymbol() + rng := mergeRanges(tokenToRange(lTok), tokenToRange(rTok)) + return Value{Range: rng, Kind: ValRegList, Regs: ®s} + } + return Value{} +} + +func buildConst(ctx antlrParser.IConst_Context) Const { + switch c := ctx.(type) { + case *antlrParser.ConstStringContext: + tok := c.STRING().GetSymbol() + s := tok.GetText() + return Const{Range: tokenToRange(tok), Kind: ConstString, StrVal: &s} + case *antlrParser.ConstIntContext: + tok := c.INT().GetSymbol() + s := tok.GetText() + rng := tokenToRange(tok) + if minus := c.MINUS(); minus != nil { + s = "-" + s + rng.Start = tokenToRange(minus.GetSymbol()).Start + } + return Const{Range: rng, Kind: ConstInt, IntVal: &s} + case *antlrParser.ConstBoolContext: + tok := c.BOOL().GetSymbol() + b := tok.GetText() == "true" + return Const{Range: tokenToRange(tok), Kind: ConstBool, BoolVal: &b} + } + return Const{} +} + +func buildRegRef(ctx antlrParser.IRegContext) RegRef { + tok := ctx.REG().GetSymbol() + name := tok.GetText() + return RegRef{Range: tokenToRange(tok), Name: name} +} + +func mergeRanges(a, b parser.Range) parser.Range { + return parser.Range{Start: a.Start, End: b.End} +} diff --git a/internal/ir/mark_test.go b/internal/ir/mark_test.go new file mode 100644 index 00000000..a7f097bd --- /dev/null +++ b/internal/ir/mark_test.go @@ -0,0 +1,83 @@ +package ir + +import ( + "testing" + + "github.com/formancehq/numscript/internal/vm" + "github.com/stretchr/testify/require" +) + +// The mark ops carry no register, so there is nothing for the typechecker to get +// wrong — and, more to the point, no operand a caller could use to name a queue +// depth the run-state never marked. What does need checking about them (that pushes +// and ends balance, and that no send or asset change sits inside a region) is a +// control-flow property, checked by vm.Verify on the assembled program. +func TestMark_Typecheck(t *testing.T) { + require.NoError(t, Typecheck([]Instr{ + MarkPush{}, + MarkEnd{Rewind: true}, + MarkPush{}, + MarkEnd{Rewind: false}, + })) + + // with no operand there is no such thing as an ill-typed mark op: even an + // unbalanced stream is well-typed, and fails at run time instead + require.NoError(t, Typecheck([]Instr{MarkEnd{Rewind: false}})) + require.NoError(t, Typecheck([]Instr{ + LoadStr{Dest: 0, Value: "USD/2"}, + MarkEnd{Rewind: true}, + })) +} + +// One instruction, two textual names — the same shape as +// assert_leftover / assert_leftover_exact. +func TestMark_Dump(t *testing.T) { + out := Dump([]Instr{ + MarkPush{}, + MarkEnd{Rewind: true}, + MarkEnd{Rewind: false}, + }) + require.Contains(t, out, "mark_push()") + require.Contains(t, out, "mark_rewind()") + require.Contains(t, out, "mark_commit()") +} + +func TestMark_Assemble(t *testing.T) { + prog, err := Assemble([]Instr{ + MarkPush{}, + MarkEnd{Rewind: true}, + MarkEnd{Rewind: false}, + }) + require.NoError(t, err) + // the rewind flag rides in A; there is no register to allocate + require.Equal(t, []vm.Instruction{ + {Opcode: byte(vm.Op_MarkPush), A: 0xFF, B: 0xFF, C: 0xFF}, + {Opcode: byte(vm.Op_MarkEnd), A: 1, B: 0xFF, C: 0xFF}, + {Opcode: byte(vm.Op_MarkEnd), A: 0, B: 0xFF, C: 0xFF}, + }, prog.Instructions) + + // no register is consumed in any bank — the old Op_Snapshot spent a big.Int one + require.Zero(t, prog.MaxRegInt) +} + +func TestMark_ParseRoundTrip(t *testing.T) { + instrs, errs := Parse(` + mark_push() + mark_rewind() + mark_push() + mark_commit() +`) + require.Empty(t, errs) + require.Equal(t, []Instr{ + MarkPush{}, + MarkEnd{Rewind: true}, + MarkPush{}, + MarkEnd{Rewind: false}, + }, instrs) + + // a Dump of the parsed stream parses back to the same program, so the two names + // round-trip through the single instruction + reparsed, errs := Parse(Dump(instrs)) + require.Empty(t, errs) + require.Equal(t, instrs, reparsed) +} diff --git a/internal/ir/parse.go b/internal/ir/parse.go new file mode 100644 index 00000000..57e12121 --- /dev/null +++ b/internal/ir/parse.go @@ -0,0 +1,642 @@ +package ir + +import ( + "fmt" + "math/big" + "strconv" + + "github.com/formancehq/numscript/internal/ir/internal/syntax" + "github.com/formancehq/numscript/internal/parser" +) + +// Error is something wrong with an IR text: either the grammar rejected it, or +// it doesn't describe a well-formed instruction stream. +type Error struct { + Range parser.Range + Msg string +} + +func (e Error) Error() string { + return fmt.Sprintf("%d:%d: %s", e.Range.Start.Line+1, e.Range.Start.Character+1, e.Msg) +} + +// Parse reads an IR text into the instruction stream it describes. +// +// It checks the grammar and everything the grammar can't express: that +// instructions exist and take the arguments they were given, that labels resolve +// and are unique, and that jumps go forward. It does not typecheck the registers +// — that's Typecheck. +func Parse(src string) ([]Instr, []Error) { + parsed := syntax.Parse(src) + if len(parsed.Errors) > 0 { + errs := make([]Error, len(parsed.Errors)) + for i, e := range parsed.Errors { + errs[i] = Error{Range: e.Range, Msg: e.Msg} + } + return nil, errs + } + return transform(parsed.Value) +} + +// transformer carries the state shared by the whole transformation. +type transformer struct { + // labelPos maps each label defined in the program to its position, so a jump + // can be checked to both resolve and go forward. + labelPos map[string]int + // stmtPos is the position of the statement being transformed. + stmtPos int + // regByName binds each register name to the logical register it got on its + // first appearance, and nameByReg maps it back for error messages. + regByName map[string]Reg + nameByReg map[Reg]string + nextReg Reg + // written records the registers an instruction has assigned to so far, so a + // read of one that was never written can be reported. + written map[Reg]bool +} + +// regName spells a register the way the text did. +func (t *transformer) regName(r Reg) string { + if name, ok := t.nameByReg[r]; ok { + return name + } + return r.String() +} + +// freshReg allocates a register bound to no name. +func (t *transformer) freshReg() Reg { + r := t.nextReg + t.nextReg++ + return r +} + +// resolveReg returns the logical register a name refers to, allocating one on +// the name's first appearance. +func (t *transformer) resolveReg(rr syntax.RegRef) Reg { + if r, ok := t.regByName[rr.Name]; ok { + return r + } + r := t.freshReg() + t.regByName[rr.Name] = r + t.nameByReg[r] = rr.Name + return r +} + +// transform converts a parsed IR AST into a slice of Instr. +// It returns all instructions and any errors encountered. +func transform(prog syntax.Program) ([]Instr, []Error) { + var instrs []Instr + var errs []Error + + t := &transformer{ + labelPos: map[string]int{}, + regByName: map[string]Reg{}, + nameByReg: map[Reg]string{}, + written: map[Reg]bool{}, + } + // First pass: collect the labels and where they sit. + for pos, stmt := range prog.Stmts { + if ls, ok := stmt.(*syntax.LabelStmt); ok { + if _, seen := t.labelPos[ls.Name]; seen { + errs = append(errs, Error{Range: ls.Range, Msg: fmt.Sprintf("duplicate label #%s", ls.Name)}) + } + t.labelPos[ls.Name] = pos + } + } + + for pos, stmt := range prog.Stmts { + switch s := stmt.(type) { + case *syntax.LabelStmt: + instrs = append(instrs, LabelMarker{Label: Label(s.Name)}) + case *syntax.InstrStmt: + t.stmtPos = pos + instr, err := t.transformInstr(s) + if err != nil { + errs = append(errs, *err) + continue + } + + // Jumps only go forward, so text order is execution order: a read with + // no earlier write can't be reached by any path. + for _, r := range instr.sources() { + if !t.written[r] { + errs = append(errs, Error{ + Range: s.Range, + Msg: fmt.Sprintf("register %s is read but never written", t.regName(r)), + }) + } + } + for _, r := range instr.dests() { + t.written[r] = true + } + + instrs = append(instrs, instr) + } + } + + if len(errs) > 0 { + return instrs, errs + } + return instrs, nil +} + +func (t *transformer) transformInstr(s *syntax.InstrStmt) (Instr, *Error) { + switch { + case s.Const != nil: + return t.transformConst(s) + case s.Call != nil: + return t.transformCall(s) + case s.Infix != nil: + return t.transformInfix(s, s.Infix, "infix") + case s.CompoundAssign != nil: + return t.transformInfix(s, s.CompoundAssign, "compound assign") + default: + return nil, &Error{Range: s.Range, Msg: "empty instruction"} + } +} + +// ---- const assignment ---- + +func (t *transformer) transformConst(s *syntax.InstrStmt) (Instr, *Error) { + if s.Dest == nil || s.Dest.Kind != syntax.DestReg { + return nil, &Error{Range: s.Range, Msg: "const assignment requires a single register dest"} + } + dest := t.resolveReg(s.Dest.Regs[0]) + + switch s.Const.Kind { + case syntax.ConstString: + str, err := strconv.Unquote(*s.Const.StrVal) + if err != nil { + return nil, &Error{Range: s.Const.Range, Msg: fmt.Sprintf("invalid string literal: %s", *s.Const.StrVal)} + } + return LoadStr{Dest: dest, Value: str}, nil + case syntax.ConstInt: + n, ok := new(big.Int).SetString(*s.Const.IntVal, 10) + if !ok { + return nil, &Error{Range: s.Const.Range, Msg: fmt.Sprintf("invalid integer: %q", *s.Const.IntVal)} + } + return LoadInt{Dest: dest, Value: *n}, nil + case syntax.ConstBool: + return ConstBool{Dest: dest, Value: *s.Const.BoolVal}, nil + default: + return nil, &Error{Range: s.Range, Msg: "unknown const kind"} + } +} + +// ---- infix / compound assign ---- + +// transformInfix handles both `$d = $l + $r` and `$d += $r`: the parser gives +// the compound form the same shape, with left repeated as the dest. +func (t *transformer) transformInfix(s *syntax.InstrStmt, infix *syntax.Infix, what string) (Instr, *Error) { + if s.Dest == nil || s.Dest.Kind != syntax.DestReg { + return nil, &Error{Range: s.Range, Msg: what + " requires a single register dest"} + } + dest := t.resolveReg(s.Dest.Regs[0]) + left := t.resolveReg(infix.Left) + right := t.resolveReg(infix.Right) + + var op BinKind + switch infix.Op { + case "+": + op = OpAddInt{} + case "-": + op = OpSubInt{} + default: + return nil, &Error{Range: infix.Range, Msg: fmt.Sprintf("unknown infix operator: %q", infix.Op)} + } + return BinaryOp{Op: op, Dest: dest, Left: left, Right: right}, nil +} + +// ---- call instructions ---- + +// argParser reads the args of one instruction call. Accessors report a bad arg +// themselves and return a zero value, so callers read args straight into an Instr +// literal and transformCall checks ap.errs once at the end. Composite literal +// operands evaluate left to right, so `f{a: ap.reg(), b: ap.reg()}` reads in order. +type argParser struct { + t *transformer + args []syntax.Arg + pos int + seenLabel map[string]bool + errs *[]Error + // callRange is where to point an error about an arg that isn't there. + callRange parser.Range +} + +func (t *transformer) newArgParser(call *syntax.InstrCall, errs *[]Error) *argParser { + return &argParser{ + t: t, + args: call.Args, + seenLabel: map[string]bool{}, + errs: errs, + callRange: call.Range, + } +} + +// next consumes the next positional arg, checking it holds the expected kind. It +// reports missing and mistyped args itself and returns nil in both cases, so a +// caller only reads what it needs and the error shows up in ap.errs. +func (ap *argParser) next(want syntax.ValueKind) *syntax.Value { + if ap.pos >= len(ap.args) { + ap.addErr(ap.callRange, "missing %s argument", valueKindStr(want)) + return nil + } + a := ap.args[ap.pos] + ap.pos++ + if a.Value.Kind != want { + ap.addErr(a.Range, "expected %s, got %s", valueKindStr(want), valueKindStr(a.Value.Kind)) + return nil + } + return &a.Value +} + +// reg consumes the next positional arg as a register. +func (ap *argParser) reg() Reg { + v := ap.next(syntax.ValReg) + if v == nil { + return 0 + } + return ap.t.resolveReg(*v.Reg) +} + +// optLabeledReg consumes an optional labeled arg with the given label as a register. +func (ap *argParser) optLabeledReg(name string) *Reg { + a, ok := ap.labeledArg(name) + if !ok { + return nil + } + if a.Value.Kind != syntax.ValReg { + ap.addErr(a.Range, "labeled arg %q: expected register, got %s", name, valueKindStr(a.Value.Kind)) + return nil + } + r := ap.t.resolveReg(*a.Value.Reg) + return &r +} + +// reqLabeledReg is optLabeledReg for a label the instruction can't do without. +func (ap *argParser) reqLabeledReg(name string) Reg { + r := ap.optLabeledReg(name) + if r == nil { + ap.addErr(ap.callRange, "missing labeled argument %q", name) + return 0 + } + return *r +} + +// labeledArg finds a labeled arg by name. Returns nil if not found. +func (ap *argParser) labeledArg(name string) (*syntax.Arg, bool) { + if ap.seenLabel[name] { + return nil, false + } + for i := ap.pos; i < len(ap.args); i++ { + if ap.args[i].Label == name { + ap.seenLabel[name] = true + return &ap.args[i], true + } + } + // also check already-consumed positional area + for i := 0; i < ap.pos; i++ { + if ap.args[i].Label == name { + ap.seenLabel[name] = true + return &ap.args[i], true + } + } + return nil, false +} + +// labelRef consumes the next positional arg as a label reference. It returns "" +// when the arg is missing or isn't one. +func (ap *argParser) labelRef() Label { + v := ap.next(syntax.ValLabel) + if v == nil { + return "" + } + return Label(*v.Label) +} + +// intLit consumes the next positional arg as an integer literal. +func (ap *argParser) intLit() uint16 { + v := ap.next(syntax.ValInt) + if v == nil { + return 0 + } + n, err := strconv.ParseUint(*v.Int, 10, 16) + if err != nil { + ap.addErr(v.Range, "integer literal out of range (0-65535): %s", *v.Int) + return 0 + } + return uint16(n) +} + +func (ap *argParser) addErr(rng parser.Range, format string, args ...any) { + *ap.errs = append(*ap.errs, Error{Range: rng, Msg: fmt.Sprintf(format, args...)}) +} + +func valueKindStr(k syntax.ValueKind) string { + switch k { + case syntax.ValReg: + return "register" + case syntax.ValLabel: + return "label" + case syntax.ValInt: + return "integer literal" + case syntax.ValRegList: + return "register list" + default: + return "unknown" + } +} + +func (t *transformer) transformCall(s *syntax.InstrStmt) (Instr, *Error) { + var errs []Error + ap := t.newArgParser(s.Call, &errs) + + // A labeled arg is looked up by name, so a repeated one would silently lose + // every occurrence but the first. + seen := map[string]bool{} + for _, a := range s.Call.Args { + if a.Label == "" { + continue + } + if seen[a.Label] { + ap.addErr(a.Range, "duplicate labeled argument %q", a.Label) + } + seen[a.Label] = true + } + + // Resolve dest + var destReg *Reg + var dests []Reg + if s.Dest != nil { + switch s.Dest.Kind { + case syntax.DestReg: + r := t.resolveReg(s.Dest.Regs[0]) + destReg = &r + case syntax.DestDiscard: + // `_` only exists in the text: desugar it to a fresh register, which + // no other statement can name and nothing reads back. + r := t.freshReg() + destReg = &r + case syntax.DestList: + dests = make([]Reg, len(s.Dest.Regs)) + for i, r := range s.Dest.Regs { + dests[i] = t.resolveReg(r) + } + } + } + // A call with no dest discards its result, as `_` does. The fresh register + // is allocated only when the instruction writes one: registers are numbered + // by first appearance, so one allocated for an effect-only call would shift + // every later register and break the dump round-trip. + dest := func() Reg { + if destReg == nil { + r := t.freshReg() + destReg = &r + } + return *destReg + } + + name, typeParam := s.Call.Name, s.Call.TypeParam + + // load_var and meta are the only instructions parameterized by a type. + if typeParam != "" && name != "load_var" && name != "meta" { + return nil, &Error{Range: s.Call.Range, Msg: fmt.Sprintf("%s doesn't take a type parameter", name)} + } + + var instr Instr + + switch name { + case "load_var": + var typ VarType + switch typeParam { + case "int": + typ = VarInt{} + case "str": + typ = VarStr{} + default: + return nil, &Error{Range: s.Call.Range, Msg: fmt.Sprintf("load_var: expected type parameter int or str, got %q", typeParam)} + } + instr = LoadVar{Dest: dest(), Typ: typ, Index: ap.intLit()} + + case "meta": + var typ MetaType + switch typeParam { + case "str": + typ = MetaStr{} + case "int": + typ = MetaInt{} + case "portion": + typ = MetaPortion{} + default: + return nil, &Error{Range: s.Call.Range, Msg: fmt.Sprintf("meta: expected type parameter str, int or portion, got %q", typeParam)} + } + instr = MetaVar{Dest: dest(), Typ: typ, Account: ap.reg(), Key: ap.reg(), Scope: ap.optLabeledReg("scope")} + + case "balance": + instr = FetchBalance{Dest: dest(), Account: ap.reg(), Asset: ap.reg(), Scope: ap.optLabeledReg("scope")} + + case "monetary_to_string": + instr = ap.BinaryOp(dest(), OpMonetaryToString{}) + case "mk_portion": + instr = ap.BinaryOp(dest(), OpMakePortion{}) + case "add_int": + instr = ap.BinaryOp(dest(), OpAddInt{}) + case "sub_int": + instr = ap.BinaryOp(dest(), OpSubInt{}) + case "add_string": + instr = ap.BinaryOp(dest(), OpAddString{}) + case "str_eq": + instr = ap.BinaryOp(dest(), OpStrEq{}) + case "sub_portion": + instr = ap.BinaryOp(dest(), OpSubPortion{}) + case "mul_portion": + instr = ap.BinaryOp(dest(), OpMulPortion{}) + case "add_portion": + instr = ap.BinaryOp(dest(), OpAddPortion{}) + case "lt_int": + instr = ap.BinaryOp(dest(), OpLtInt{}) + case "eq_int": + instr = ap.BinaryOp(dest(), OpEqInt{}) + case "lt_portion": + instr = ap.BinaryOp(dest(), OpLtPortion{}) + case "eq_portion": + instr = ap.BinaryOp(dest(), OpEqPortion{}) + + case "int_copy": + instr = UnaryOp{Dest: dest(), Op: OpIntCopy{}, Arg: ap.reg()} + case "portion_copy": + instr = UnaryOp{Dest: dest(), Op: OpPortionCopy{}, Arg: ap.reg()} + case "str_copy": + instr = UnaryOp{Dest: dest(), Op: OpStrCopy{}, Arg: ap.reg()} + case "bool_copy": + instr = UnaryOp{Dest: dest(), Op: OpBoolCopy{}, Arg: ap.reg()} + case "neg_int": + instr = UnaryOp{Dest: dest(), Op: OpNegInt{}, Arg: ap.reg()} + case "int_to_string": + instr = UnaryOp{Dest: dest(), Op: OpIntToString{}, Arg: ap.reg()} + case "is_zero": + instr = UnaryOp{Dest: dest(), Op: OpIsZero{}, Arg: ap.reg()} + case "not": + instr = UnaryOp{Dest: dest(), Op: OpNot{}, Arg: ap.reg()} + case "portion_to_string": + instr = UnaryOp{Dest: dest(), Op: OpPortionToString{}, Arg: ap.reg()} + case "int_to_portion": + instr = UnaryOp{Dest: dest(), Op: OpIntToPortion{}, Arg: ap.reg()} + case "portion_to_int": + instr = UnaryOp{Dest: dest(), Op: OpPortionToInt{}, Arg: ap.reg()} + + case "pull_account": + instr = PullAccount{ + Dest: dest(), + Account: ap.reqLabeledReg("account"), + Cap: ap.optLabeledReg("cap"), + Overdraft: ap.optLabeledReg("overdraft"), + Color: ap.optLabeledReg("color"), + Scope: ap.optLabeledReg("scope"), + } + case "send_to_account": + instr = SendToAccount{ + Account: ap.optLabeledReg("account"), + Cap: ap.optLabeledReg("cap"), + Scope: ap.optLabeledReg("scope"), + } + case "save": + instr = Save{ + Account: ap.reqLabeledReg("account"), + Asset: ap.reqLabeledReg("asset"), + Amount: ap.optLabeledReg("amount"), + Scope: ap.optLabeledReg("scope"), + } + + case "meta_monetary": + if len(dests) != 2 { + ap.addErr(s.Range, "meta_monetary requires a dest list of 2 registers (asset, amount)") + break + } + instr = MetaMonetary{ + DestAsset: dests[0], + DestAmount: dests[1], + Account: ap.reg(), + Key: ap.reg(), + Scope: ap.optLabeledReg("scope"), + } + + case "check_enough_funds": + instr = CheckEnoughFunds{Got: ap.reg(), Needed: ap.reg()} + case "assert_leftover": + instr = AssertLeftover{Portion: ap.reg(), Exact: false} + case "assert_leftover_exact": + instr = AssertLeftover{Portion: ap.reg(), Exact: true} + case "set_current_asset": + instr = SetCurrentAsset{Asset: ap.reg()} + case "assert_same_asset": + instr = AssertSameAsset{Left: ap.reg(), Right: ap.reg()} + case "assert_valid_account": + instr = AssertValidAccount{Account: ap.reg()} + case "assert_valid_color": + instr = AssertValidColor{Color: ap.reg()} + case "assert_valid_scope": + instr = AssertValidScope{Scope: ap.reg()} + case "assert_unscoped": + instr = AssertUnscoped{Scope: ap.reg(), Account: ap.reg()} + case "assert_non_negative_balance": + instr = AssertNonNegativeBalance{Balance: ap.reg(), Account: ap.reg()} + case "assert_non_negative_amount": + instr = AssertNonNegativeAmount{Amount: ap.reg()} + case "assert_non_negative_portion": + instr = AssertNonNegativePortion{Portion: ap.reg()} + + // two names for one instruction, as assert_leftover/assert_leftover_exact are + case "mark_push": + instr = MarkPush{} + case "mark_rewind": + instr = MarkEnd{Rewind: true} + case "mark_commit": + instr = MarkEnd{Rewind: false} + + case "set_tx_meta": + instr = SetTxMeta{Key: ap.reg(), Value: ap.reg()} + case "set_account_meta": + instr = SetAccountMeta{Account: ap.reg(), Key: ap.reg(), Value: ap.reg(), Scope: ap.optLabeledReg("scope")} + + case "jmp_if_false": + cond, target := ap.reg(), ap.labelRef() + if t.checkJmpTarget(ap, s, name, target) { + instr = JmpIfFalse{Cond: cond, Target: target} + } + + case "jmp_if_true": + cond, target := ap.reg(), ap.labelRef() + if t.checkJmpTarget(ap, s, name, target) { + instr = JmpIfTrue{Cond: cond, Target: target} + } + + case "jmp": + target := ap.labelRef() + if t.checkJmpTarget(ap, s, name, target) { + instr = Jmp{Target: target} + } + + default: + return nil, &Error{Range: s.Call.Range, Msg: fmt.Sprintf("unknown instruction: %s", name)} + } + + // meta_monetary checks its own two-register dest list above + if instr != nil && name != "meta_monetary" { + switch n := len(instr.dests()); { + case n == 0 && s.Dest != nil: + ap.addErr(s.Range, "%s produces no value, so it takes no destination", name) + case n == 1 && s.Dest != nil && s.Dest.Kind == syntax.DestList: + ap.addErr(s.Range, "%s writes one register, not a list", name) + } + } + + // Check for unconsumed args (skip labeled args that were already seen) + for ap.pos < len(ap.args) { + a := ap.args[ap.pos] + if a.Label != "" && ap.seenLabel[a.Label] { + ap.pos++ + continue + } + ap.addErr(a.Range, "unexpected extra argument") + ap.pos++ + } + // Also check labeled args that weren't consumed + for i := range ap.args { + if ap.args[i].Label != "" && !ap.seenLabel[ap.args[i].Label] { + ap.addErr(ap.args[i].Range, "unknown labeled argument %q", ap.args[i].Label) + } + } + + if len(errs) > 0 { + // Return first error along with the instruction + return instr, &errs[0] + } + return instr, nil +} + +// checkJmpTarget reports whether target is a label the jump named name may +// reach: one that is defined, and not behind the jump. The VM only allows +// jumping forward — that's what makes every program terminate — so a backward +// jump is rejected here rather than assembled. +func (t *transformer) checkJmpTarget(ap *argParser, s *syntax.InstrStmt, name string, target Label) bool { + labelPos, defined := t.labelPos[string(target)] + switch { + case target == "": + // labelRef already said what was wrong + return false + case !defined: + ap.addErr(s.Range, "%s: label %s is not defined in the program", name, target) + return false + case labelPos < t.stmtPos: + ap.addErr(s.Range, "%s: label %s is behind the jump (jumps must go forward)", name, target) + return false + default: + return true + } +} + +// BinaryOp reads the two register args every binary instruction takes. +func (ap *argParser) BinaryOp(dest Reg, op BinKind) BinaryOp { + return BinaryOp{Dest: dest, Op: op, Left: ap.reg(), Right: ap.reg()} +} diff --git a/internal/ir/parse_test.go b/internal/ir/parse_test.go new file mode 100644 index 00000000..c79b0bea --- /dev/null +++ b/internal/ir/parse_test.go @@ -0,0 +1,922 @@ +package ir + +import ( + "math/big" + "testing" + + "github.com/stretchr/testify/require" +) + +// parseAndDump parses an IR text and re-dumps it. +func parseAndDump(t *testing.T, source string) ([]Instr, string) { + t.Helper() + instrs, errs := Parse(source) + require.Empty(t, errs, "IR errors: %v", errs) + return instrs, "\n" + Dump(instrs) +} + +// TestParseErrors checks what Parse rejects. +func TestParseErrors(t *testing.T) { + t.Run("unknown instruction", func(t *testing.T) { + _, errs := Parse(` + $r0 = no_such_instr($r1) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "unknown instruction") + }) + + t.Run("invalid string escape", func(t *testing.T) { + _, errs := Parse(` + $r0 = "a\q" +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "invalid string literal") + }) + + t.Run("valid string escapes", func(t *testing.T) { + instrs, errs := Parse(` + $r0 = "a\"b\\c\n" +`) + require.Empty(t, errs) + require.Equal(t, "a\"b\\c\n", instrs[0].(LoadStr).Value) + }) + + t.Run("destination shape must match what the instruction writes", func(t *testing.T) { + for name, src := range map[string]string{ + "single-result call with a dest list": " $a = 1\n [$b, $c] = neg_int($a)\n", + "effect-only call with a dest": " $a = \"USD/2\"\n $b = set_current_asset($a)\n", + "effect-only call with a discard dest": " $a = \"USD/2\"\n _ = set_current_asset($a)\n", + } { + t.Run(name, func(t *testing.T) { + _, errs := Parse(src) + require.Len(t, errs, 1) + require.Equal(t, 1, errs[0].Range.Start.Line) // 0-based: the second line + }) + } + + _, errs := Parse(" $a = 1\n _ = neg_int($a)\n") + require.Empty(t, errs) + }) + + // as if written `_ = load_var(0)`: the value lands in a fresh register + t.Run("a value-producing call without a dest discards its result", func(t *testing.T) { + instrs, errs := Parse(" $a = 1\n load_var(0)\n $b = neg_int($a)\n") + require.Empty(t, errs) + require.Equal(t, []Instr{ + LoadInt{Dest: 0, Value: *big.NewInt(1)}, + LoadVar{Dest: 1, Typ: VarInt{}, Index: 0}, + UnaryOp{Dest: 2, Op: OpNegInt{}, Arg: 0}, + }, instrs) + }) + + // an effect-only call takes no register, so the numbering after it is unchanged + t.Run("an effect-only call allocates no register", func(t *testing.T) { + instrs, errs := Parse(" $a = \"USD/2\"\n set_current_asset($a)\n $b = 1\n") + require.Empty(t, errs) + require.Equal(t, LoadInt{Dest: 1, Value: *big.NewInt(1)}, instrs[2]) + }) + + t.Run("negative int constant round-trips", func(t *testing.T) { + instrs := []Instr{LoadInt{Dest: 0, Value: *big.NewInt(-1)}} + parsed, errs := Parse(Dump(instrs)) + require.Empty(t, errs) + require.Equal(t, instrs, parsed) + }) + + t.Run("invalid arg type", func(t *testing.T) { + _, errs := Parse(` + $r0 = neg_int(42) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "expected register") + }) + + t.Run("unbound jmp label", func(t *testing.T) { + _, errs := Parse(` + jmp_if_false($r0, #missing_label) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "not defined") + }) + + t.Run("forward jmp", func(t *testing.T) { + _, errs := Parse(` + $r0 = 0 + jmp_if_false($r0, #my_label) +#my_label +`) + require.Empty(t, errs) + }) + + t.Run("backward jmp", func(t *testing.T) { + _, errs := Parse(` +#my_label + jmp_if_false($r0, #my_label) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "must go forward") + }) + + t.Run("forward unconditional jmp", func(t *testing.T) { + _, errs := Parse(` + jmp(#my_label) +#my_label +`) + require.Empty(t, errs) + }) + + t.Run("backward unconditional jmp", func(t *testing.T) { + _, errs := Parse(` +#my_label + jmp(#my_label) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "must go forward") + }) + + t.Run("unconditional jmp to an undefined label", func(t *testing.T) { + _, errs := Parse(` + jmp(#nope) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "not defined") + }) + + t.Run("missing required labeled arg", func(t *testing.T) { + _, errs := Parse(` + $r0 = "acc" + $r1 = 1 + $r2 = pull_account(cap: $r1) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, `missing labeled argument "account"`) + }) + + t.Run("labeled arg no instruction takes", func(t *testing.T) { + // a stray label on an arg that was consumed positionally: the extra-args + // loop can't see it, so it's the leftover-label check that reports it + _, errs := Parse(` + $r0 = 1 + $r1 = 2 + check_enough_funds(foo: $r0, $r1) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, `unknown labeled argument "foo"`) + }) + + t.Run("labeled arg an instruction can't place", func(t *testing.T) { + _, errs := Parse(` + $r0 = "acc" + $r1 = pull_account(account: $r0, nope: $r0) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "unexpected extra argument") + }) + + t.Run("labels are case sensitive", func(t *testing.T) { + // IDENTIFIER is lowercase-only, so this doesn't even lex + _, errs := Parse(` + $r0 = "acc" + $r1 = pull_account(Account: $r0) +`) + require.NotEmpty(t, errs) + }) + + t.Run("duplicate labeled arg", func(t *testing.T) { + _, errs := Parse(` + $r0 = "acc" + $r1 = 1 + $r2 = pull_account(account: $r0, cap: $r1, cap: $r1) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "duplicate labeled argument") + }) + + t.Run("duplicate label", func(t *testing.T) { + _, errs := Parse(` + jmp_if_false($r0, #l) +#l +#l +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "duplicate label") + }) + + t.Run("const assigned to a dest list", func(t *testing.T) { + _, errs := Parse(` + [$r0, $r1] = 1 +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "const assignment requires a single register dest") + }) + + t.Run("infix assigned to a dest list", func(t *testing.T) { + _, errs := Parse(` + $r0 = 1 + [$r1, $r2] = $r0 + $r0 +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "infix requires a single register dest") + }) + + t.Run("compound assign to a dest list", func(t *testing.T) { + // the grammar only allows a single register left of `+=`, so this one is + // rejected before the transform sees it + _, errs := Parse(` + $r0 = 1 + [$r1, $r2] += $r0 +`) + require.NotEmpty(t, errs) + }) + + t.Run("labeled arg of the wrong kind", func(t *testing.T) { + _, errs := Parse(` + $r0 = pull_account(account: 42) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, `labeled arg "account": expected register, got integer literal`) + }) + + t.Run("register where a label is expected", func(t *testing.T) { + _, errs := Parse(` + $r0 = 1 + jmp_if_false($r0, $r0) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "expected label, got register") + }) + + t.Run("register list where a register is expected", func(t *testing.T) { + _, errs := Parse(` + $r0 = "USD/2" + set_current_asset([$r0, $r0]) +`) + require.NotEmpty(t, errs) + require.Contains(t, errs[0].Msg, "expected register, got register list") + }) +} + +// The infix forms are sugar: the call form of every binary op parses too, even +// though ir.Dump never prints it for add_int / sub_int. +func TestBinaryOpCallForms(t *testing.T) { + instrs, errs := Parse(` + $i = 1 + $s = "x" + $p = mk_portion($i, $i) + $sum = add_int($i, $i) + $diff = sub_int($i, $i) + $cat = add_string($s, $s) + $rest = sub_portion($p, $p) + $tot = add_portion($p, $p) + $label = monetary_to_string($s, $i) +`) + require.Empty(t, errs) + require.NoError(t, Typecheck(instrs)) + + // the dump switches the two int ops back to their infix spelling + require.Contains(t, Dump(instrs), "$r3 = $r0 + $r0") + require.Contains(t, Dump(instrs), "$r4 = $r0 - $r0") + require.Contains(t, Dump(instrs), "$r5 = add_string($r1, $r1)") + // only the int ops have infix sugar: `+` and `-` bind to add_int/sub_int, and + // ir.Parse doesn't typecheck, so it couldn't dispatch on operand type anyway + require.Contains(t, Dump(instrs), "$r6 = sub_portion($r2, $r2)") + require.Contains(t, Dump(instrs), "$r7 = add_portion($r2, $r2)") +} + +func TestParseErrorMessageFormat(t *testing.T) { + _, errs := Parse(` + $r0 = no_such_instr($r1) +`) + require.NotEmpty(t, errs) + // 1-based line:character, then the reason + require.Regexp(t, `^\d+:\d+: .*unknown instruction`, errs[0].Error()) +} + +// TestReadBeforeWrite checks that reading a register nothing ever assigned to is +// rejected, and reported under the name the text used. +func TestReadBeforeWrite(t *testing.T) { + t.Run("never written at all", func(t *testing.T) { + _, errs := Parse(` + $a = 42 + $y = lt_int($a, $b) +`) + require.Len(t, errs, 1) + require.Contains(t, errs[0].Msg, "$b is read but never written") + }) + + t.Run("written only after the read", func(t *testing.T) { + _, errs := Parse(` + $a = 42 + $y = lt_int($a, $b) + $b = 1 +`) + require.Len(t, errs, 1) + require.Contains(t, errs[0].Msg, "$b is read but never written") + }) + + t.Run("compound assign reads its own dest", func(t *testing.T) { + _, errs := Parse(` + $b = 1 + $acc += $b +`) + require.Len(t, errs, 1) + require.Contains(t, errs[0].Msg, "$acc is read but never written") + }) + + t.Run("labeled args are reads too", func(t *testing.T) { + _, errs := Parse(` + $acc = "src" + $pulled = pull_account(account: $acc, cap: $missing) +`) + require.Len(t, errs, 1) + require.Contains(t, errs[0].Msg, "$missing is read but never written") + }) + + t.Run("a dest is written for later instructions", func(t *testing.T) { + _, errs := Parse(` + $a = 1 + $b = 2 + $sum = $a + $b + $twice = $sum + $sum +`) + require.Empty(t, errs) + }) + + t.Run("dest list entries count as written", func(t *testing.T) { + _, errs := Parse(` + $acct = "acct" + $key = "key" + [$asset, $amount] = meta_monetary($acct, $key) + check_enough_funds($amount, $amount) +`) + require.Empty(t, errs) + }) +} + +// TestMalformedInputIsRejected checks that text → Instr reports errors on +// invalid input rather than panicking. It doesn't typecheck: only syntax and +// the structural rules (labels resolve, jumps go forward) are checked here. +func TestMalformedInputIsRejected(t *testing.T) { + sources := []struct { + name string + ir string + }{ + {"comment", "// not a comment in this format\n $r0 = 1\n"}, + {"no args at all", " $r0 = neg_int()\n"}, + {"too few args", " $r0 = balance($r1)\n"}, + {"too many args", " $r0 = neg_int($r1, $r2)\n"}, + {"missing required labeled arg", " $r0 = pull_account(cap: $r1)\n"}, + {"unknown labeled arg", " $r0 = pull_account(account: $r1, nope: $r2)\n"}, + {"capitalised label", " $r0 = pull_account(Account: $r1)\n"}, + {"load_var index out of range", " $r0 = load_var(70000)\n"}, + {"load_var without type param", " $r0 = load_var(0)\n"}, + {"load_var with a type it doesn't have", " $r0 = load_var(0)\n"}, + {"meta without type param", " $r0 = meta($r1, $r2)\n"}, + {"meta_monetary without dest list", " $r0 = meta_monetary($r1, $r2)\n"}, + {"reg to reg copy", " $r0 = $r1\n"}, + {"garbage", "$$$ !!!"}, + {"unclosed paren", " $r0 = neg_int($r1"}, + {"uppercase instr name", " $r0 = NEG_INT($r1)"}, + {"empty dest list", " [] = meta_monetary($r0, $r1)\n"}, + {"missing dest", " = neg_int($r0)\n"}, + {"unterminated string", " $r0 = \"oops\n"}, + {"stray operator", " $r0 = $r1 * $r2\n"}, + {"type param on plain instr", " $r0 = neg_int($r1)\n"}, + {"label as instr arg", " set_current_asset(#lbl)\n"}, + // a bool is a const, never an operand: no instruction takes one inline + {"bool as instr arg", " set_current_asset(true)\n"}, + {"bool as labeled arg", " $r0 = pull_account(account: false)\n"}, + {"bool in a dest list", " [$r0, $r1] = true\n"}, + {"capitalised bool", " $r0 = True\n"}, + } + + for _, s := range sources { + t.Run(s.name, func(t *testing.T) { + instrs, errs := Parse(s.ir) + require.NotEmpty(t, errs, "neither the parser nor the transform rejected it") + // and whatever it did return must be usable: a nil instruction in the + // stream would blow up in dump or assemble instead + for _, instr := range instrs { + require.NotNil(t, instr) + } + }) + } +} + +// TestRegNamesBindInOrder checks how a name becomes a logical register: the +// first appearance allocates the next one, later appearances reuse it. The name +// itself carries no meaning — `$r` is a convention, not an index. +func TestRegNamesBindInOrder(t *testing.T) { + _, dumped := parseAndDump(t, ` + $asset = "USD/2" + $amount = 10 + $label = monetary_to_string($asset, $amount) + $twice = add_int($amount, $amount) + $r99 = int_to_string($twice) +`) + + require.Equal(t, ` + $r0 = "USD/2" + $r1 = 10 + $r2 = monetary_to_string($r0, $r1) + $r3 = $r1 + $r1 + $r4 = int_to_string($r3) +`, dumped) +} + +// TestDiscardDestDesugarsToFreshReg checks that `_` becomes a register no +// statement can name, and that two discards don't alias — otherwise they'd be +// forced to share a type. +func TestDiscardDestDesugarsToFreshReg(t *testing.T) { + instrs, dumped := parseAndDump(t, ` + $r0 = "acc" + $r1 = "USD/2" + _ = pull_account(account: $r0) + _ = balance($r0, $r1) +`) + + pulled := instrs[2].dests()[0] + balance := instrs[3].dests()[0] + require.NotEqual(t, pulled, balance) + // above every register the text refers to ($r0, $r1) + require.Greater(t, uint(pulled), uint(1)) + require.Greater(t, uint(balance), uint(1)) + + // so the typechecker doesn't see one register written with two types + require.NoError(t, Typecheck(instrs)) + + // the IR has no notion of a discard: it dumps as the register it desugared to + require.Equal(t, ` + $r0 = "acc" + $r1 = "USD/2" + $r2 = pull_account(account: $r0) + $r3 = balance($r0, $r1) +`, dumped) +} + +// TestRoundtripAllInstructions tests every instruction in isolation for roundtrip. +func TestRoundtripAllInstructions(t *testing.T) { + tests := []struct { + name string + ir string + }{ + { + name: "LoadStr", + ir: ` + $r0 = "hello" +`, + }, + { + name: "LoadInt", + ir: ` + $r0 = 42 +`, + }, + { + name: "mk_portion", + ir: ` + $r0 = 1 + $r1 = 2 + $r2 = mk_portion($r0, $r1) +`, + }, + { + name: "add_int via infix", + ir: ` + $r0 = 1 + $r1 = 2 + $r2 = $r0 + $r1 +`, + }, + { + name: "infix add", + ir: ` + $r0 = 1 + $r1 = 2 + $r2 = $r0 + $r1 +`, + }, + { + name: "compound add", + ir: ` + $r0 = 0 + $r1 = 1 + $r0 += $r1 +`, + }, + { + name: "infix sub", + ir: ` + $r0 = 5 + $r1 = 3 + $r2 = $r0 - $r1 +`, + }, + { + name: "compound sub", + ir: ` + $r0 = 5 + $r1 = 3 + $r0 -= $r1 +`, + }, + { + name: "unary ops", + ir: ` + $r0 = 10 + $r1 = neg_int($r0) + $r2 = int_copy($r0) + $r3 = int_to_string($r0) +`, + }, + { + name: "pull_account with all labeled args", + ir: ` + $r0 = "src" + $r1 = 100 + $r2 = 0 + $r3 = "red" + $r4 = pull_account(account: $r0, cap: $r1, overdraft: $r2, color: $r3) +`, + }, + { + name: "pull_account minimal", + ir: ` + $r0 = "src" + $r1 = pull_account(account: $r0) +`, + }, + { + name: "send_to_account", + ir: ` + $r0 = "dest" + send_to_account(account: $r0) +`, + }, + { + name: "send_to_account with cap", + ir: ` + $r0 = "dest" + $r1 = 50 + send_to_account(account: $r0, cap: $r1) +`, + }, + { + name: "save with amount", + ir: ` + $r0 = "acct" + $r1 = "USD/2" + $r2 = 100 + save(account: $r0, asset: $r1, amount: $r2) +`, + }, + { + name: "save all", + ir: ` + $r0 = "acct" + $r1 = "USD/2" + save(account: $r0, asset: $r1) +`, + }, + { + name: "check_enough_funds", + ir: ` + $r0 = 50 + $r1 = 100 + check_enough_funds($r0, $r1) +`, + }, + { + name: "assert_leftover", + ir: ` + $r0 = 1 + $r1 = 1 + $r2 = mk_portion($r0, $r1) + assert_leftover($r2) +`, + }, + { + name: "assert_leftover_exact", + ir: ` + $r0 = 1 + $r1 = 1 + $r2 = mk_portion($r0, $r1) + assert_leftover_exact($r2) +`, + }, + { + name: "assert_same_asset", + ir: ` + $r0 = "USD/2" + $r1 = "EUR/2" + assert_same_asset($r0, $r1) +`, + }, + { + name: "assert_valid_account", + ir: ` + $r0 = "users:alice" + assert_valid_account($r0) +`, + }, + { + name: "assert_valid_color", + ir: ` + $r0 = "RED" + assert_valid_color($r0) +`, + }, + { + name: "assert_non_negative_balance", + ir: ` + $r0 = 100 + $r1 = "src" + assert_non_negative_balance($r0, $r1) +`, + }, + { + name: "set_tx_meta", + ir: ` + $r0 = "key" + $r1 = "value" + set_tx_meta($r0, $r1) +`, + }, + { + name: "set_account_meta", + ir: ` + $r0 = "acct" + $r1 = "key" + $r2 = "value" + set_account_meta($r0, $r1, $r2) +`, + }, + { + name: "set_current_asset", + ir: ` + $r0 = "USD/2" + set_current_asset($r0) +`, + }, + { + name: "balance", + ir: ` + $r0 = "src" + $r1 = "USD/2" + $r2 = balance($r0, $r1) +`, + }, + { + name: "meta str", + ir: ` + $r0 = "acct" + $r1 = "key" + $r2 = meta($r0, $r1) +`, + }, + { + name: "meta int", + ir: ` + $r0 = "acct" + $r1 = "key" + $r2 = meta($r0, $r1) +`, + }, + { + name: "meta portion", + ir: ` + $r0 = "acct" + $r1 = "key" + $r2 = meta($r0, $r1) +`, + }, + { + name: "meta_monetary", + ir: ` + $r0 = "acct" + $r1 = "key" + [$r2, $r3] = meta_monetary($r0, $r1) +`, + }, + { + name: "mul_portion", + ir: ` + $r0 = 1 + $r1 = 2 + $r2 = mk_portion($r0, $r1) + $r3 = mul_portion($r2, $r2) +`, + }, + { + name: "int_to_portion", + ir: ` + $r0 = 100 + $r1 = int_to_portion($r0) +`, + }, + { + name: "portion_to_int", + ir: ` + $r0 = 1 + $r1 = 2 + $r2 = mk_portion($r0, $r1) + $r3 = portion_to_int($r2) +`, + }, + { + name: "load_var int", + ir: ` + $r0 = load_var(0) +`, + }, + { + name: "load_var str", + ir: ` + $r0 = load_var(1) +`, + }, + { + name: "mark push, rewind and commit", + ir: ` + mark_push() + mark_rewind() + mark_push() + mark_commit() +`, + }, + { + name: "jmp_if_false and label", + ir: ` + $r0 = true + jmp_if_false($r0, #my_label) +#my_label +`, + }, + { + name: "jmp_if_true and label", + ir: ` + $r0 = false + jmp_if_true($r0, #my_label) +#my_label +`, + }, + { + name: "is_zero", + ir: ` + $r0 = 1 + $r1 = is_zero($r0) +`, + }, + { + name: "jmp and label", + ir: ` + jmp(#my_label) +#my_label +`, + }, + { + name: "str_eq", + ir: ` + $r0 = "a" + $r1 = "b" + $r2 = str_eq($r0, $r1) +`, + }, + { + name: "sub_int via infix", + ir: ` + $r0 = 5 + $r1 = 3 + $r2 = $r0 - $r1 +`, + }, + { + name: "add_string", + ir: ` + $r0 = "hello" + $r1 = "world" + $r2 = add_string($r0, $r1) +`, + }, + { + name: "lt_int", + ir: ` + $r0 = 10 + $r1 = 5 + $r2 = lt_int($r0, $r1) +`, + }, + { + name: "eq_int", + ir: ` + $r0 = 10 + $r1 = 5 + $r2 = eq_int($r0, $r1) +`, + }, + { + name: "add_portion", + ir: ` + $r0 = 1 + $r1 = 2 + $r2 = mk_portion($r0, $r1) + $r3 = add_portion($r2, $r2) +`, + }, + { + name: "sub_portion", + ir: ` + $r0 = 1 + $r1 = 2 + $r2 = mk_portion($r0, $r1) + $r3 = sub_portion($r2, $r2) +`, + }, + { + name: "lt_portion", + ir: ` + $r0 = 1 + $r1 = 2 + $r2 = mk_portion($r0, $r1) + $r3 = lt_portion($r2, $r2) +`, + }, + { + name: "eq_portion", + ir: ` + $r0 = 1 + $r1 = 2 + $r2 = mk_portion($r0, $r1) + $r3 = eq_portion($r2, $r2) +`, + }, + { + name: "not", + ir: ` + $r0 = true + $r1 = not($r0) +`, + }, + { + name: "portion_copy", + ir: ` + $r0 = 1 + $r1 = 1 + $r2 = mk_portion($r0, $r1) + $r3 = portion_copy($r2) +`, + }, + { + name: "str_copy", + ir: ` + $r0 = "USD/2" + $r1 = str_copy($r0) +`, + }, + { + name: "bool_copy", + ir: ` + $r0 = true + $r1 = bool_copy($r0) +`, + }, + { + name: "portion_to_string", + ir: ` + $r0 = 1 + $r1 = 1 + $r2 = mk_portion($r0, $r1) + $r3 = portion_to_string($r2) +`, + }, + { + name: "monetary_to_string", + ir: ` + $r0 = "USD/2" + $r1 = 100 + $r2 = monetary_to_string($r0, $r1) +`, + }, + { + name: "bool true", + ir: ` + $r0 = true +`, + }, + { + name: "bool false", + ir: ` + $r0 = false +`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + // source already has leading newline and indentation + source := tt.ir + _, roundtripped := parseAndDump(t, source) + require.Equal(t, source, roundtripped) + }) + } +} diff --git a/internal/ir/typecheck.go b/internal/ir/typecheck.go new file mode 100644 index 00000000..83a247e4 --- /dev/null +++ b/internal/ir/typecheck.go @@ -0,0 +1,291 @@ +package ir + +import "fmt" + +// regType is the type of a virtual register. It mirrors the VM register banks; +// every register has exactly one type for its whole life. A monetary is not one +// of them: it is a (regStr asset, regInt amount) pair. +type regType int + +const ( + regInt regType = iota + regStr + regPortion + regBool +) + +func (t regType) String() string { + switch t { + case regInt: + return "int" + case regStr: + return "string" + case regPortion: + return "portion" + case regBool: + return "bool" + default: + return "?" + } +} + +// bytecodeTypechecker validates an IR stream one instruction at a time, +// remembering the type each register was written with. A later read of that +// register with a different type, or a read before any write, is a bug in whatever +// produced the instructions. +type bytecodeTypechecker struct { + types map[Reg]regType +} + +func newBytecodeTypechecker() *bytecodeTypechecker { + return &bytecodeTypechecker{types: map[Reg]regType{}} +} + +// use asserts r was already written with type want. +func (tc *bytecodeTypechecker) use(r Reg, want regType) error { + got, ok := tc.types[r] + if !ok { + return fmt.Errorf("register %s read as %s before being written", r, want) + } + if got != want { + return fmt.Errorf("register %s read as %s but holds %s", r, want, got) + } + return nil +} + +func (tc *bytecodeTypechecker) useOpt(r *Reg, want regType) error { + if r == nil { + return nil + } + return tc.use(*r, want) +} + +// def records that r now holds type t, rejecting a write that changes its type. +func (tc *bytecodeTypechecker) def(r Reg, t regType) error { + if got, ok := tc.types[r]; ok && got != t { + return fmt.Errorf("register %s written as %s but already holds %s", r, t, got) + } + tc.types[r] = t + return nil +} + +// check typechecks a single instruction, updating the state on success. +func (tc *bytecodeTypechecker) check(instr Instr) error { + switch i := instr.(type) { + case LoadInt: + return tc.def(i.Dest, regInt) + case LoadStr: + return tc.def(i.Dest, regStr) + case ConstBool: + return tc.def(i.Dest, regBool) + case LoadVar: + t, err := varRegType(i.Typ) + if err != nil { + return err + } + return tc.def(i.Dest, t) + + case UnaryOp: + dest, arg, err := unOpRegTypes(i.Op) + if err != nil { + return err + } + return firstErr(tc.use(i.Arg, arg), tc.def(i.Dest, dest)) + case BinaryOp: + dest, left, right, err := binOpRegTypes(i.Op) + if err != nil { + return err + } + return firstErr(tc.use(i.Left, left), tc.use(i.Right, right), tc.def(i.Dest, dest)) + + case PullAccount: + return firstErr( + tc.use(i.Account, regStr), + tc.useOpt(i.Cap, regInt), + tc.useOpt(i.Overdraft, regInt), + tc.useOpt(i.Color, regStr), + tc.useOpt(i.Scope, regStr), + tc.def(i.Dest, regInt), + ) + case SendToAccount: + return firstErr(tc.useOpt(i.Account, regStr), tc.useOpt(i.Cap, regInt), tc.useOpt(i.Scope, regStr)) + case Save: + return firstErr( + tc.use(i.Account, regStr), + tc.use(i.Asset, regStr), + tc.useOpt(i.Amount, regInt), + tc.useOpt(i.Scope, regStr), + ) + + case CheckEnoughFunds: + return firstErr(tc.use(i.Got, regInt), tc.use(i.Needed, regInt)) + case AssertLeftover: + return tc.use(i.Portion, regPortion) + case SetCurrentAsset: + return tc.use(i.Asset, regStr) + case AssertSameAsset: + return firstErr(tc.use(i.Left, regStr), tc.use(i.Right, regStr)) + case AssertValidAccount: + return tc.use(i.Account, regStr) + case AssertValidColor: + return tc.use(i.Color, regStr) + case AssertValidScope: + return tc.use(i.Scope, regStr) + case AssertUnscoped: + return firstErr(tc.use(i.Scope, regStr), tc.use(i.Account, regStr)) + case AssertNonNegativeBalance: + return firstErr(tc.use(i.Balance, regInt), tc.use(i.Account, regStr)) + case AssertNonNegativeAmount: + return tc.use(i.Amount, regInt) + case AssertNonNegativePortion: + return tc.use(i.Portion, regPortion) + + case SetTxMeta: + return firstErr(tc.use(i.Key, regStr), tc.use(i.Value, regStr)) + case SetAccountMeta: + return firstErr( + tc.use(i.Account, regStr), + tc.use(i.Key, regStr), + tc.use(i.Value, regStr), + tc.useOpt(i.Scope, regStr), + ) + case MetaVar: + t, err := metaRegType(i.Typ) + if err != nil { + return err + } + return firstErr( + tc.use(i.Account, regStr), + tc.use(i.Key, regStr), + tc.useOpt(i.Scope, regStr), + tc.def(i.Dest, t), + ) + case MetaMonetary: + return firstErr( + tc.use(i.Account, regStr), + tc.use(i.Key, regStr), + tc.useOpt(i.Scope, regStr), + tc.def(i.DestAsset, regStr), + tc.def(i.DestAmount, regInt), + ) + case FetchBalance: + return firstErr( + tc.use(i.Account, regStr), + tc.use(i.Asset, regStr), + tc.useOpt(i.Scope, regStr), + tc.def(i.Dest, regInt), + ) + + case JmpIfFalse: + return tc.use(i.Cond, regBool) + case JmpIfTrue: + return tc.use(i.Cond, regBool) + case Jmp: + return nil + case LabelMarker: + return nil + + // the mark ops touch no register. What does need checking about them is a + // control-flow property, not a register one; see the note on MarkPush in instr.go + case MarkPush, MarkEnd: + return nil + + default: + return fmt.Errorf("bytecode typechecker: unhandled instruction %T", instr) + } +} + +func Typecheck(instrs []Instr) error { + tc := newBytecodeTypechecker() + for pos, instr := range instrs { + if err := tc.check(instr); err != nil { + return fmt.Errorf("at instruction %d (%s): %w", pos, instr, err) + } + } + return nil +} + +func firstErr(errs ...error) error { + for _, e := range errs { + if e != nil { + return e + } + } + return nil +} + +func varRegType(t VarType) (regType, error) { + switch t.(type) { + case VarInt: + return regInt, nil + case VarStr: + return regStr, nil + default: + return 0, fmt.Errorf("bytecode typechecker: unknown var type %T", t) + } +} + +func metaRegType(t MetaType) (regType, error) { + switch t.(type) { + case MetaStr: + return regStr, nil + case MetaInt: + return regInt, nil + case MetaPortion: + return regPortion, nil + default: + return 0, fmt.Errorf("bytecode typechecker: unknown meta type %T", t) + } +} + +func unOpRegTypes(op UnKind) (dest, arg regType, err error) { + switch op.(type) { + case OpIntCopy: + return regInt, regInt, nil + case OpPortionCopy: + return regPortion, regPortion, nil + case OpStrCopy: + return regStr, regStr, nil + case OpBoolCopy: + return regBool, regBool, nil + case OpNegInt: + return regInt, regInt, nil + case OpIntToString: + return regStr, regInt, nil + case OpIsZero: + return regBool, regInt, nil + case OpNot: + return regBool, regBool, nil + case OpPortionToString: + return regStr, regPortion, nil + case OpIntToPortion: + return regPortion, regInt, nil + case OpPortionToInt: + return regInt, regPortion, nil + default: + return 0, 0, fmt.Errorf("bytecode typechecker: unknown unary op %T", op) + } +} + +func binOpRegTypes(op BinKind) (dest, left, right regType, err error) { + switch op.(type) { + case OpAddInt, OpSubInt: + return regInt, regInt, regInt, nil + case OpAddString: + return regStr, regStr, regStr, nil + case OpStrEq: + return regBool, regStr, regStr, nil + case OpLtInt, OpEqInt: + return regBool, regInt, regInt, nil + case OpLtPortion, OpEqPortion: + return regBool, regPortion, regPortion, nil + case OpAddPortion, OpSubPortion, OpMulPortion: + return regPortion, regPortion, regPortion, nil + case OpMakePortion: + return regPortion, regInt, regInt, nil + case OpMonetaryToString: + return regStr, regStr, regInt, nil + default: + return 0, 0, 0, fmt.Errorf("bytecode typechecker: unknown binary op %T", op) + } +} diff --git a/internal/ir/typecheck_test.go b/internal/ir/typecheck_test.go new file mode 100644 index 00000000..880fd114 --- /dev/null +++ b/internal/ir/typecheck_test.go @@ -0,0 +1,321 @@ +package ir + +import ( + "math/big" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestBytecodeTypecheck_Valid(t *testing.T) { + // $0 = 1; $1 = 2; $2 = $0 + $1 (all int) + instrs := []Instr{ + LoadInt{Dest: 0, Value: *big.NewInt(1)}, + LoadInt{Dest: 1, Value: *big.NewInt(2)}, + BinaryOp{Op: OpAddInt{}, Dest: 2, Left: 0, Right: 1}, + } + require.NoError(t, Typecheck(instrs)) +} + +func TestBytecodeTypecheck_UseBeforeWrite(t *testing.T) { + // reads $0 as int before it is ever written + instrs := []Instr{ + UnaryOp{Op: OpNegInt{}, Dest: 1, Arg: 0}, + } + require.ErrorContains(t, Typecheck(instrs), "read as int before being written") +} + +func TestBytecodeTypecheck_WrongType(t *testing.T) { + // $0 is a string, then used where an int is expected + instrs := []Instr{ + LoadStr{Dest: 0, Value: "USD/2"}, + UnaryOp{Op: OpNegInt{}, Dest: 1, Arg: 0}, + } + require.ErrorContains(t, Typecheck(instrs), "read as int but holds string") +} + +func TestBytecodeTypecheck_RedefinedWithDifferentType(t *testing.T) { + // $0 written as int, then overwritten as string + instrs := []Instr{ + LoadInt{Dest: 0, Value: *big.NewInt(1)}, + LoadStr{Dest: 0, Value: "x"}, + } + require.ErrorContains(t, Typecheck(instrs), "written as string but already holds int") +} + +func TestBytecodeTypecheck_ErrorLocatesTheInstruction(t *testing.T) { + instrs := []Instr{ + LoadStr{Dest: 0, Value: "src"}, + LoadStr{Dest: 1, Value: "dest"}, + CheckEnoughFunds{Got: 0, Needed: 1}, + } + err := Typecheck(instrs) + require.ErrorContains(t, err, "at instruction 2") + require.ErrorContains(t, err, "check_enough_funds($r0, $r1)") +} + +// Every register operand must be rejected when it names a register of the wrong +// bank. Each case is the prelude plus one instruction with exactly one bad operand. +func TestBytecodeTypecheck_OperandTypes(t *testing.T) { + intReg, strReg, portionReg := Reg(0), Reg(1), Reg(2) + prelude := []Instr{ + LoadInt{Dest: intReg, Value: *big.NewInt(1)}, + LoadStr{Dest: strReg, Value: "USD/2"}, + BinaryOp{Op: OpMakePortion{}, Dest: portionReg, Left: intReg, Right: intReg}, + } + + testCases := []struct { + name string + instr Instr + }{ + {"pull_account account", PullAccount{Dest: 9, Account: intReg}}, + {"pull_account cap", PullAccount{Dest: 9, Account: strReg, Cap: &strReg}}, + {"pull_account overdraft", PullAccount{Dest: 9, Account: strReg, Overdraft: &strReg}}, + {"pull_account color", PullAccount{Dest: 9, Account: strReg, Color: &intReg}}, + {"send_to_account account", SendToAccount{Account: &intReg}}, + {"send_to_account cap", SendToAccount{Account: &strReg, Cap: &strReg}}, + {"save account", Save{Account: intReg, Asset: strReg}}, + {"save asset", Save{Account: strReg, Asset: intReg}}, + {"save amount", Save{Account: strReg, Asset: strReg, Amount: &strReg}}, + {"check_enough_funds got", CheckEnoughFunds{Got: strReg, Needed: intReg}}, + {"check_enough_funds needed", CheckEnoughFunds{Got: intReg, Needed: strReg}}, + {"assert_leftover", AssertLeftover{Portion: intReg}}, + {"set_current_asset", SetCurrentAsset{Asset: intReg}}, + {"assert_same_asset left", AssertSameAsset{Left: intReg, Right: strReg}}, + {"assert_same_asset right", AssertSameAsset{Left: strReg, Right: intReg}}, + {"assert_valid_account", AssertValidAccount{Account: intReg}}, + {"assert_valid_color", AssertValidColor{Color: intReg}}, + {"assert_non_negative_balance balance", AssertNonNegativeBalance{Balance: strReg, Account: strReg}}, + {"assert_non_negative_balance account", AssertNonNegativeBalance{Balance: intReg, Account: intReg}}, + {"set_tx_meta key", SetTxMeta{Key: intReg, Value: strReg}}, + {"set_tx_meta value", SetTxMeta{Key: strReg, Value: intReg}}, + {"set_account_meta account", SetAccountMeta{Account: intReg, Key: strReg, Value: strReg}}, + {"set_account_meta key", SetAccountMeta{Account: strReg, Key: intReg, Value: strReg}}, + {"set_account_meta value", SetAccountMeta{Account: strReg, Key: strReg, Value: intReg}}, + {"meta account", MetaVar{Dest: 9, Account: intReg, Key: strReg, Typ: MetaStr{}}}, + {"meta key", MetaVar{Dest: 9, Account: strReg, Key: intReg, Typ: MetaStr{}}}, + {"meta_monetary account", MetaMonetary{DestAsset: 9, DestAmount: 10, Account: intReg, Key: strReg}}, + {"meta_monetary key", MetaMonetary{DestAsset: 9, DestAmount: 10, Account: strReg, Key: intReg}}, + {"meta_monetary dest asset", MetaMonetary{DestAsset: intReg, DestAmount: 10, Account: strReg, Key: strReg}}, + {"meta_monetary dest amount", MetaMonetary{DestAsset: 9, DestAmount: strReg, Account: strReg, Key: strReg}}, + {"balance account", FetchBalance{Dest: 9, Account: intReg, Asset: strReg}}, + {"balance asset", FetchBalance{Dest: 9, Account: strReg, Asset: intReg}}, + {"balance dest", FetchBalance{Dest: strReg, Account: strReg, Asset: strReg}}, + // a quantity is not a condition: that's the guarantee the bool bank buys + {"jmp_if_false cond", JmpIfFalse{Cond: intReg, Target: "end"}}, + {"jmp_if_true cond", JmpIfTrue{Cond: strReg, Target: "end"}}, + {"is_zero arg", UnaryOp{Op: OpIsZero{}, Dest: 9, Arg: strReg}}, + {"str_eq left", BinaryOp{Op: OpStrEq{}, Dest: 9, Left: intReg, Right: strReg}}, + // each comparison takes its own bank and yields a bool; not takes a bool + {"lt_int left", BinaryOp{Op: OpLtInt{}, Dest: 9, Left: strReg, Right: intReg}}, + {"lt_int right", BinaryOp{Op: OpLtInt{}, Dest: 9, Left: intReg, Right: portionReg}}, + {"eq_int left", BinaryOp{Op: OpEqInt{}, Dest: 9, Left: strReg, Right: intReg}}, + {"lt_portion left", BinaryOp{Op: OpLtPortion{}, Dest: 9, Left: intReg, Right: portionReg}}, + {"eq_portion right", BinaryOp{Op: OpEqPortion{}, Dest: 9, Left: portionReg, Right: intReg}}, + {"not arg", UnaryOp{Op: OpNot{}, Dest: 9, Arg: intReg}}, + {"add_portion left", BinaryOp{Op: OpAddPortion{}, Dest: 9, Left: intReg, Right: portionReg}}, + {"add_portion right", BinaryOp{Op: OpAddPortion{}, Dest: 9, Left: portionReg, Right: intReg}}, + // a copy never crosses banks + {"int_copy arg", UnaryOp{Op: OpIntCopy{}, Dest: 9, Arg: strReg}}, + {"portion_copy arg", UnaryOp{Op: OpPortionCopy{}, Dest: 9, Arg: intReg}}, + {"str_copy arg", UnaryOp{Op: OpStrCopy{}, Dest: 9, Arg: portionReg}}, + {"bool_copy arg", UnaryOp{Op: OpBoolCopy{}, Dest: 9, Arg: intReg}}, + {"unary arg", UnaryOp{Op: OpPortionToString{}, Dest: 9, Arg: intReg}}, + {"mul_portion left", BinaryOp{Op: OpMulPortion{}, Dest: 9, Left: intReg, Right: portionReg}}, + {"mul_portion right", BinaryOp{Op: OpMulPortion{}, Dest: 9, Left: portionReg, Right: intReg}}, + {"int_to_portion arg", UnaryOp{Op: OpIntToPortion{}, Dest: 9, Arg: portionReg}}, + {"portion_to_int arg", UnaryOp{Op: OpPortionToInt{}, Dest: 9, Arg: intReg}}, + {"binary left", BinaryOp{Op: OpAddString{}, Dest: 9, Left: intReg, Right: strReg}}, + {"binary right", BinaryOp{Op: OpAddString{}, Dest: 9, Left: strReg, Right: intReg}}, + {"monetary_to_string asset", BinaryOp{Op: OpMonetaryToString{}, Dest: 9, Left: intReg, Right: intReg}}, + {"monetary_to_string amount", BinaryOp{Op: OpMonetaryToString{}, Dest: 9, Left: strReg, Right: strReg}}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + require.Error(t, Typecheck(append(append([]Instr{}, prelude...), tc.instr))) + }) + } +} + +// The dest bank of these instructions comes from a type tag, not from an operand. +// Reading the dest back as an int is what tells the two apart. +func TestBytecodeTypecheck_TaggedDests(t *testing.T) { + str := Reg(0) + prelude := []Instr{LoadStr{Dest: str, Value: "k"}} + + testCases := []struct { + name string + instr Instr + destIsInt bool + }{ + {"load_var", LoadVar{Dest: 9, Typ: VarInt{}}, true}, + {"load_var", LoadVar{Dest: 9, Typ: VarStr{}}, false}, + {"meta", MetaVar{Dest: 9, Account: str, Key: str, Typ: MetaStr{}}, false}, + {"meta", MetaVar{Dest: 9, Account: str, Key: str, Typ: MetaInt{}}, true}, + {"meta", MetaVar{Dest: 9, Account: str, Key: str, Typ: MetaPortion{}}, false}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + // assert_non_negative_balance only accepts an int register + instrs := append(append([]Instr{}, prelude...), tc.instr, + AssertNonNegativeBalance{Balance: 9, Account: str}) + if tc.destIsInt { + require.NoError(t, Typecheck(instrs)) + } else { + require.Error(t, Typecheck(instrs)) + } + }) + } +} + +// A bool register is its own bank: nothing that takes an int accepts one, and it +// can't be rewritten as another type. +func TestBytecodeTypecheck_Bool(t *testing.T) { + t.Run("const_true and const_false define a bool", func(t *testing.T) { + require.NoError(t, Typecheck([]Instr{ + ConstBool{Dest: 0, Value: true}, + ConstBool{Dest: 1, Value: false}, + })) + }) + + t.Run("a bool is not an int", func(t *testing.T) { + err := Typecheck([]Instr{ + ConstBool{Dest: 0, Value: true}, + AssertNonNegativeBalance{Balance: 0, Account: 0}, + }) + require.ErrorContains(t, err, "read as int but holds bool") + }) + + t.Run("an int is not a bool", func(t *testing.T) { + err := Typecheck([]Instr{ + LoadInt{Dest: 0, Value: *big.NewInt(1)}, + ConstBool{Dest: 0, Value: true}, + }) + require.ErrorContains(t, err, "written as bool but already holds int") + }) + + t.Run("rewriting a bool with the same type is allowed", func(t *testing.T) { + require.NoError(t, Typecheck([]Instr{ + ConstBool{Dest: 0, Value: true}, + ConstBool{Dest: 0, Value: false}, + })) + }) + + // every comparison feeds a jump directly, and `not` composes with all of them + // — which is what makes the derived operators expressible without opcodes + t.Run("comparisons yield branchable bools", func(t *testing.T) { + // operand register per bank, so each op is fed its own type + prelude := []Instr{ + LoadInt{Dest: 0, Value: *big.NewInt(1)}, + BinaryOp{Op: OpMakePortion{}, Dest: 1, Left: 0, Right: 0}, + LoadStr{Dest: 2, Value: "x"}, + } + ops := map[BinKind]Reg{ + OpLtInt{}: 0, + OpEqInt{}: 0, + OpLtPortion{}: 1, + OpEqPortion{}: 1, + OpStrEq{}: 2, + } + for op, operand := range ops { + t.Run(op.String(), func(t *testing.T) { + require.NoError(t, Typecheck(append(append([]Instr{}, prelude...), + BinaryOp{Op: op, Dest: 9, Left: operand, Right: operand}, + UnaryOp{Op: OpNot{}, Dest: 10, Arg: 9}, + JmpIfTrue{Cond: 9, Target: "end"}, + JmpIfFalse{Cond: 10, Target: "end"}, + LabelMarker{Label: "end"}, + ))) + }) + } + }) + + // is_zero is in the comparison group too, and is the one unary member + t.Run("is_zero yields a branchable bool", func(t *testing.T) { + require.NoError(t, Typecheck([]Instr{ + LoadInt{Dest: 0, Value: *big.NewInt(1)}, + UnaryOp{Op: OpIsZero{}, Dest: 1, Arg: 0}, + JmpIfTrue{Cond: 1, Target: "end"}, + LabelMarker{Label: "end"}, + })) + }) + + t.Run("a comparison result is not an int", func(t *testing.T) { + err := Typecheck([]Instr{ + LoadInt{Dest: 0, Value: *big.NewInt(1)}, + BinaryOp{Op: OpEqInt{}, Dest: 1, Left: 0, Right: 0}, + AssertNonNegativeBalance{Balance: 1, Account: 0}, + }) + require.ErrorContains(t, err, "read as int but holds bool") + }) +} + +func TestBytecodeTypecheck_LabelMarker(t *testing.T) { + require.NoError(t, Typecheck([]Instr{LabelMarker{Label: "end"}})) + require.Empty(t, LabelMarker{Label: "end"}.dests()) + require.Empty(t, LabelMarker{Label: "end"}.sources()) +} + +// --- An unknown instruction or type tag is a bug in whatever built the stream: +// reported as an error, never panicked. + +type unknownInstr struct{} + +func (unknownInstr) dests() []Reg { return nil } +func (unknownInstr) sources() []Reg { return nil } +func (unknownInstr) assemble(*assembler) error { return nil } +func (unknownInstr) String() string { return "unknown_instr" } + +type unknownUnOp struct{} + +func (unknownUnOp) String() string { return "unknown_un_op" } +func (unknownUnOp) sig() unaryOpSig { return unaryOpSig{} } + +type unknownBinOp struct{} + +func (unknownBinOp) String() string { return "unknown_bin_op" } +func (unknownBinOp) sig() binaryOpSig { return binaryOpSig{} } + +type unknownVarType struct{} + +func (unknownVarType) String() string { return "unknown_var_type" } +func (unknownVarType) assembleLoad(*assembler, Reg, uint16) error { return nil } + +type unknownMetaType struct{} + +func (unknownMetaType) String() string { return "unknown_meta_type" } +func (unknownMetaType) assembleMeta(*assembler, Reg, Reg, Reg, *Reg) error { return nil } + +func TestBytecodeTypecheck_UnknownTags(t *testing.T) { + str := Reg(0) + prelude := []Instr{LoadStr{Dest: str, Value: "k"}} + + testCases := []struct { + name string + instr Instr + msg string + }{ + {"instruction", unknownInstr{}, "unhandled instruction"}, + {"unary op", UnaryOp{Op: unknownUnOp{}, Dest: 9, Arg: str}, "unknown unary op"}, + {"binary op", BinaryOp{Op: unknownBinOp{}, Dest: 9, Left: str, Right: str}, "unknown binary op"}, + {"var type", LoadVar{Dest: 9, Typ: unknownVarType{}}, "unknown var type"}, + {"meta type", MetaVar{Dest: 9, Account: str, Key: str, Typ: unknownMetaType{}}, "unknown meta type"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + require.ErrorContains(t, Typecheck(append(append([]Instr{}, prelude...), tc.instr)), tc.msg) + }) + } +} + +func TestRegTypeString(t *testing.T) { + require.Equal(t, "int", regInt.String()) + require.Equal(t, "string", regStr.String()) + require.Equal(t, "portion", regPortion.String()) + require.Equal(t, "bool", regBool.String()) + require.Equal(t, "?", regType(99).String()) + require.Equal(t, "?", regType(42).String()) +} diff --git a/internal/oracle/DIVERGENCES.md b/internal/oracle/DIVERGENCES.md index 621561b7..8607a922 100644 --- a/internal/oracle/DIVERGENCES.md +++ b/internal/oracle/DIVERGENCES.md @@ -376,10 +376,11 @@ without the harness. Only scripts both engines run to completion have their postings compared. Of the rest, some are rejected by the oracle at compile time (the generator's cleanup pass is best-effort; counted as `b-side compile rejection`), some are -tolerated under #1, and the remainder fail on both engines for the same -missing-funds reason (the sweep's reach table and tolerated counts print the -live numbers). A scenario block glued to a random program is often wasted this -way, which is why the generator has a scenario-only strategy. +tolerated under #1, some are numscript-only (below) and skip the oracle +altogether, and the remainder fail on both engines for the same missing-funds +reason (the sweep's reach table and tolerated counts print the live numbers). +A scenario block glued to a random program is often wasted this way, which is +why the generator has a scenario-only strategy. Since 2026-09-25 the generator also reaches allotment `remaining` clauses, portions written through vars (only ever inside a `remaining` block, summing @@ -396,13 +397,14 @@ Any other one-sided runtime failure is a mismatch. A quarter of generated scripts may draw numscript-only shapes — `oneof`, colored sources, division-expression portions (`$n/3`, any sign, the only route to negative or over-one portions), caps in a different asset than their -statement — and any script actually containing one skips the oracle leg +statement — and any script actually containing one skips the oracle legs (counted by name, `numscript-only script, oracle skipped`): the oracle has no -grammar for them. Until the compiler+VM leg lands, such a script is executed -by the interpreter alone — a panic still fails the fuzz target, and the reach -table counts the shapes — and the fixture corpus pins their semantics; the VM -leg is what will compare them engine-against-engine. Asset scaling and account -interpolation still have no generator coverage. +grammar for them. Those scripts are compared on `new vs vm` only, where +nothing is tolerated. Asset scaling and account interpolation still have no +generator coverage. Account interpolation is checked by the fixture corpus +(which both numscript engines run) and `TestNumscriptOnlyShapeAgreements`. +Asset scaling has no compiler lowering yet: its fixtures run on the +interpreter only, and `internal/compiler/scripts_test.go` skips them. --- diff --git a/internal/typecheck/typecheck.go b/internal/typecheck/typecheck.go new file mode 100644 index 00000000..41e36653 --- /dev/null +++ b/internal/typecheck/typecheck.go @@ -0,0 +1,449 @@ +// Package typecheck is the shared, side-effect-free type checker for numscript. +// It synthesizes the (base) type of every expression, resolves variable types +// from their declarations, and reports type/name/arity errors — fault-tolerantly +// (it collects all errors instead of bailing, using TypeAny to avoid cascades). +// +// It is the compiler's checker. The analysis module (LSP/CI) keeps its own +// unification-based checker: asset-identity inference, feature-version gating +// and lint-style warnings all stay there, and it does not consume this package. +package typecheck + +import ( + "fmt" + "slices" + "strings" + + "github.com/formancehq/numscript/internal/builtins" + "github.com/formancehq/numscript/internal/parser" +) + +type Type = string + +const ( + TypeNumber Type = "number" + TypeString Type = "string" + TypeAsset Type = "asset" + TypeMonetary Type = "monetary" + TypeAccount Type = "account" + TypePortion Type = "portion" + + // TypeAny is the type of an expression whose type couldn't be determined + // (e.g. an unbound variable or an unknown function). It's compatible with + // everything, so it suppresses cascading errors. + TypeAny Type = "any" +) + +// order mirrors analysis.AllowedTypes so the InvalidType message reads identically +var allowedTypes = []Type{TypeMonetary, TypeAccount, TypePortion, TypeAsset, TypeNumber, TypeString} + +func isTypeAllowed(t string) bool { return slices.Contains(allowedTypes, t) } + +// --- errors + +// Severity mirrors the LSP DiagnosticSeverity spec (and analysis.Severity, an +// alias of byte too) so an ErrorKind directly satisfies analysis.DiagnosticKind. +type Severity = byte + +const severityError Severity = 1 + +// ErrorKind is both a typecheck error and a renderable diagnostic (Message + +// Severity), so callers can push it as a diagnostic without a translation layer. +type ErrorKind interface { + errorKind() + Message() string + Severity() Severity +} + +type ( + TypeMismatch struct{ Expected, Got string } + UnboundVariable struct{ Name, Type string } + InvalidType struct{ Name string } + BadArity struct{ Expected, Actual int } + // UnknownFunction is either a truly-unknown name (WrongContext == "") or a + // known builtin used in the wrong context (WrongContext is the context it + // belongs to, e.g. "statement"). typecheck only ever emits the former. + UnknownFunction struct{ Name, WrongContext string } + DuplicateVariable struct{ Name string } +) + +func (TypeMismatch) errorKind() {} +func (UnboundVariable) errorKind() {} +func (InvalidType) errorKind() {} +func (BadArity) errorKind() {} +func (UnknownFunction) errorKind() {} +func (DuplicateVariable) errorKind() {} + +func (e TypeMismatch) Message() string { + return fmt.Sprintf("Type mismatch (expected '%s', got '%s' instead)", e.Expected, e.Got) +} + +func (e UnboundVariable) Message() string { + return fmt.Sprintf("The variable '$%s' was not declared", e.Name) +} + +func (e InvalidType) Message() string { + return fmt.Sprintf("'%s' is not a valid type. Allowed types are: %s", e.Name, strings.Join(allowedTypes, ", ")) +} + +func (e BadArity) Message() string { + return fmt.Sprintf("Wrong number of arguments (expected %d, got %d instead)", e.Expected, e.Actual) +} + +func (e UnknownFunction) Message() string { + if e.WrongContext != "" { + return fmt.Sprintf("You cannot use this function here (try to use it in a %s context)", e.WrongContext) + } + return fmt.Sprintf("The function '%s' does not exist", e.Name) +} + +func (e DuplicateVariable) Message() string { + return fmt.Sprintf("A variable with the name '$%s' was already declared", e.Name) +} + +func (TypeMismatch) Severity() Severity { return severityError } +func (UnboundVariable) Severity() Severity { return severityError } +func (InvalidType) Severity() Severity { return severityError } +func (BadArity) Severity() Severity { return severityError } +func (UnknownFunction) Severity() Severity { return severityError } +func (DuplicateVariable) Severity() Severity { return severityError } + +type Error struct { + Range parser.Range + Kind ErrorKind +} + +// --- builtin function signatures + +type fnSig struct { + params []Type + ret Type // "" for statement functions (no return) +} + +var builtinSigs = map[string]fnSig{ + builtins.SetTxMeta: {params: []Type{TypeString, TypeAny}}, + builtins.SetAccountMeta: {params: []Type{TypeAccount, TypeString, TypeAny}}, + builtins.Meta: {params: []Type{TypeAccount, TypeString}, ret: TypeAny}, + builtins.Balance: {params: []Type{TypeAccount, TypeAsset}, ret: TypeMonetary}, + builtins.Overdraft: {params: []Type{TypeAccount, TypeAsset}, ret: TypeMonetary}, + builtins.GetAsset: {params: []Type{TypeMonetary}, ret: TypeAsset}, + builtins.GetAmount: {params: []Type{TypeMonetary}, ret: TypeNumber}, + builtins.Scoped: {params: []Type{TypeAccount, TypeString}, ret: TypeAccount}, +} + +// --- Result / entrypoint + +type Result struct { + ExprTypes map[parser.ValueExpr]Type + VarTypes map[string]Type + Errors []Error +} + +func Check(program parser.Program) Result { + c := checker{ + exprTypes: map[parser.ValueExpr]Type{}, + varTypes: map[string]Type{}, + declared: map[string]struct{}{}, + } + c.checkProgram(program) + return Result{ExprTypes: c.exprTypes, VarTypes: c.varTypes, Errors: c.errors} +} + +type checker struct { + exprTypes map[parser.ValueExpr]Type + varTypes map[string]Type + declared map[string]struct{} + errors []Error +} + +func (c *checker) push(rng parser.Range, kind ErrorKind) { + c.errors = append(c.errors, Error{Range: rng, Kind: kind}) +} + +func (c *checker) checkProgram(program parser.Program) { + if program.Vars != nil { + for _, varDecl := range program.Vars.Declarations { + if varDecl.Type != nil && !isTypeAllowed(varDecl.Type.Name) { + c.push(varDecl.Type.Range, InvalidType{Name: varDecl.Type.Name}) + } + + if varDecl.Name != nil { + if _, dup := c.declared[varDecl.Name.Name]; dup { + c.push(varDecl.Name.Range, DuplicateVariable{Name: varDecl.Name.Name}) + } else { + c.declared[varDecl.Name.Name] = struct{}{} + if varDecl.Type != nil && isTypeAllowed(varDecl.Type.Name) { + c.varTypes[varDecl.Name.Name] = varDecl.Type.Name + } + } + } + + if varDecl.Origin != nil && varDecl.Type != nil { + c.checkExpr(*varDecl.Origin, varDecl.Type.Name) + } + } + } + + for _, statement := range program.Statements { + c.checkStatement(statement) + } +} + +func (c *checker) checkStatement(statement parser.Statement) { + switch statement := statement.(type) { + case *parser.SaveStatement: + c.checkSentValue(statement.SentValue) + c.checkExpr(statement.Account, TypeAccount) + + case *parser.SendStatement: + c.checkSentValue(statement.SentValue) + c.checkSource(statement.Source) + c.checkDestination(statement.Destination) + + case *parser.FnCall: + c.checkFnCallArity(statement) + } +} + +func (c *checker) checkSentValue(sentValue parser.SentValue) { + switch sentValue := sentValue.(type) { + case *parser.SentValueAll: + c.checkExpr(sentValue.Asset, TypeAsset) + case *parser.SentValueLiteral: + c.checkExpr(sentValue.Monetary, TypeMonetary) + } +} + +func (c *checker) checkSource(source parser.Source) { + if source == nil { + return + } + switch source := source.(type) { + case *parser.SourceAccount: + c.checkExpr(source.ValueExpr, TypeAccount) + c.checkExpr(source.Color, TypeString) + + case *parser.SourceOverdraft: + c.checkExpr(source.Address, TypeAccount) + c.checkExpr(source.Color, TypeString) + if source.Bounded != nil { + c.checkExpr(*source.Bounded, TypeMonetary) + } + + case *parser.SourceWithScaling: + c.checkExpr(source.Address, TypeAccount) + c.checkExpr(source.Through, TypeAccount) + + case *parser.SourceInorder: + for _, sub := range source.Sources { + c.checkSource(sub) + } + + case *parser.SourceOneof: + for _, sub := range source.Sources { + c.checkSource(sub) + } + + case *parser.SourceCapped: + c.checkExpr(source.Cap, TypeMonetary) + c.checkSource(source.From) + + case *parser.SourceAllotment: + for _, item := range source.Items { + if al, ok := item.Allotment.(*parser.ValueExprAllotment); ok { + c.checkExpr(al.Value, TypePortion) + } + c.checkSource(item.From) + } + } +} + +func (c *checker) checkDestination(destination parser.Destination) { + if destination == nil { + return + } + switch destination := destination.(type) { + case *parser.DestinationAccount: + c.checkExpr(destination.ValueExpr, TypeAccount) + + case *parser.DestinationInorder: + for _, clause := range destination.Clauses { + c.checkExpr(clause.Cap, TypeMonetary) + c.checkKeptOrDestination(clause.To) + } + c.checkKeptOrDestination(destination.Remaining) + + case *parser.DestinationOneof: + for _, clause := range destination.Clauses { + c.checkExpr(clause.Cap, TypeMonetary) + c.checkKeptOrDestination(clause.To) + } + c.checkKeptOrDestination(destination.Remaining) + + case *parser.DestinationAllotment: + for _, item := range destination.Items { + if al, ok := item.Allotment.(*parser.ValueExprAllotment); ok { + c.checkExpr(al.Value, TypePortion) + } + c.checkKeptOrDestination(item.To) + } + } +} + +func (c *checker) checkKeptOrDestination(keptOrDest parser.KeptOrDestination) { + if dest, ok := keptOrDest.(*parser.DestinationTo); ok { + c.checkDestination(dest.Destination) + } +} + +// checkExpr synthesizes lit's type, records it, and asserts it matches want. +func (c *checker) checkExpr(lit parser.ValueExpr, want Type) { + got := c.synthType(lit, want) + if want != TypeAny && got != TypeAny && want != got { + c.push(lit.GetRange(), TypeMismatch{Expected: want, Got: got}) + } +} + +// synthType synthesizes lit's type. hint is the type expected by the context; it +// is only used to annotate an unbound-variable error (matching the interpreter's +// diagnostic), never to influence the synthesized type. +func (c *checker) synthType(lit parser.ValueExpr, hint Type) Type { + if lit == nil { + return TypeAny + } + t := c.synthTypeInner(lit, hint) + c.exprTypes[lit] = t + return t +} + +func (c *checker) synthTypeInner(lit parser.ValueExpr, hint Type) Type { + switch lit := lit.(type) { + case *parser.Variable: + t, ok := c.varTypes[lit.Name] + if !ok { + if _, declared := c.declared[lit.Name]; !declared { + c.push(lit.Range, UnboundVariable{Name: lit.Name, Type: hint}) + } + return TypeAny + } + return t + + case *parser.MonetaryLiteral: + c.checkExpr(lit.Asset, TypeAsset) + c.checkExpr(lit.Amount, TypeNumber) + return TypeMonetary + + case *parser.BinaryInfix: + switch lit.Operator { + case parser.InfixOperatorPlus, parser.InfixOperatorMinus: + return c.checkInfixOverload(lit, []Type{TypeNumber, TypeMonetary}) + case parser.InfixOperatorDiv: + c.checkExpr(lit.Left, TypeNumber) + c.checkExpr(lit.Right, TypeNumber) + return TypePortion + default: + c.checkExpr(lit.Left, TypeAny) + c.checkExpr(lit.Right, TypeAny) + return TypeAny + } + + case *parser.Prefix: + switch lit.Operator { + case parser.PrefixOperatorMinus: + return c.checkHasOneOfTypes(lit.Expr, []Type{TypeNumber, TypeMonetary}) + default: + return TypeAny + } + + case *parser.AccountInterpLiteral: + for _, part := range lit.Parts { + if v, ok := part.(*parser.Variable); ok { + c.checkExpr(v, TypeAny) + } + } + return TypeAccount + + case *parser.PercentageLiteral: + return TypePortion + case *parser.AssetLiteral: + return TypeAsset + case *parser.NumberLiteral: + return TypeNumber + case *parser.StringLiteral: + return TypeString + + case *parser.FnCall: + return c.checkFnCall(lit) + + default: + return TypeAny + } +} + +func (c *checker) checkInfixOverload(bin *parser.BinaryInfix, allowed []Type) Type { + leftType := c.synthType(bin.Left, allowed[0]) + if leftType == TypeAny || slices.Contains(allowed, leftType) { + c.checkExpr(bin.Right, leftType) + return leftType + } + c.push(bin.Left.GetRange(), TypeMismatch{Expected: strings.Join(allowed, "|"), Got: leftType}) + c.checkExpr(bin.Right, TypeAny) + return TypeAny +} + +func (c *checker) checkHasOneOfTypes(expr parser.ValueExpr, allowed []Type) Type { + exprType := c.synthType(expr, allowed[0]) + if exprType == TypeAny || slices.Contains(allowed, exprType) { + return exprType + } + c.push(expr.GetRange(), TypeMismatch{Expected: strings.Join(allowed, "|"), Got: exprType}) + return TypeAny +} + +func (c *checker) checkFnCall(fnCall *parser.FnCall) Type { + ret := TypeAny + if sig, ok := builtinSigs[fnCall.Caller.Name]; ok { + ret = sig.ret + if ret == "" { + ret = TypeAny + } + } + c.checkFnCallArity(fnCall) + return ret +} + +func (c *checker) checkFnCallArity(fnCall *parser.FnCall) { + var validArgs []parser.ValueExpr + for _, arg := range fnCall.Args { + if arg != nil { + validArgs = append(validArgs, arg) + } + } + + sig, resolved := builtinSigs[fnCall.Caller.Name] + if !resolved { + for _, arg := range validArgs { + c.checkExpr(arg, TypeAny) + } + c.push(fnCall.Caller.Range, UnknownFunction{Name: fnCall.Caller.Name}) + return + } + + expected := len(sig.params) + actual := len(validArgs) + if actual < expected { + c.push(fnCall.Range, BadArity{Expected: expected, Actual: actual}) + } else if actual > expected { + first := validArgs[expected] + last := validArgs[len(validArgs)-1] + c.push(parser.Range{Start: first.GetRange().Start, End: last.GetRange().End}, + BadArity{Expected: expected, Actual: actual}) + } + + for i, arg := range validArgs { + want := TypeAny + if i < len(sig.params) { + want = sig.params[i] + } + c.checkExpr(arg, want) + } +} diff --git a/internal/typecheck/typecheck_test.go b/internal/typecheck/typecheck_test.go new file mode 100644 index 00000000..6a47871a --- /dev/null +++ b/internal/typecheck/typecheck_test.go @@ -0,0 +1,97 @@ +package typecheck_test + +import ( + "testing" + + "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/typecheck" + + "github.com/stretchr/testify/require" +) + +func check(t *testing.T, src string) typecheck.Result { + t.Helper() + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + return typecheck.Check(parsed.Value) +} + +func kinds(res typecheck.Result) []typecheck.ErrorKind { + out := make([]typecheck.ErrorKind, len(res.Errors)) + for i, e := range res.Errors { + out[i] = e.Kind + } + return out +} + +func TestValidProgram(t *testing.T) { + res := check(t, ` + vars { account $acc = @src } + send [USD/2 10] (source = $acc destination = @dest) + `) + require.Empty(t, res.Errors) + require.Equal(t, typecheck.TypeAccount, res.VarTypes["acc"]) +} + +func TestInvalidType(t *testing.T) { + res := check(t, `vars { invalid $x }`) + require.Equal(t, []typecheck.ErrorKind{typecheck.InvalidType{Name: "invalid"}}, kinds(res)) +} + +func TestDuplicateVariable(t *testing.T) { + res := check(t, `vars { account $x account $x }`) + require.Equal(t, []typecheck.ErrorKind{typecheck.DuplicateVariable{Name: "x"}}, kinds(res)) +} + +func TestUnboundVariable(t *testing.T) { + res := check(t, `send [C 10] (source = $nope destination = @d)`) + require.Equal(t, []typecheck.ErrorKind{typecheck.UnboundVariable{Name: "nope", Type: typecheck.TypeAccount}}, kinds(res)) +} + +func TestTypeMismatch(t *testing.T) { + // a string var used where an account is expected + res := check(t, `vars { string $s } send [C 10] (source = $s destination = @d)`) + require.Equal(t, []typecheck.ErrorKind{ + typecheck.TypeMismatch{Expected: typecheck.TypeAccount, Got: typecheck.TypeString}, + }, kinds(res)) +} + +func TestUnknownFunction(t *testing.T) { + res := check(t, `vars { number $n = nope() }`) + require.Equal(t, []typecheck.ErrorKind{typecheck.UnknownFunction{Name: "nope"}}, kinds(res)) +} + +func TestBadArity(t *testing.T) { + res := check(t, `vars { monetary $m = balance(@a) }`) + require.Equal(t, []typecheck.ErrorKind{typecheck.BadArity{Expected: 2, Actual: 1}}, kinds(res)) +} + +func TestExprTypes(t *testing.T) { + res := check(t, `send [USD/2 10] (source = @a destination = @b)`) + // the monetary literal is typed + send := res // just assert no errors + monetary present via a scan + require.Empty(t, send.Errors) + found := false + for _, ty := range res.ExprTypes { + if ty == typecheck.TypeMonetary { + found = true + } + } + require.True(t, found, "expected a monetary-typed expr") +} + +func TestInfixInvalidLeftStillChecksRight(t *testing.T) { + res := check(t, `set_tx_meta("k", "text" + $missing)`) + require.Equal(t, []typecheck.ErrorKind{ + typecheck.TypeMismatch{Expected: typecheck.TypeNumber + "|" + typecheck.TypeMonetary, Got: typecheck.TypeString}, + typecheck.UnboundVariable{Name: "missing", Type: typecheck.TypeAny}, + }, kinds(res)) +} + +func TestBadArityStillChecksSurplusArgs(t *testing.T) { + res := check(t, `vars { monetary $m = balance(@a, USD/2, $missing) }`) + require.Equal(t, []typecheck.ErrorKind{ + typecheck.BadArity{Expected: 2, Actual: 3}, + typecheck.UnboundVariable{Name: "missing", Type: typecheck.TypeAny}, + }, kinds(res)) +} diff --git a/internal/vm/execution_err.go b/internal/vm/execution_err.go new file mode 100644 index 00000000..0f575fce --- /dev/null +++ b/internal/vm/execution_err.go @@ -0,0 +1,210 @@ +package vm + +import ( + "fmt" + "math/big" + + "github.com/formancehq/numscript/internal/funds" +) + +type ( + ExecutionError interface { + error + execErr() + } + + MissingFundsError struct { + Asset string + Needed *big.Int + Got *big.Int + } + + AssetMismatchError struct { + Expected string + Got string + } + + InvalidUncappedSource struct { + Account string + } + + InvalidAllotmentSum struct { + ActualSum big.Rat + } + + MetadataNotFoundError struct { + Account string + Key string + } + + BadMetaValueError struct { + Account string + Key string + Raw string + } + + InvalidAccountName struct { + Name string + } + + InvalidColor struct { + Color string + } + + InvalidScope struct { + Scope string + } + + CannotCastScopedAccountToString struct { + Account string + Scope string + } + + NegativeBalanceError struct { + Account string + Amount big.Int + } + + // NegativeAmountError is a sent/saved amount that evaluated to negative — + // unlike NegativeBalanceError, it is not tied to any account. + NegativeAmountError struct { + Amount big.Int + } + + DivideByZeroError struct { + Numerator big.Int + } + + NegativePortionError struct { + Portion big.Rat + } + + // InternalError signals a malformed program the VM cannot execute: a bug in + // whatever produced the bytecode, never a user-script error. Returned rather + // than panicked so the VM never crashes its host. + // + // The mark violations land here rather than getting user-facing error types, + // since both are properties a verifier can decide from the instruction stream + // alone and neither is a legitimate outcome of a well-formed program. + InternalError struct { + Err error + } + + // StoreError wraps an error returned by the host Store (balance or metadata + // fetch): neither a script error nor a bytecode bug, so the wrapped error is + // preserved for the host to inspect. + StoreError struct { + Wrapped error + } +) + +func (e MissingFundsError) Error() string { + return fmt.Sprintf("missing funds for asset %s: needed %s, got %s", e.Asset, e.Needed, e.Got) +} + +func (e AssetMismatchError) Error() string { + return fmt.Sprintf("asset mismatch: expected %s, got %s", e.Expected, e.Got) +} + +func (e InvalidUncappedSource) Error() string { + return fmt.Sprintf("unbounded source is not allowed here: @%s", e.Account) +} + +func (e InternalError) Error() string { + return "internal error: " + e.Err.Error() +} + +func (e InternalError) Unwrap() error { return e.Err } + +// InvalidPostingError is wrapped in an InternalError when a posting fails +// funds.ValidatePosting. +type InvalidPostingError struct { + Posting funds.Posting +} + +func (e InvalidPostingError) Error() string { + return fmt.Sprintf("produced a posting with invalid values: %+v", e.Posting) +} + +func (e DivideByZeroError) Error() string { + return fmt.Sprintf("cannot divide by zero (in %s/0)", e.Numerator.String()) +} + +func (e InvalidAccountName) Error() string { + return fmt.Sprintf("invalid account name: %q", e.Name) +} + +func (e InvalidColor) Error() string { + return fmt.Sprintf("invalid color name: %q", e.Color) +} + +func (e InvalidScope) Error() string { + return fmt.Sprintf("invalid scope name: %q", e.Scope) +} + +func (e CannotCastScopedAccountToString) Error() string { + return fmt.Sprintf("cannot cast a scoped account to string (account %q has scope %q)", e.Account, e.Scope) +} + +func (e NegativeBalanceError) Error() string { + return fmt.Sprintf("cannot fetch negative balance from account @%s", e.Account) +} + +func (e NegativeAmountError) Error() string { + return fmt.Sprintf("cannot send negative amount: %s", e.Amount.String()) +} + +func (e InvalidAllotmentSum) Error() string { + return fmt.Sprintf("invalid allotment: portions must sum to 1, got %s", e.ActualSum.String()) +} + +func (e NegativePortionError) Error() string { + return fmt.Sprintf("invalid allotment: portions cannot be negative, got %s", e.Portion.String()) +} + +func (e MetadataNotFoundError) Error() string { + return fmt.Sprintf("metadata not found: %s[%q]", e.Account, e.Key) +} + +func (e BadMetaValueError) Error() string { + return fmt.Sprintf("invalid metadata value for %s[%q]: %q", e.Account, e.Key, e.Raw) +} + +func (e StoreError) Error() string { return "store error: " + e.Wrapped.Error() } +func (e StoreError) Unwrap() error { return e.Wrapped } + +func (MissingFundsError) execErr() {} +func (AssetMismatchError) execErr() {} +func (InvalidUncappedSource) execErr() {} +func (InvalidAllotmentSum) execErr() {} +func (NegativePortionError) execErr() {} +func (MetadataNotFoundError) execErr() {} +func (BadMetaValueError) execErr() {} +func (InvalidAccountName) execErr() {} +func (InvalidColor) execErr() {} +func (InvalidScope) execErr() {} +func (CannotCastScopedAccountToString) execErr() {} +func (NegativeBalanceError) execErr() {} +func (NegativeAmountError) execErr() {} +func (DivideByZeroError) execErr() {} +func (InternalError) execErr() {} +func (StoreError) execErr() {} + +var ( + _ ExecutionError = (*MissingFundsError)(nil) + _ ExecutionError = (*AssetMismatchError)(nil) + _ ExecutionError = (*InvalidUncappedSource)(nil) + _ ExecutionError = (*InvalidAllotmentSum)(nil) + _ ExecutionError = (*NegativePortionError)(nil) + _ ExecutionError = (*MetadataNotFoundError)(nil) + _ ExecutionError = (*BadMetaValueError)(nil) + _ ExecutionError = (*InvalidAccountName)(nil) + _ ExecutionError = (*InvalidColor)(nil) + _ ExecutionError = (*InvalidScope)(nil) + _ ExecutionError = (*CannotCastScopedAccountToString)(nil) + _ ExecutionError = (*NegativeBalanceError)(nil) + _ ExecutionError = (*NegativeAmountError)(nil) + _ ExecutionError = (*DivideByZeroError)(nil) + _ ExecutionError = (*InternalError)(nil) + _ ExecutionError = (*StoreError)(nil) +) diff --git a/internal/vm/export_test.go b/internal/vm/export_test.go new file mode 100644 index 00000000..77388479 --- /dev/null +++ b/internal/vm/export_test.go @@ -0,0 +1,3 @@ +package vm + +var ErrSendWhileMarkOpen = errSendWhileMarkOpen diff --git a/internal/vm/fuzz_test.go b/internal/vm/fuzz_test.go new file mode 100644 index 00000000..647b894c --- /dev/null +++ b/internal/vm/fuzz_test.go @@ -0,0 +1,58 @@ +package vm + +import ( + "context" + "math/big" + "testing" +) + +// FuzzExec reads arbitrary bytes as an instruction stream and asserts that +// anything the verifier accepts, the VM can run without panicking. +// +// This is the verifier's contract stated as a test. It says nothing about +// programs the verifier rejects: Exec is entitled to crash on those, which is +// why it is Exec's caller that has to decide whether a program is trusted. +// +// Note VerifyWithVars rather than Verify — a program that loads a variable is +// only safe against the vars it will actually be given, and Verify alone does +// not look at them. +func FuzzExec(f *testing.F) { + f.Add([]byte{}) + f.Add([]byte{byte(Op_LoadInt), 0, 0, 0}) + f.Add([]byte{byte(Op_PullAccount), 0, 0, nilReg}) // truncated: no ext word + f.Add([]byte{byte(Op_LoadStr), 0, 0, 0, byte(Op_Jmp), 0, 1, 0}) // jump past the end + f.Add([]byte{byte(Op_ConstTrue), 0, 0, 0, byte(Op_JmpIfTrue), 0, 0, 0}) + f.Add([]byte{byte(Op_MarkPush), 0, 0, 0, byte(Op_MarkEnd), 1, 0, 0}) + + stringsPool := []string{"world", "dest", "USD/2"} + intsPool := []big.Int{*big.NewInt(0), *big.NewInt(7)} + vars := &Vars{ + StringsPool: []string{"var0"}, + IntsPool: []big.Int{*big.NewInt(3)}, + } + + f.Fuzz(func(t *testing.T, data []byte) { + instrs := make([]Instruction, len(data)/4) + for i := range instrs { + off := i * 4 + instrs[i] = Instruction{data[off], data[off+1], data[off+2], data[off+3]} + } + + prog := fullBanks(Program{ + Instructions: instrs, + StringsPool: stringsPool, + IntsPool: intsPool, + }) + + if _, err := VerifyWithVars(prog, vars); err != nil { + return + } + + defer func() { + if r := recover(); r != nil { + t.Fatalf("verified program panicked in Exec: %v", r) + } + }() + _, _ = Exec(context.Background(), NewVm(prog), vars, mockStore{}) + }) +} diff --git a/internal/vm/instruction.go b/internal/vm/instruction.go new file mode 100644 index 00000000..c46f280c --- /dev/null +++ b/internal/vm/instruction.go @@ -0,0 +1,264 @@ +package vm + +import "encoding/binary" + +type Instruction struct { + Opcode byte + A byte + B byte + C byte +} + +// Little endian view of the b and c fields +func (i Instruction) GetBC() uint16 { + return uint16(i.B) | uint16(i.C)<<8 +} + +func NewBC( + opcode Opcode, + a byte, + bc uint16, +) Instruction { + var bcBytes [2]byte + binary.LittleEndian.PutUint16(bcBytes[:], bc) + + return Instruction{ + Opcode: byte(opcode), + A: a, + B: bcBytes[0], + C: bcBytes[1], + } +} + +type Opcode byte + +// Opcodes are grouped by category with gaps, so new instructions can be added to +// a category without renumbering. See instruction-encoding.md. +const ( + // --- state & assertions (0x00) --- + Op_SetCurrentAsset Opcode = 0x00 + + Op_AssertSameAsset Opcode = 0x01 + + // errors if the account name in str reg A is not well-formed + Op_AssertValidAccount Opcode = 0x02 + + // errors (NegativeBalanceError) if the amount in int reg A is negative; + // B = account str reg (for the error) + Op_AssertNonNegativeBalance Opcode = 0x03 + + // checks the allotment leftover portion in reg A: errors if negative (portions + // summing to > 1), and — when B == 1 (no `remaining` clause) — if non-zero + Op_AssertLeftover Opcode = 0x04 + + Op_CheckEnoughFunds Opcode = 0x05 + + // errors if the color in str reg A is not well-formed + Op_AssertValidColor Opcode = 0x06 + + // errors (NegativeAmountError) if the amount in int reg A is negative — + // a sent/saved amount, not tied to any account (unlike + // Op_AssertNonNegativeBalance, which is a balance() read) + Op_AssertNonNegativeAmount Opcode = 0x07 + + // errors if the scope in str reg A is not well-formed + Op_AssertValidScope Opcode = 0x08 + + // errors (NegativePortionError) if the portion in reg A is negative — an + // allotment clause portion + Op_AssertNonNegativePortion Opcode = 0x09 + + // errors (CannotCastScopedAccountToString) if the scope in str reg A is not + // empty; B = account str reg (for the error) + Op_AssertUnscoped Opcode = 0x0A + + // --- constants & variables (0x10) --- + // may split into one opcode per expr_typ later + Op_LoadInt Opcode = 0x10 // LoadConst (`Int) -> b_c = const-pool index + Op_LoadStr Opcode = 0x11 // LoadConst (`String) -> b_c = const-pool index + + Op_LoadVarInt Opcode = 0x12 // b_c = int-var index + Op_LoadVarStr Opcode = 0x13 // b_c = string-var index + + // 0x14 Op_LoadIntImmediate: inline i16 literal in b_c. NOT IMPLEMENTED (reserved) + + // A = dest (bool reg); one opcode per constant, so there is no operand to decode + Op_ConstTrue Opcode = 0x15 + Op_ConstFalse Opcode = 0x16 + + // --- metadata (0x20) --- + // A = key (str reg), B = value (str reg) + Op_SetTxMeta Opcode = 0x20 + + // A = account (str reg), B = key (str reg), C = value (str reg); ext.A = + // scope (str reg, 0xFF = unscoped) + Op_SetAccountMeta Opcode = 0x21 + + // meta(account, key) read, dispatched on the target type. + // A = dest, B = account (str reg), C = key (str reg); ext.A = scope (str reg, + // 0xFF = unscoped) + Op_MetaStr Opcode = 0x22 + Op_MetaInt Opcode = 0x23 + Op_MetaPortion Opcode = 0x24 + + // as above, but a monetary needs two destinations, so the amount's goes in an + // ext word: A = dest asset (str reg), ext.A = dest amount (int reg), ext.B = + // scope (str reg, 0xFF = unscoped) + Op_MetaMonetary Opcode = 0x25 + + // --- arithmetic & constructors (0x30) --- + Op_AddInt Opcode = 0x30 + Op_SubInt Opcode = 0x31 + // 0x32 was Op_MinInt: a min is a comparison and a copy, so it is Op_LtInt + // plus a branch. Reserved, do not reuse. + Op_SubPortion Opcode = 0x33 + Op_MkPortion Opcode = 0x34 + // 0x35 was Op_MkMonetary: a monetary is a (str asset, int amount) register + // pair, so there is nothing to construct. Reserved, do not reuse. + Op_AddString Opcode = 0x36 + + // 0x37 was Op_StrEq: moved to the comparison group, now 0x62. Reserved, do + // not reuse. + + // not adjacent to Op_SubPortion (0x33) because 0x32 is burned and 0x34..0x37 + // are taken + Op_AddPortion Opcode = 0x38 + + // an allotment share is a mul plus a floor (Op_PortionToInt) + Op_MulPortion Opcode = 0x39 + + // 0x3A..0x3F reserved + + // --- unary & conversions (0x40) --- + // 0x40 was Op_GetAmount and 0x41 was Op_GetAsset: projecting a monetary is now + // just naming one of its two registers. Reserved, do not reuse. + // + // One copy per register bank: A = dest, B = src, both in that bank. There is no + // monetary copy — a monetary is a (str asset, int amount) pair, so copy the two + // halves. The family is split across 0x42..0x43 and 0x4A..0x4B because + // 0x44..0x49 were already spoken for. + Op_IntCopy Opcode = 0x42 + Op_PortionCopy Opcode = 0x43 + + Op_NegInt Opcode = 0x44 + Op_IntToString Opcode = 0x45 + Op_PortionToString Opcode = 0x46 + + // A = dest (str reg), B = asset (str reg), C = amount (int reg) + Op_MonetaryToString Opcode = 0x47 + + // 0x48 was Op_IsZero: moved to the comparison group, now 0x63. Reserved, do + // not reuse. + + // 0x49 was Op_Not: moved to the bool-ops group, now 0x70. Reserved, do not + // reuse. + + // the other two bank copies; see Op_IntCopy above + Op_StrCopy Opcode = 0x4A + Op_BoolCopy Opcode = 0x4B + + // Op_IntToPortion is exact; Op_PortionToInt floors. + Op_IntToPortion Opcode = 0x4C + Op_PortionToInt Opcode = 0x4D + + // 0x4E..0x4F reserved + + // --- funds & postings (0x50) --- + + // The most general form: account,cap,overdraft,color,scope + // The 0xFF special register means NULL for cap,overdraft,color and scope + // ext.A = overdraft (int reg), ext.B = color (str reg), ext.C = scope (str reg) + Op_PullAccount Opcode = 0x50 + + // account?, cap?, scope? + Op_SendToAccount Opcode = 0x51 + + // save: reduce balance of account A for asset B by amount C (C == nilReg => + // save all), floored at 0; ext.A = scope (str reg, 0xFF = unscoped) + Op_Save Opcode = 0x52 + + // 0x53 was Op_MkAllotment: an allotment share is now built out of pure ops + // (Op_IntToPortion, Op_MulPortion, Op_PortionToInt plus the leftover fixup), + // so there is no variadic domain instruction. Reserved, do not reuse. + + // reads the account balance from the run-state; A = dest, B = account, C = + // asset; ext.A = scope (str reg, 0xFF = unscoped) + Op_Balance Opcode = 0x54 + + // --- marks (oneof backtracking) --- + // + // 0x55 was Op_Snapshot and 0x56 was Op_Restore: the source-queue mark used to + // travel through an int register, which let any int be passed to a restore. The + // pair below takes no register; the mark lives on a LIFO owned by the run-state. + // + // There is no "rewind but keep the mark" opcode: a retry is Op_MarkEnd with the + // rewind flag followed by a fresh Op_MarkPush, so pushes and ends match strictly + // and mark depth is a function of position in the instruction stream. Verify + // proves pushes and ends balance and that no Op_SendToAccount / + // Op_SetCurrentAsset / Op_Save sits inside a region; Exec enforces the same at + // execution time, since it does not require a verified program. + + // opens a region at the current source-queue depth and posting count + Op_MarkPush Opcode = 0x55 + + // A = rewind flag. Always pops the innermost mark; A == 1 additionally repays + // everything pulled and reverses everything posted since the matching + // Op_MarkPush, while A == 0 commits it. + Op_MarkEnd Opcode = 0x56 + + // reserved (0x57..0x5F) for PullAccount specializations, e.g.: + // // cap=None, overdraft=BoundedZero + // Op_PullAccountBoundedZero + // // cap=None, overdraft=Bounded r + // Op_PullAccountOverdraft + // // cap=Some, overdraft=BoundedZero + // Op_PullAccountCap + // // cap=Some, overdraft=Unbounded + // Op_PullAccountUnboundedOverdraft + // + // This block used to run to 0x8F; the comparison and bool-ops groups below took + // 0x60..0x7F out of it. + + // --- comparisons (0x60) --- + // A = dest (bool reg) for all of them; the operand banks are what the opcode + // implies. Op_IsZero is unary and the rest binary, but they are one group + // because they are the whole set of bool *producers*. + // + // Only `<` and `==` exist, per type. The other surface operators are normalised + // by the front end: + // + // a < b -> Lt(a, b) + // a > b -> Lt(b, a) operands swapped + // a <= b -> Not(Lt(b, a)) + // a >= b -> Not(Lt(a, b)) + // a == b -> Eq(a, b) + // a != b -> Not(Eq(a, b)) + Op_LtInt Opcode = 0x60 + Op_EqInt Opcode = 0x61 + Op_StrEq Opcode = 0x62 // was 0x37 + Op_IsZero Opcode = 0x63 // was 0x48 + Op_LtPortion Opcode = 0x64 + Op_EqPortion Opcode = 0x65 + + // reserved (0x66..0x6F) for `<` and `==` on types that don't exist yet. Str gets + // equality only, never ordering. Bool equality and structural comparison of + // tuples/arrays are front-end expansions rather than opcodes. + + // --- bool ops (0x70) --- + // A = dest (bool reg), B = src (bool reg). + Op_Not Opcode = 0x70 // was 0x49 + + // reserved (0x71..0x7F) for and/or; both are expressible as branches, so + // neither is needed for completeness + + // --- control flow (0x90) --- + // A = cond (bool reg); b_c = unsigned forward delta, added to the pc of the + // next instruction. A quantity is not a condition: project it with Op_IsZero. + Op_JmpIfFalse Opcode = 0x90 + // unconditional; b_c = unsigned forward delta, as above + Op_Jmp Opcode = 0x91 + // the dual of Op_JmpIfFalse, so either edge of a bool can be the branch without + // a negation instruction + Op_JmpIfTrue Opcode = 0x92 + // Label emits no instruction; it only feeds the symbol table at assemble time +) diff --git a/internal/vm/ir_test.go b/internal/vm/ir_test.go new file mode 100644 index 00000000..80c36a5f --- /dev/null +++ b/internal/vm/ir_test.go @@ -0,0 +1,1371 @@ +package vm_test + +import ( + "context" + "errors" + "fmt" + "math/big" + "testing" + + "github.com/formancehq/numscript/internal/funds" + "github.com/formancehq/numscript/internal/ir" + "github.com/formancehq/numscript/internal/vm" + "github.com/stretchr/testify/require" +) + +// These tests drive the VM from the IR textual format, without the compiler, so +// they can also cover instruction sequences the compiler doesn't emit. + +// irStore is a vm.Store backed by plain maps. A non-nil err fails every lookup. +type irStore struct { + balances map[funds.PairKey]*big.Int + metadata map[string]map[string]string + err error +} + +func (s irStore) GetBalance(_ context.Context, account, scope, asset, color string) (*big.Int, error) { + if s.err != nil { + return nil, s.err + } + if v, ok := s.balances[funds.PairKey{Account: account, Scope: scope, Asset: asset, Color: color}]; ok { + return new(big.Int).Set(v), nil + } + return new(big.Int), nil +} + +func (s irStore) GetMetadata(_ context.Context, account, scope, key string) (string, bool, error) { + if s.err != nil { + return "", false, s.err + } + v, ok := s.metadata[account][key] + return v, ok, nil +} + +func meta(rows map[string]map[string]string) irStore { + return irStore{metadata: rows} +} + +func balances(pairs map[string]int64) irStore { + b := map[funds.PairKey]*big.Int{} + for account, amount := range pairs { + b[funds.PairKey{Account: account, Asset: "USD/2"}] = big.NewInt(amount) + } + return irStore{balances: b} +} + +// allot2IR is the sequence the compiler emits to split an amount two ways: +// floor each share, then hand the flooring leftover to the earliest. There is +// no allotment instruction — the split is built out of pure ops — so the three +// tests that need one share this rather than spelling it out each time. +// +// Only one fixup block: flooring loses under a unit per share, so with the two +// portions summing to 1 the shortfall is at most 1 and the second share never +// receives it. +func allot2IR(amount, portion1, portion2, share1, share2 string) string { + return fmt.Sprintf(` + $allot_amt = int_to_portion($%[1]s) + $allot_prod = mul_portion($%[2]s, $allot_amt) + $%[4]s = portion_to_int($allot_prod) + $allot_total = int_copy($%[4]s) + $allot_prod = mul_portion($%[3]s, $allot_amt) + $%[5]s = portion_to_int($allot_prod) + $allot_total = add_int($allot_total, $%[5]s) + $allot_one = 1 + $allot_short = lt_int($allot_total, $%[1]s) + jmp_if_false($allot_short, #allot_end) + $%[4]s = add_int($%[4]s, $allot_one) + $allot_total = add_int($allot_total, $allot_one) +#allot_end +`, amount, portion1, portion2, share1, share2) +} + +// assembleUnverifiedIR is assembleIR without the verifier, for the malformed +// programs that exercise Exec's own guards. +func assembleUnverifiedIR(t *testing.T, src string) vm.Program { + t.Helper() + + instrs, errs := ir.Parse(src) + require.Empty(t, errs, "IR errors: %v", errs) + require.NoError(t, ir.Typecheck(instrs)) + + program, err := ir.Assemble(instrs) + require.NoError(t, err) + return program +} + +// assembleIR turns an IR text into a runnable program, failing the test on any +// error the format's own layers report. +func assembleIR(t *testing.T, src string) vm.Program { + t.Helper() + + program := assembleUnverifiedIR(t, src) + + // every assembled program in this file goes through the verifier, so these + // tests double as its corpus for sequences the compiler never emits + require.NoError(t, vm.Verify(program)) + + return program +} + +// runIR assembles and runs an IR text, requiring it to succeed. +func runIR(t *testing.T, src string, store irStore, vars *vm.Vars) funds.ExecutionResult { + t.Helper() + + res, execErr := vm.Exec(context.Background(), vm.NewVm(assembleIR(t, src)), vars, store) + require.Nil(t, execErr, "unexpected execution error: %v", execErr) + return res +} + +// runUnverifiedIRExpectingError is runIRExpectingError for a program the +// verifier rejects with an error containing verifyErr. +func runUnverifiedIRExpectingError(t *testing.T, src string, verifyErr string, store irStore, vars *vm.Vars) vm.ExecutionError { + t.Helper() + + program := assembleUnverifiedIR(t, src) + require.ErrorContains(t, vm.Verify(program), verifyErr) + + _, execErr := vm.Exec(context.Background(), vm.NewVm(program), vars, store) + require.NotNil(t, execErr, "expected an execution error") + return execErr +} + +// runIRExpectingError is runIR for the cases that must fail at run time. +func runIRExpectingError(t *testing.T, src string, store irStore, vars *vm.Vars) vm.ExecutionError { + t.Helper() + + _, execErr := vm.Exec(context.Background(), vm.NewVm(assembleIR(t, src)), vars, store) + require.NotNil(t, execErr, "expected an execution error") + return execErr +} + +func requirePostings(t *testing.T, want, got []funds.Posting) { + t.Helper() + + require.Len(t, got, len(want)) + for i := range want { + w, g := want[i], got[i] + require.Equal(t, w.Source, g.Source, "posting[%d].Source", i) + require.Equal(t, w.Destination, g.Destination, "posting[%d].Destination", i) + require.Equal(t, w.Asset, g.Asset, "posting[%d].Asset", i) + require.Equal(t, w.Color, g.Color, "posting[%d].Color", i) + require.Zero(t, g.Amount.Cmp(w.Amount), "posting[%d].Amount: got %s, want %s", i, g.Amount, w.Amount) + } +} + +func posting(source, destination string, amount int64) funds.Posting { + return funds.Posting{Source: source, Destination: destination, Asset: "USD/2", Amount: big.NewInt(amount)} +} + +func TestIRSend(t *testing.T) { + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 10 + $src = "src" + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) + $dest = "dest" + send_to_account(account: $dest) +`, balances(map[string]int64{"src": 100}), nil) + + requirePostings(t, []funds.Posting{posting("src", "dest", 10)}, res.Postings) +} + +// The `max [USD/2 20] from @src` shape: the cap is the smaller of the two. There +// is no min opcode — it is lt_int plus a branch, so both arms need covering, and +// the ties too (lt_int is strict). +func TestIRSourceCappedByMin(t *testing.T) { + // $cap = min($max, $amount), by copying $max and overwriting it unless it + // already won + src := ` + $asset = "USD/2" + set_current_asset($asset) + $amount = load_var(0) + $max = load_var(1) + $cap = int_copy($max) + $lt = lt_int($max, $amount) + jmp_if_true($lt, #min_end) + $cap = int_copy($amount) +#min_end + $src = "src" + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $cap, overdraft: $overdraft) + $dest = "dest" + send_to_account(account: $dest) +` + + testCases := []struct { + name string + amount, max int64 + wantSent int64 + }{ + {"right operand is smaller", 20, 50, 20}, + {"left operand is smaller", 50, 20, 20}, + {"equal operands", 20, 20, 20}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + vars := &vm.Vars{IntsPool: []big.Int{*big.NewInt(tc.amount), *big.NewInt(tc.max)}} + res := runIR(t, src, balances(map[string]int64{"src": 100}), vars) + requirePostings(t, []funds.Posting{posting("src", "dest", tc.wantSent)}, res.Postings) + }) + } +} + +// A comparison drives a real branch end to end, including the `!=` spelling that +// has no opcode of its own (eq_int + not). +func TestIRComparisonBranch(t *testing.T) { + // send the whole balance only when it differs from the requested amount, + // otherwise send the amount — a shape numscript can't express yet, which is + // the point of testing it here + src := ` + $asset = "USD/2" + set_current_asset($asset) + $src = "src" + $amount = 10 + $bal = balance($src, $asset) + $same = eq_int($bal, $amount) + $differs = not($same) + $cap = int_copy($amount) + jmp_if_false($differs, #end) + $cap = int_copy($bal) +#end + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $cap, overdraft: $overdraft) + $dest = "dest" + send_to_account(account: $dest) +` + + t.Run("balance differs, so it is sent whole", func(t *testing.T) { + res := runIR(t, src, balances(map[string]int64{"src": 4}), nil) + requirePostings(t, []funds.Posting{posting("src", "dest", 4)}, res.Postings) + }) + + t.Run("balance equals the amount, so the amount is sent", func(t *testing.T) { + res := runIR(t, src, balances(map[string]int64{"src": 10}), nil) + requirePostings(t, []funds.Posting{posting("src", "dest", 10)}, res.Postings) + }) +} + +func TestIRInorderSourcesStopAtFirstThatCovers(t *testing.T) { + // @a holds enough, so the forward jump must skip @b entirely + src := ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 10 + $pulled = 0 + $remaining = int_copy($amount) + $a = "a" + $overdraft = 0 + $from_a = pull_account(account: $a, cap: $remaining, overdraft: $overdraft) + $pulled += $from_a + $remaining -= $from_a + $exhausted = is_zero($remaining) + jmp_if_true($exhausted, #inorder_end) + $b = "b" + $from_b = pull_account(account: $b, cap: $remaining, overdraft: $overdraft) + $pulled += $from_b +#inorder_end + check_enough_funds($pulled, $amount) + $dest = "dest" + send_to_account(account: $dest) +` + + t.Run("first source covers it", func(t *testing.T) { + res := runIR(t, src, balances(map[string]int64{"a": 100, "b": 100}), nil) + requirePostings(t, []funds.Posting{posting("a", "dest", 10)}, res.Postings) + }) + + t.Run("falls through to the second", func(t *testing.T) { + res := runIR(t, src, balances(map[string]int64{"a": 4, "b": 100}), nil) + requirePostings(t, []funds.Posting{ + posting("a", "dest", 4), + posting("b", "dest", 6), + }, res.Postings) + }) + + t.Run("neither covers it", func(t *testing.T) { + execErr := runIRExpectingError(t, src, balances(map[string]int64{"a": 4, "b": 3}), nil) + require.IsType(t, vm.MissingFundsError{}, execErr) + }) +} + +func TestIRAllotmentDestination(t *testing.T) { + // 1/4 to @small, the remaining 3/4 to @big + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 100 + $world = "world" + $overdraft = 100 + $pulled = pull_account(account: $world, cap: $amount, overdraft: $overdraft) + $one = 1 + $four = 4 + $quarter = mk_portion($one, $four) + $whole = mk_portion($one, $one) + $leftover = sub_portion($whole, $quarter) + assert_leftover($leftover) +`+allot2IR("amount", "quarter", "leftover", "small_share", "big_share")+` + $small = "small" + send_to_account(account: $small, cap: $small_share) + $big = "big" + send_to_account(account: $big, cap: $big_share) +`, balances(nil), nil) + + requirePostings(t, []funds.Posting{ + posting("world", "small", 25), + posting("world", "big", 75), + }, res.Postings) +} + +func TestIRBalanceReadFromStore(t *testing.T) { + // send exactly what @src holds, read at run time + res := runIR(t, ` + $src = "src" + $asset = "USD/2" + $bal = balance($src, $asset) + assert_non_negative_balance($bal, $src) + set_current_asset($asset) + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $bal, overdraft: $overdraft) + check_enough_funds($pulled, $bal) + $dest = "dest" + send_to_account(account: $dest) +`, balances(map[string]int64{"src": 42}), nil) + + requirePostings(t, []funds.Posting{posting("src", "dest", 42)}, res.Postings) +} + +func TestIRUnsentFundsAreReturnedToTheSource(t *testing.T) { + // send_to_account with no account: the `kept` destination. The funds are + // released back and no posting is emitted for them. + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 100 + $src = "src" + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + $half = 50 + $dest = "dest" + send_to_account(account: $dest, cap: $half) + send_to_account() +`, balances(map[string]int64{"src": 100}), nil) + + requirePostings(t, []funds.Posting{posting("src", "dest", 50)}, res.Postings) +} + +func TestIRMarkBacktracks(t *testing.T) { + // the `oneof` shape as the compiler emits it: a region per branch, each failed + // one closed with a rewind and immediately reopened, committed once at the join + src := ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 10 + mark_push() + $a = "a" + $overdraft = 0 + $from_a = pull_account(account: $a, cap: $amount, overdraft: $overdraft) + $result = int_copy($from_a) + $missing = $amount - $from_a + $covered = is_zero($missing) + jmp_if_true($covered, #oneof_end) + mark_rewind() + mark_push() + $b = "b" + $from_b = pull_account(account: $b, cap: $amount, overdraft: $overdraft) + $result = int_copy($from_b) +#oneof_end + mark_commit() + check_enough_funds($result, $amount) + $dest = "dest" + send_to_account(account: $dest) +` + + t.Run("first branch covers it", func(t *testing.T) { + res := runIR(t, src, balances(map[string]int64{"a": 10, "b": 10}), nil) + requirePostings(t, []funds.Posting{posting("a", "dest", 10)}, res.Postings) + }) + + t.Run("rewinds to the second branch", func(t *testing.T) { + // @a can only cover part of it, so its partial funding must be discarded + res := runIR(t, src, balances(map[string]int64{"a": 3, "b": 10}), nil) + requirePostings(t, []funds.Posting{posting("b", "dest", 10)}, res.Postings) + }) +} + +// A rewind must undo only what its own region pulled: funds queued before the +// mark_push survive it and are still sendable afterwards. +func TestIRMarkRewindKeepsFundsQueuedBeforeThePush(t *testing.T) { + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $overdraft = 0 + $keep_amt = 4 + $kept = "kept" + $from_kept = pull_account(account: $kept, cap: $keep_amt, overdraft: $overdraft) + mark_push() + $spec_amt = 7 + $spec = "spec" + $from_spec = pull_account(account: $spec, cap: $spec_amt, overdraft: $overdraft) + mark_rewind() + $dest = "dest" + send_to_account(account: $dest) +`, balances(map[string]int64{"kept": 100, "spec": 100}), nil) + + // only the pre-mark pull reaches the destination; @spec was repaid + requirePostings(t, []funds.Posting{posting("kept", "dest", 4)}, res.Postings) +} + +// Nested regions must rewind independently: the inner one leaves the outer one's +// funds alone, and the outer rewind then discards everything. +func TestIRMarkNestedRegions(t *testing.T) { + src := ` + $asset = "USD/2" + set_current_asset($asset) + $overdraft = 0 + $ten = 10 + mark_push() + $outer = "outer" + $from_outer = pull_account(account: $outer, cap: $ten, overdraft: $overdraft) + mark_push() + $inner = "inner" + $from_inner = pull_account(account: $inner, cap: $ten, overdraft: $overdraft) + mark_rewind() +` + balances := balances(map[string]int64{"outer": 100, "inner": 100}) + + t.Run("outer region commits", func(t *testing.T) { + res := runIR(t, src+` + mark_commit() + $dest = "dest" + send_to_account(account: $dest) +`, balances, nil) + // the inner rewind dropped @inner; @outer survived it and is committed + requirePostings(t, []funds.Posting{posting("outer", "dest", 10)}, res.Postings) + }) + + t.Run("outer region rewinds too", func(t *testing.T) { + res := runIR(t, src+` + mark_rewind() + $dest = "dest" + send_to_account(account: $dest) +`, balances, nil) + requirePostings(t, []funds.Posting{}, res.Postings) + }) +} + +// A mark op with nothing to act on is a malformed program, not a script outcome: +// it is a bug in whatever produced the bytecode, so Verify rejects it and, run +// unverified, it surfaces as an InternalError. The point is that it never panics — the old index-valued restore +// truncated the source queue to an arbitrary int, which panicked out of range or, +// worse, resurrected already-consumed entries. +func TestIRMarkWithNoOpenRegionIsAnInternalError(t *testing.T) { + cases := map[string]string{ + "rewind with no push": ` + mark_rewind() +`, + "commit with no push": ` + mark_commit() +`, + "one push, two ends": ` + mark_push() + mark_commit() + mark_commit() +`, + "rewind after the region closed": ` + mark_push() + mark_commit() + mark_rewind() +`, + // a rewind closes too, so a second end has nothing left to act on + "rewind then commit": ` + mark_push() + mark_rewind() + mark_commit() +`, + } + + for name, src := range cases { + t.Run(name, func(t *testing.T) { + execErr := runUnverifiedIRExpectingError(t, src, "mark end with no open mark", balances(nil), nil) + require.IsType(t, vm.InternalError{}, execErr) + require.ErrorContains(t, execErr, "no open mark") + }) + } +} + +// A mark is a source-queue depth, so it only means anything while nothing drains +// the queue from the front and the asset a repay lands on is fixed. All three ops +// that would break it are rejected inside a region rather than silently corrupting +// balances. Compiled numscript never emits any of them inside one — sources only +// pull, and `save` is a statement — so this is reachable only from hand-written IR +// (or a hand-crafted .numb), and Verify rejects it statically. +func TestIRSendAndSetAssetAreRejectedInsideARegion(t *testing.T) { + prelude := ` + $asset = "USD/2" + set_current_asset($asset) + $overdraft = 0 + $ten = 10 + $src = "src" + mark_push() + $pulled = pull_account(account: $src, cap: $ten, overdraft: $overdraft) +` + store := balances(map[string]int64{"src": 100}) + + t.Run("send inside a region", func(t *testing.T) { + execErr := runUnverifiedIRExpectingError(t, prelude+` + $dest = "dest" + send_to_account(account: $dest) +`, "while a mark is open", store, nil) + require.IsType(t, vm.InternalError{}, execErr) + require.ErrorContains(t, execErr, "send while a mark is open") + }) + + t.Run("uncapped send inside a region", func(t *testing.T) { + execErr := runUnverifiedIRExpectingError(t, prelude+` + send_to_account() +`, "while a mark is open", store, nil) + require.IsType(t, vm.InternalError{}, execErr) + require.ErrorContains(t, execErr, "send while a mark is open") + }) + + t.Run("set_current_asset inside a region", func(t *testing.T) { + execErr := runUnverifiedIRExpectingError(t, prelude+` + $other = "EUR/2" + set_current_asset($other) +`, "while a mark is open", store, nil) + require.IsType(t, vm.InternalError{}, execErr) + require.ErrorContains(t, execErr, "set_current_asset while a mark is open") + }) + + // save reduces a balance, and a rewind only repays queued sources — so a save in + // an abandoned branch would persist. It is the one op on this list that stays + // forbidden no matter how much rollback is added later: its floor at zero is not + // invertible from a delta. + t.Run("save inside a region", func(t *testing.T) { + execErr := runUnverifiedIRExpectingError(t, prelude+` + $five = 5 + save(account: $src, asset: $asset, amount: $five) +`, "while a mark is open", store, nil) + require.IsType(t, vm.InternalError{}, execErr) + require.ErrorContains(t, execErr, "save while a mark is open") + }) + + t.Run("save-all inside a region", func(t *testing.T) { + execErr := runUnverifiedIRExpectingError(t, prelude+` + save(account: $src, asset: $asset) +`, "while a mark is open", store, nil) + require.IsType(t, vm.InternalError{}, execErr) + require.ErrorContains(t, execErr, "save while a mark is open") + }) + + // all three are fine once the region has closed, so the check is scoped to the + // region and not a blanket ban + t.Run("all three are allowed after the region closes", func(t *testing.T) { + res := runIR(t, prelude+` + mark_commit() + $dest = "dest" + send_to_account(account: $dest) + $five = 5 + save(account: $src, asset: $asset, amount: $five) + $other = "EUR/2" + set_current_asset($other) +`, store, nil) + requirePostings(t, []funds.Posting{posting("src", "dest", 10)}, res.Postings) + }) +} + +// Interleaving pulls with mark ops across a jump: the region spans a branch, and +// both paths through it reach the same single mark_commit. This is the shape the +// compiler emits, and the reason mark depth stays a function of position. +func TestIRMarkAcrossAJump(t *testing.T) { + src := ` + $asset = "USD/2" + set_current_asset($asset) + $overdraft = 0 + $ten = 10 + $zero = 0 + mark_push() + $a = "a" + $from_a = pull_account(account: $a, cap: $ten, overdraft: $overdraft) + $took_nothing = is_zero($from_a) + jmp_if_false($took_nothing, #done) + $b = "b" + $from_b = pull_account(account: $b, cap: $ten, overdraft: $overdraft) +#done + mark_commit() + $dest = "dest" + send_to_account(account: $dest) +` + + t.Run("branch taken", func(t *testing.T) { + // @a is empty, so the jump falls through to the @b pull + res := runIR(t, src, balances(map[string]int64{"a": 0, "b": 10}), nil) + requirePostings(t, []funds.Posting{posting("b", "dest", 10)}, res.Postings) + }) + + t.Run("branch skipped", func(t *testing.T) { + res := runIR(t, src, balances(map[string]int64{"a": 10, "b": 10}), nil) + requirePostings(t, []funds.Posting{posting("a", "dest", 10)}, res.Postings) + }) +} + +// A run that dies inside a region must not leak the open mark into the next run +// on the same Vm: the reused RunState drops it, so the second run's send works. +func TestIRMarkDoesNotLeakAcrossRuns(t *testing.T) { + // var 0 != 0 opens a mark before the send, which fails inside the region + machine := vm.NewVm(assembleUnverifiedIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $overdraft = 0 + $ten = 10 + $src = "src" + $dest = "dest" + $open = load_var(0) + $skip = is_zero($open) + jmp_if_true($skip, #send) + mark_push() +#send + $pulled = pull_account(account: $src, cap: $ten, overdraft: $overdraft) + send_to_account(account: $dest) +`)) + store := balances(map[string]int64{"src": 100}) + withMark := &vm.Vars{IntsPool: []big.Int{*big.NewInt(1)}} + withoutMark := &vm.Vars{IntsPool: []big.Int{*big.NewInt(0)}} + + _, execErr := vm.Exec(context.Background(), machine, withMark, store) + require.Equal(t, vm.InternalError{Err: vm.ErrSendWhileMarkOpen}, execErr) + + res, execErr := vm.Exec(context.Background(), machine, withoutMark, store) + require.Nil(t, execErr) + requirePostings(t, []funds.Posting{posting("src", "dest", 10)}, res.Postings) + + res, execErr = vm.Exec(context.Background(), machine, withMark, store) + require.Equal(t, vm.InternalError{Err: vm.ErrSendWhileMarkOpen}, execErr) + require.Empty(t, res.Postings) +} + +func TestIRStrEqAndJmp(t *testing.T) { + // the if/else shape: str_eq is the only way to branch on a string, and jmp is + // what skips the else arm. Here the taken arm decides which account is pulled + // from, which is how @world's unboundedness is expressed in bytecode. + src := ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 10 + $overdraft = 0 + $expected = "yes" + $probe = load_var(0) + $eq = str_eq($probe, $expected) + jmp_if_false($eq, #else) + $then_acc = "a" + $pulled = pull_account(account: $then_acc, cap: $amount, overdraft: $overdraft) + jmp(#end) +#else + $else_acc = "b" + $pulled = pull_account(account: $else_acc, cap: $amount, overdraft: $overdraft) +#end + $dest = "dest" + send_to_account(account: $dest) +` + + store := balances(map[string]int64{"a": 10, "b": 10}) + + t.Run("equal strings take the then arm", func(t *testing.T) { + res := runIR(t, src, store, &vm.Vars{StringsPool: []string{"yes"}}) + requirePostings(t, []funds.Posting{posting("a", "dest", 10)}, res.Postings) + }) + + t.Run("and jmp skips it otherwise", func(t *testing.T) { + res := runIR(t, src, store, &vm.Vars{StringsPool: []string{"no"}}) + requirePostings(t, []funds.Posting{posting("b", "dest", 10)}, res.Postings) + }) +} + +func TestIRSaveWithholdsFunds(t *testing.T) { + // save reserves part of the balance, so the pull can't reach it + src := ` + $asset = "USD/2" + set_current_asset($asset) + $src = "src" + $reserved = 30 + save(account: $src, asset: $asset, amount: $reserved) + $amount = 100 + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + $dest = "dest" + send_to_account(account: $dest) +` + + res := runIR(t, src, balances(map[string]int64{"src": 100}), nil) + requirePostings(t, []funds.Posting{posting("src", "dest", 70)}, res.Postings) +} + +func TestIROverdraftAllowsNegativeBalance(t *testing.T) { + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 40 + $src = "src" + $overdraft = 25 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) + $dest = "dest" + send_to_account(account: $dest) +`, balances(map[string]int64{"src": 15}), nil) + + // 15 on the account plus 25 of allowed overdraft + requirePostings(t, []funds.Posting{posting("src", "dest", 40)}, res.Postings) +} + +func TestIRMetadata(t *testing.T) { + res := runIR(t, ` + $account = "acc" + $key = "k" + $value = "v" + set_account_meta($account, $key, $value) + $tx_key = "tx" + $tx_value = "yes" + set_tx_meta($tx_key, $tx_value) +`, balances(nil), nil) + + require.Equal(t, map[string]string{"tx": "yes"}, res.Metadata) + require.Equal(t, funds.AccountsMetadata{{Account: "acc", Key: "k", Value: "v"}}, res.AccountsMetadata) +} + +func TestIRReadsMetadataFromStore(t *testing.T) { + // the amount to send is an int read out of @src's metadata + store := irStore{ + balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2"}: big.NewInt(100), + }, + metadata: map[string]map[string]string{ + "src": {"quota": "7"}, + }, + } + + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $src = "src" + $key = "quota" + $amount = meta($src, $key) + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) + $dest = "dest" + send_to_account(account: $dest) +`, store, nil) + + requirePostings(t, []funds.Posting{posting("src", "dest", 7)}, res.Postings) +} + +func TestIRMissingMetadataIsAnError(t *testing.T) { + execErr := runIRExpectingError(t, ` + $src = "src" + $key = "nope" + $value = meta($src, $key) +`, balances(nil), nil) + + require.IsType(t, vm.MetadataNotFoundError{}, execErr) +} + +func TestIRLoadsVars(t *testing.T) { + // vars come in as pools, indexed by the load_var instructions + vars := &vm.Vars{ + StringsPool: []string{"USD/2", "src", "dest"}, + IntsPool: []big.Int{*big.NewInt(10), *big.NewInt(0)}, + } + + res := runIR(t, ` + $asset = load_var(0) + set_current_asset($asset) + $amount = load_var(0) + $src = load_var(1) + $overdraft = load_var(1) + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) + $dest = load_var(2) + send_to_account(account: $dest) +`, balances(map[string]int64{"src": 100}), vars) + + requirePostings(t, []funds.Posting{posting("src", "dest", 10)}, res.Postings) +} + +func TestIRAssertions(t *testing.T) { + t.Run("invalid account name", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $account = "not a valid account!" + assert_valid_account($account) +`, balances(nil), nil) + require.IsType(t, vm.InvalidAccountName{}, execErr) + }) + + t.Run("invalid color", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $color = "not a color" + assert_valid_color($color) +`, balances(nil), nil) + require.IsType(t, vm.InvalidColor{}, execErr) + }) + + t.Run("empty color is valid", func(t *testing.T) { + res := runIR(t, ` + $color = "" + assert_valid_color($color) +`, balances(nil), nil) + require.Empty(t, res.Postings) + }) + + t.Run("mismatched assets", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $usd = "USD/2" + $eur = "EUR/2" + assert_same_asset($usd, $eur) +`, balances(nil), nil) + require.IsType(t, vm.AssetMismatchError{}, execErr) + }) + + t.Run("negative balance", func(t *testing.T) { + store := irStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2"}: big.NewInt(-1), + }} + execErr := runIRExpectingError(t, ` + $src = "src" + $asset = "USD/2" + $bal = balance($src, $asset) + assert_non_negative_balance($bal, $src) +`, store, nil) + require.IsType(t, vm.NegativeBalanceError{}, execErr) + }) + + t.Run("allotment portions over 100%", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $one = 1 + $two = 2 + $half = mk_portion($one, $two) + $whole = mk_portion($one, $one) + $leftover = sub_portion($whole, $half) + $negative = sub_portion($leftover, $whole) + assert_leftover($negative) +`, balances(nil), nil) + require.IsType(t, vm.InvalidAllotmentSum{}, execErr) + }) + + t.Run("negative allotment portion", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $zero = 0 + $one = 1 + $minusOne = sub_int($zero, $one) + $three = 3 + $portion = mk_portion($minusOne, $three) + assert_non_negative_portion($portion) +`, balances(nil), nil) + require.IsType(t, vm.NegativePortionError{}, execErr) + }) +} + +func TestIRRejectsInvalidPostings(t *testing.T) { + for name, tc := range map[string]struct{ asset, dest string }{ + "account": {asset: "USD/2", dest: "not valid"}, + "asset": {asset: "usd", dest: "dest"}, + } { + t.Run(name, func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $asset = "`+tc.asset+`" + set_current_asset($asset) + $src = "src" + $amount = 10 + $overdraft = 100 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + $dest = "`+tc.dest+`" + send_to_account(account: $dest) +`, balances(map[string]int64{"src": 70}), nil) + + var invalid vm.InvalidPostingError + require.IsType(t, vm.InternalError{}, execErr) + require.ErrorAs(t, execErr, &invalid) + require.Equal(t, tc.dest, invalid.Posting.Destination) + }) + } +} + +func TestIRAssertNonNegativePortionAcceptsZero(t *testing.T) { + runIR(t, ` + $zero = 0 + $one = 1 + $portion = mk_portion($zero, $one) + assert_non_negative_portion($portion) +`, balances(nil), nil) +} + +func TestIRUncappedPull(t *testing.T) { + // no cap: the pull is bounded only by the overdraft, i.e. `send *` + t.Run("with an overdraft", func(t *testing.T) { + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $src = "src" + $overdraft = 0 + $pulled = pull_account(account: $src, overdraft: $overdraft) + $dest = "dest" + send_to_account(account: $dest) +`, balances(map[string]int64{"src": 70}), nil) + + requirePostings(t, []funds.Posting{posting("src", "dest", 70)}, res.Postings) + }) + + t.Run("without one it is unbounded and rejected", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $asset = "USD/2" + set_current_asset($asset) + $src = "src" + $pulled = pull_account(account: $src) +`, balances(map[string]int64{"src": 70}), nil) + + require.IsType(t, vm.InvalidUncappedSource{}, execErr) + }) +} + +func TestIRSaveAll(t *testing.T) { + // save with no amount withholds the whole balance, so the pull finds nothing + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $src = "src" + save(account: $src, asset: $asset) + $amount = 100 + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + $dest = "dest" + send_to_account(account: $dest) +`, balances(map[string]int64{"src": 100}), nil) + + require.Empty(t, res.Postings) +} + +func TestIRAssertLeftoverExact(t *testing.T) { + // the no-`remaining` form: portions must cover exactly 1 + execErr := runIRExpectingError(t, ` + $one = 1 + $two = 2 + $half = mk_portion($one, $two) + $whole = mk_portion($one, $one) + $leftover = sub_portion($whole, $half) + assert_leftover_exact($leftover) +`, balances(nil), nil) + + require.IsType(t, vm.InvalidAllotmentSum{}, execErr) +} + +func TestIRMetaTypes(t *testing.T) { + store := meta(map[string]map[string]string{ + "acc": { + "portion": "1/4", + "monetary": "USD/2 250", + "oops": "not a number", + }, + }) + + t.Run("portion", func(t *testing.T) { + // the portion drives an allotment, so the split proves it parsed + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 100 + $world = "world" + $overdraft = 100 + $pulled = pull_account(account: $world, cap: $amount, overdraft: $overdraft) + $acc = "acc" + $key = "portion" + $quarter = meta($acc, $key) + $one = 1 + $whole = mk_portion($one, $one) + $rest = sub_portion($whole, $quarter) + assert_leftover($rest) +`+allot2IR("amount", "quarter", "rest", "a_share", "b_share")+` + $a = "a" + send_to_account(account: $a, cap: $a_share) + $b = "b" + send_to_account(account: $b, cap: $b_share) +`, store, nil) + + requirePostings(t, []funds.Posting{ + posting("world", "a", 25), + posting("world", "b", 75), + }, res.Postings) + }) + + t.Run("monetary", func(t *testing.T) { + res := runIR(t, ` + $acc = "acc" + $key = "monetary" + [$asset, $amount] = meta_monetary($acc, $key) + set_current_asset($asset) + $overdraft = 300 + $pulled = pull_account(account: $acc, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) + $dest = "dest" + send_to_account(account: $dest) +`, store, nil) + + requirePostings(t, []funds.Posting{posting("acc", "dest", 250)}, res.Postings) + }) + + t.Run("a value of the wrong shape is an error", func(t *testing.T) { + for _, read := range []string{ + ` $v = meta($acc, $key)`, + ` $v = meta($acc, $key)`, + ` [$a, $n] = meta_monetary($acc, $key)`, + } { + execErr := runIRExpectingError(t, ` + $acc = "acc" + $key = "oops" +`+read+"\n", store, nil) + require.IsType(t, vm.BadMetaValueError{}, execErr, "%s", read) + } + }) +} + +func TestIRStoreErrorsPropagate(t *testing.T) { + failing := irStore{err: errors.New("store is down")} + + t.Run("on a balance read", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $src = "src" + $asset = "USD/2" + $bal = balance($src, $asset) +`, failing, nil) + require.IsType(t, vm.StoreError{}, execErr) + require.ErrorContains(t, execErr, "store is down") + }) + + t.Run("on a pull", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 10 + $src = "src" + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) +`, failing, nil) + require.IsType(t, vm.StoreError{}, execErr) + }) + + t.Run("on a metadata read", func(t *testing.T) { + for _, read := range []string{ + ` $v = meta($acc, $key)`, + ` $v = meta($acc, $key)`, + ` $v = meta($acc, $key)`, + ` [$a, $n] = meta_monetary($acc, $key)`, + } { + execErr := runIRExpectingError(t, ` + $acc = "acc" + $key = "k" +`+read+"\n", failing, nil) + require.IsType(t, vm.StoreError{}, execErr, "%s", read) + } + }) + + t.Run("on an uncapped pull", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $asset = "USD/2" + set_current_asset($asset) + $src = "src" + $overdraft = 0 + $pulled = pull_account(account: $src, overdraft: $overdraft) +`, failing, nil) + require.IsType(t, vm.StoreError{}, execErr) + }) + + t.Run("on a save", func(t *testing.T) { + execErr := runIRExpectingError(t, ` + $acc = "acc" + $asset = "USD/2" + $amount = 10 + save(account: $acc, asset: $asset, amount: $amount) +`, failing, nil) + require.IsType(t, vm.StoreError{}, execErr) + }) + + t.Run("but not on a send: crediting a destination reads nothing", func(t *testing.T) { + // the pull has no overdraft operand, so it is unbounded and reads no balance + // either — the whole send runs against a store that fails every call. This + // is the arm the compiler emits for @world. + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $world = "world" + $amount = 10 + $pulled = pull_account(account: $world, cap: $amount) + $dest = "dest" + send_to_account(account: $dest) +`, failing, nil) + requirePostings(t, []funds.Posting{posting("world", "dest", 10)}, res.Postings) + }) +} + +// The bogus-mark case the old index-valued restore needed a guard for is gone: +// with no operand there is no value to pass, so an out-of-range mark is not +// expressible in the IR at all. Misuse can only be an unbalanced stack, covered by +// TestIRMarkWithNoOpenRegionIsAnInternalError. +func TestIRMarkTakesNoOperand(t *testing.T) { + _, errs := ir.Parse(` + $mark = 3 + mark_rewind($mark) +`) + require.NotEmpty(t, errs, "mark_rewind must not accept an operand") +} + +func TestIRVmIsReusableAcrossRuns(t *testing.T) { + // a second run must not see the first one's funds or postings + program := assembleIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 10 + $src = "src" + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) + $dest = "dest" + send_to_account(account: $dest) +`) + machine := vm.NewVm(program) + store := balances(map[string]int64{"src": 100}) + + want := []funds.Posting{posting("src", "dest", 10)} + for run := 1; run <= 3; run++ { + res, execErr := vm.Exec(context.Background(), machine, nil, store) + require.Nil(t, execErr, "run %d", run) + requirePostings(t, want, res.Postings) + } +} + +// TestIRSurvivesTheWireFormat runs one program in memory and again after a trip +// through Encode/DecodeProgram. Nothing else ties the encoder to the VM. +// No instruction reads a bool yet, so this only pins down that the bank is +// allocated separately from the others and that the two ops run. +func TestIRConstBool(t *testing.T) { + program := assembleIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $t = true + $f = false + $amount = 10 + $overdraft = 0 + $src = "src" + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + $dest = "dest" + send_to_account(account: $dest) +`) + + // $t and $f are never read again after being set, so the allocator frees + // $t's slot right after it and reuses it for $f — both share bool bank + // index 0. + require.Equal(t, byte(1), program.MaxRegBool) + // $amount, $overdraft, $pulled — the two bools are not among them + require.Equal(t, byte(3), program.MaxRegInt, "bools don't consume int registers") + + decoded, err := vm.DecodeProgram(program.Encode()) + require.NoError(t, err) + require.Equal(t, program, decoded) + + res, execErr := vm.Exec(context.Background(), vm.NewVm(decoded), nil, balances(map[string]int64{"src": 10})) + require.Nil(t, execErr, "unexpected execution error: %v", execErr) + requirePostings(t, []funds.Posting{posting("src", "dest", 10)}, res.Postings) +} + +func TestIRSurvivesTheWireFormat(t *testing.T) { + program := assembleIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 100 + $src = "src" + $overdraft = 0 + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) + $one = 1 + $two = 2 + $half = mk_portion($one, $two) + $whole = mk_portion($one, $one) + $rest = sub_portion($whole, $half) + assert_leftover($rest) +`+allot2IR("amount", "half", "rest", "a_share", "b_share")+` + $a = "a" + send_to_account(account: $a, cap: $a_share) + $b = "b" + send_to_account(account: $b, cap: $b_share) +`) + + decoded, err := vm.DecodeProgram(program.Encode()) + require.NoError(t, err) + require.Equal(t, program, decoded, "the program changed shape on the way through") + + want := []funds.Posting{posting("src", "a", 50), posting("src", "b", 50)} + for name, prog := range map[string]vm.Program{"in memory": program, "decoded": decoded} { + res, execErr := vm.Exec(context.Background(), vm.NewVm(prog), nil, balances(map[string]int64{"src": 100})) + require.Nil(t, execErr, "%s: %v", name, execErr) + requirePostings(t, want, res.Postings) + } +} + +// TestIRVarsSurviveTheWireFormat is the same for the vars payload. +func TestIRVarsSurviveTheWireFormat(t *testing.T) { + vars := vm.Vars{ + StringsPool: []string{"USD/2", "src", "dest"}, + IntsPool: []big.Int{*big.NewInt(10), *big.NewInt(0)}, + } + decoded, err := vm.DecodeVars(vars.Encode()) + require.NoError(t, err) + + program := assembleIR(t, ` + $asset = load_var(0) + set_current_asset($asset) + $amount = load_var(0) + $src = load_var(1) + $overdraft = load_var(1) + $pulled = pull_account(account: $src, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) + $dest = load_var(2) + send_to_account(account: $dest) +`) + + res, execErr := vm.Exec(context.Background(), vm.NewVm(program), &decoded, balances(map[string]int64{"src": 100})) + require.Nil(t, execErr) + requirePostings(t, []funds.Posting{posting("src", "dest", 10)}, res.Postings) +} + +// --- The int/portion boundary ops ------------------------------------------- + +func TestIRPortionToIntFloors(t *testing.T) { + // 7/2 of nothing in particular: the projection floors, it does not round + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $seven = 7 + $two = 2 + $p = mk_portion($seven, $two) + $amount = portion_to_int($p) + $world = "world" + $overdraft = 100 + $pulled = pull_account(account: $world, cap: $amount, overdraft: $overdraft) + $dest = "dest" + send_to_account(account: $dest) +`, balances(nil), nil) + + requirePostings(t, []funds.Posting{posting("world", "dest", 3)}, res.Postings) +} + +func TestIRIntToPortionAndMul(t *testing.T) { + // 1/4 * 100 == 25, computed as mul_portion(int_to_portion(100), 1/4) + res := runIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $hundred = 100 + $one = 1 + $four = 4 + $quarter = mk_portion($one, $four) + $ap = int_to_portion($hundred) + $prod = mul_portion($quarter, $ap) + $amount = portion_to_int($prod) + $world = "world" + $overdraft = 100 + $pulled = pull_account(account: $world, cap: $amount, overdraft: $overdraft) + $dest = "dest" + send_to_account(account: $dest) +`, balances(nil), nil) + + requirePostings(t, []funds.Posting{posting("world", "dest", 25)}, res.Postings) +} + +// A three-way split written out of pure ops: floor each share, then hand the +// flooring leftover to the earliest shares one unit at a time. This is the +// lowering compileAllotmentSplit emits, pinned here on the 34/33/33 case. +// +// Only n-1 fixup blocks: each floor loses < 1, so with portions summing to 1 the +// shortfall is <= n-1 and the last share never receives a unit. +const allotThirdsIR = ` + $asset = "USD/2" + set_current_asset($asset) + $amount = 100 + $world = "world" + $overdraft = 100 + $pulled = pull_account(account: $world, cap: $amount, overdraft: $overdraft) + + $one = 1 + $three = 3 + $third = mk_portion($one, $three) + $ap = int_to_portion($amount) + + $prod = mul_portion($third, $ap) + $out0 = portion_to_int($prod) + $total = int_copy($out0) + $prod = mul_portion($third, $ap) + $out1 = portion_to_int($prod) + $total = add_int($total, $out1) + $prod = mul_portion($third, $ap) + $out2 = portion_to_int($prod) + $total = add_int($total, $out2) + + $lt = lt_int($total, $amount) + jmp_if_false($lt, #done) + $out0 = add_int($out0, $one) + $total = add_int($total, $one) + $lt = lt_int($total, $amount) + jmp_if_false($lt, #done) + $out1 = add_int($out1, $one) + $total = add_int($total, $one) +#done + + $a = "a" + send_to_account(account: $a, cap: $out0) + $b = "b" + send_to_account(account: $b, cap: $out1) + $c = "c" + send_to_account(account: $c, cap: $out2) +` + +func TestIRAllotmentFromPureOps(t *testing.T) { + res := runIR(t, allotThirdsIR, balances(nil), nil) + + requirePostings(t, []funds.Posting{ + posting("world", "a", 34), + posting("world", "b", 33), + posting("world", "c", 33), + }, res.Postings) +} + +// An error returned by one run must not change when the same Vm runs again. +func TestIRErrorsDoNotAliasRegisters(t *testing.T) { + t.Run("missing funds", func(t *testing.T) { + machine := vm.NewVm(assembleIR(t, ` + $asset = "USD/2" + set_current_asset($asset) + $amount = load_var(0) + $a = "a" + $overdraft = 0 + $pulled = pull_account(account: $a, cap: $amount, overdraft: $overdraft) + check_enough_funds($pulled, $amount) +`)) + _, first := vm.Exec(context.Background(), machine, &vm.Vars{IntsPool: []big.Int{*big.NewInt(10)}}, balances(map[string]int64{"a": 4})) + _, second := vm.Exec(context.Background(), machine, &vm.Vars{IntsPool: []big.Int{*big.NewInt(20)}}, balances(map[string]int64{"a": 7})) + + require.Equal(t, vm.MissingFundsError{Asset: "USD/2", Needed: big.NewInt(10), Got: big.NewInt(4)}, first) + require.Equal(t, vm.MissingFundsError{Asset: "USD/2", Needed: big.NewInt(20), Got: big.NewInt(7)}, second) + }) + + t.Run("negative amount", func(t *testing.T) { + machine := vm.NewVm(assembleIR(t, ` + $amount = load_var(0) + assert_non_negative_amount($amount) +`)) + _, first := vm.Exec(context.Background(), machine, &vm.Vars{IntsPool: []big.Int{*big.NewInt(-10)}}, balances(nil)) + _, second := vm.Exec(context.Background(), machine, &vm.Vars{IntsPool: []big.Int{*big.NewInt(-20)}}, balances(nil)) + + require.Equal(t, vm.NegativeAmountError{Amount: *big.NewInt(-10)}, first) + require.Equal(t, vm.NegativeAmountError{Amount: *big.NewInt(-20)}, second) + }) +} diff --git a/internal/vm/meta_test.go b/internal/vm/meta_test.go new file mode 100644 index 00000000..ef16247e --- /dev/null +++ b/internal/vm/meta_test.go @@ -0,0 +1,58 @@ +package vm + +import ( + "context" + "math/big" + "testing" + + "github.com/formancehq/numscript/internal/funds" + "github.com/stretchr/testify/require" +) + +func TestSetAccountMeta(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadStr, 0, 0), // r_s0 = "acc" + bc(Op_LoadStr, 1, 1), // r_s1 = "k" + bc(Op_LoadStr, 2, 2), // r_s2 = "v" + abc(Op_SetAccountMeta, 0, 1, 2), // set_account_meta(acc, k, v) + abc(0, nilReg, nilReg, nilReg), // ext: no scope + }, + StringsPool: []string{"acc", "k", "v"}, + } + + res, execErr := Exec(context.Background(), newTestVm(prog), nil, mockStore{}) + require.Nil(t, execErr) + require.Equal(t, funds.AccountsMetadata{{Account: "acc", Key: "k", Value: "v"}}, res.AccountsMetadata) +} + +func TestMetaStr(t *testing.T) { + // meta("config", "beneficiary") == "alice"; then send [USD/2 100] from world to it. + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadStr, 0, 0), // s0 = "USD/2" + abc(Op_SetCurrentAsset, 0, 0, 0), + bc(Op_LoadStr, 1, 1), // s1 = "config" + bc(Op_LoadStr, 2, 2), // s2 = "beneficiary" + abc(Op_MetaStr, 3, 1, 2), // s3 = meta(config, beneficiary) = "alice" + abc(0, nilReg, nilReg, nilReg), // ext: no scope + bc(Op_LoadStr, 4, 3), // s4 = "world" + bc(Op_LoadInt, 0, 0), // i0 = 100 (cap) + abc(Op_PullAccount, 1, 4, 0), // i1 = pull(world, cap i0) + abc(0, nilReg, nilReg, nilReg), // ext: no overdraft, no color, no scope + abc(Op_SendToAccount, 3, nilReg, nilReg), // send to s3 (alice) + }, + StringsPool: []string{"USD/2", "config", "beneficiary", "world"}, + IntsPool: []big.Int{*big.NewInt(100)}, + } + + store := mockStore{meta: map[string]map[string]string{ + "config": {"beneficiary": "alice"}, + }} + + res, execErr := Exec(context.Background(), newTestVm(prog), nil, store) + require.Nil(t, execErr) + require.Equal(t, []funds.Posting{ + {Source: "world", Destination: "alice", Asset: "USD/2", Amount: big.NewInt(100)}, + }, res.Postings) +} diff --git a/internal/vm/program.go b/internal/vm/program.go new file mode 100644 index 00000000..c40de2c9 --- /dev/null +++ b/internal/vm/program.go @@ -0,0 +1,238 @@ +package vm + +import ( + "encoding/binary" + "fmt" + "math/big" +) + +type Program struct { + Instructions []Instruction + + StringsPool []string + IntsPool []big.Int + + MaxRegString byte + MaxRegInt byte + MaxRegPortion byte + MaxRegBool byte + + // Version is the bytecode version this program was assembled for, or (after + // DecodeProgram) the version it was actually encoded with. Encode always + // writes CurrentBytecodeVersion regardless of this field. + Version BytecodeVersion +} + +var le = binary.LittleEndian + +// TODO review AI blob +func (p Program) Encode() []byte { + instrs := make([]byte, 4*len(p.Instructions)) + for i, ins := range p.Instructions { + instrs[i*4], instrs[i*4+1], instrs[i*4+2], instrs[i*4+3] = ins.Opcode, ins.A, ins.B, ins.C + } + + strs := encodeStringsPool(p.StringsPool) + ints := encodeIntsPool(p.IntsPool) + maxRegs := encodeMaxRegs(p) + + buf := make([]byte, 0, formatHeaderLen+4*6+len(instrs)+len(strs)+len(ints)+len(maxRegs)) + buf = appendFormatHeader(buf, "NUMB", 4) + buf = appendSection(buf, SectionInstructions, instrs) + buf = appendSection(buf, SectionStringsPool, strs) + buf = appendSection(buf, SectionIntsPool, ints) + buf = appendSection(buf, SectionMaxRegisters, maxRegs) + return buf +} + +// These fields hold the per-bank register *count* (== highest index + 1), as +// emitted by the assembler. Real indices are 0..0xFE (0xFF is the nil sentinel), +// so the largest possible count is 255. When the max-registers section is absent +// we assume the bank uses every usable register, i.e. this default. +const maxRegDefault byte = 255 + +func encodeMaxRegs(p Program) []byte { + return []byte{p.MaxRegString, p.MaxRegInt, p.MaxRegPortion, p.MaxRegBool} +} + +// One byte per bank, positional, append-only order. The section length is the +// number of banks the writer knew. +// +// - absent (len 0): no info, so every bank defaults to maxRegDefault (safe). +// - present: bank i uses buf[i] when i < len; banks beyond len default to 0, +// since a bank the (older) writer didn't know is a type the program predates +// and provably uses none of. +// +// Extra trailing bytes (a newer writer) are ignored; a program that actually uses +// such a bank is rejected later via its unknown opcodes. +func parseMaxRegs(buf []byte) (str, i, portion, bool_ byte) { + if len(buf) == 0 { + return maxRegDefault, maxRegDefault, maxRegDefault, maxRegDefault + } + at := func(idx int) byte { + if idx < len(buf) { + return buf[idx] + } + return 0 + } + return at(0), at(1), at(2), at(3) +} + +func encodeStringsPool(strings []string) []byte { + buf := make([]byte, 4) + le.PutUint32(buf, uint32(len(strings))) + var lenBuf [4]byte + for _, s := range strings { + le.PutUint32(lenBuf[:], uint32(len(s))) + buf = append(buf, lenBuf[:]...) + buf = append(buf, s...) + } + return buf +} + +func encodeIntsPool(ints []big.Int) []byte { + buf := make([]byte, 4) + le.PutUint32(buf, uint32(len(ints))) + var lenBuf [4]byte + for i := range ints { + n := &ints[i] + sign := byte(0) + if n.Sign() < 0 { + sign = 1 + } + mag := n.Bytes() // absolute value, big-endian (big.Int's native form) + buf = append(buf, sign) + le.PutUint32(lenBuf[:], uint32(len(mag))) + buf = append(buf, lenBuf[:]...) + buf = append(buf, mag...) + } + return buf +} + +func parseInstructions(buf []byte) ([]Instruction, error) { + if len(buf)%4 != 0 { + return nil, fmt.Errorf("instructions section size %d not a multiple of 4", len(buf)) + } + instructions := make([]Instruction, len(buf)/4) + for i := range instructions { + off := i * 4 + instructions[i] = Instruction{ + buf[off], + buf[off+1], + buf[off+2], + buf[off+3], + } + } + return instructions, nil +} + +func parseStringsPool(buf []byte) ([]string, error) { + if len(buf) == 0 { + return nil, nil + } + if len(buf) < 4 { + return nil, fmt.Errorf("strings pool: count truncated") + } + n := le.Uint32(buf) + total := uint64(len(buf)) + if uint64(n)*4 > total-4 { // every record is at least a 4B length prefix + return nil, fmt.Errorf("strings pool: count %d exceeds buffer size %d", n, total) + } + idx := uint64(4) + out := make([]string, n) + for i := range out { + if idx+4 > total { + return nil, fmt.Errorf("string %d: length prefix out of bounds", i) + } + strLen := uint64(le.Uint32(buf[idx:])) + idx += 4 + end := idx + strLen // operands <= ~4.3e9, sum fits in uint64 + if end > total { + return nil, fmt.Errorf("string %d: body [%d:%d] out of bounds (%d)", i, idx, end, total) + } + out[i] = string(buf[idx:end]) // copies; Program no longer references buf + idx = end + } + return out, nil +} + +func parseIntsPool(buf []byte) ([]big.Int, error) { + if len(buf) == 0 { + return nil, nil + } + if len(buf) < 4 { + return nil, fmt.Errorf("ints pool: count truncated") + } + n := le.Uint32(buf) + total := uint64(len(buf)) + if uint64(n)*5 > total-4 { // every record is at least a 5B header + return nil, fmt.Errorf("ints pool: count %d exceeds buffer size %d", n, total) + } + idx := uint64(4) + out := make([]big.Int, n) + for i := range out { + if idx+5 > total { + return nil, fmt.Errorf("int %d: header out of bounds", i) + } + sign := buf[idx] + magLen := uint64(le.Uint32(buf[idx+1:])) + idx += 5 + end := idx + magLen + if end > total { + return nil, fmt.Errorf("int %d: magnitude [%d:%d] out of bounds (%d)", i, idx, end, total) + } + out[i].SetBytes(buf[idx:end]) // big-endian, unsigned magnitude + switch sign { + case 0: + // non-negative + case 1: + out[i].Neg(&out[i]) + default: + return nil, fmt.Errorf("int %d: invalid sign byte %d", i, sign) + } + idx = end + } + return out, nil +} + +// PeekProgramVersion checks that buf starts with a valid NUMB header and +// returns its bytecode version, without parsing the sections that follow and +// without checking that this build can read it. +func PeekProgramVersion(buf []byte) (BytecodeVersion, error) { + return peekVersion("NUMB", buf) +} + +func DecodeProgram(buf []byte) (Program, error) { + sections, version, err := decodeSections("NUMB", buf, SectionInstructions, SectionStringsPool, SectionIntsPool, SectionMaxRegisters) + if err != nil { + return Program{}, err + } + + instructions, err := parseInstructions(sections[SectionInstructions]) + if err != nil { + return Program{}, err + } + + stringsPool, err := parseStringsPool(sections[SectionStringsPool]) + if err != nil { + return Program{}, err + } + + intsPool, err := parseIntsPool(sections[SectionIntsPool]) + if err != nil { + return Program{}, err + } + + maxStr, maxInt, maxPortion, maxBool := parseMaxRegs(sections[SectionMaxRegisters]) + + return Program{ + Instructions: instructions, + StringsPool: stringsPool, + IntsPool: intsPool, + MaxRegString: maxStr, + MaxRegInt: maxInt, + MaxRegPortion: maxPortion, + MaxRegBool: maxBool, + Version: version, + }, nil +} diff --git a/internal/vm/program_encode_test.go b/internal/vm/program_encode_test.go new file mode 100644 index 00000000..0979d92b --- /dev/null +++ b/internal/vm/program_encode_test.go @@ -0,0 +1,198 @@ +package vm + +import ( + "math/big" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestProgramEncodeDecodeRoundTrip(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + abc(Op_LoadStr, 0, 1, 2), + bc(Op_LoadInt, 3, 1), + abc(Op_AddInt, 4, 3, 3), + }, + StringsPool: []string{"world", "dest", "USD/2"}, + IntsPool: []big.Int{*big.NewInt(0), *big.NewInt(-42)}, + } + got, err := DecodeProgram(prog.Encode()) + require.NoError(t, err) + require.Equal(t, prog.Instructions, got.Instructions) + require.Equal(t, prog.StringsPool, got.StringsPool) + for i := range prog.IntsPool { + require.Zero(t, got.IntsPool[i].Cmp(&prog.IntsPool[i])) + } +} + +func TestEmptyProgramRoundTrip(t *testing.T) { + _, err := DecodeProgram(Program{}.Encode()) + require.NoError(t, err) +} + +func TestDecodeSkipsUnknownSection(t *testing.T) { + prog := Program{ + Instructions: []Instruction{bc(Op_LoadInt, 0, 0)}, + IntsPool: []big.Int{*big.NewInt(7)}, + } + buf := prog.Encode() + // bump the section count and append an unknown (skippable) section + le.PutUint16(buf[8:], le.Uint16(buf[8:])+1) + buf = appendSection(buf, 0x0999, []byte("future")) + + got, err := DecodeProgram(buf) + require.NoError(t, err) + require.Equal(t, prog.Instructions, got.Instructions) +} + +func TestDecodeRejectsUnknownRequiredSection(t *testing.T) { + buf := Program{}.Encode() + le.PutUint16(buf[8:], le.Uint16(buf[8:])+1) + buf = appendSection(buf, mustUnderstandBit|0x0999, []byte("required")) + + _, err := DecodeProgram(buf) + require.Error(t, err) +} + +func TestDecodeRejectsNewerVersion(t *testing.T) { + buf := Program{}.Encode() + le.PutUint16(buf[6:], CurrentBytecodeVersion.Minor+1) // the minor field + + _, err := DecodeProgram(buf) + require.ErrorAs(t, err, new(UnsupportedBytecodeVersionError)) +} + +func TestDecodeRejectsTruncatedSection(t *testing.T) { + buf := Program{StringsPool: []string{"abc"}}.Encode() + _, err := DecodeProgram(buf[:len(buf)-2]) + require.Error(t, err) +} + +func TestDecodeRejectsDuplicateSection(t *testing.T) { + buf := Program{}.Encode() + le.PutUint16(buf[8:], le.Uint16(buf[8:])+1) + buf = appendSection(buf, SectionStringsPool, nil) + + _, err := DecodeProgram(buf) + require.Error(t, err) +} + +func TestRoundTripEdgeValues(t *testing.T) { + prog := Program{ + StringsPool: []string{"", "héllo", "x"}, + IntsPool: []big.Int{*big.NewInt(0), *new(big.Int).Lsh(big.NewInt(1), 300), *big.NewInt(-1)}, + } + got, err := DecodeProgram(prog.Encode()) + require.NoError(t, err) + require.Equal(t, prog.StringsPool, got.StringsPool) + for i := range prog.IntsPool { + require.Zero(t, got.IntsPool[i].Cmp(&prog.IntsPool[i])) + } +} + +func TestMaxRegRoundTrip(t *testing.T) { + prog := Program{MaxRegString: 3, MaxRegInt: 7, MaxRegPortion: 12, MaxRegBool: 5} + got, err := DecodeProgram(prog.Encode()) + require.NoError(t, err) + require.Equal(t, prog.MaxRegString, got.MaxRegString) + require.Equal(t, prog.MaxRegInt, got.MaxRegInt) + require.Equal(t, prog.MaxRegPortion, got.MaxRegPortion) + require.Equal(t, prog.MaxRegBool, got.MaxRegBool) +} + +func TestMaxRegDefaultsWhenAbsent(t *testing.T) { + var buf []byte + buf = appendFormatHeader(buf, "NUMB", 0) // no sections at all + got, err := DecodeProgram(buf) + require.NoError(t, err) + require.Equal(t, maxRegDefault, got.MaxRegString) + require.Equal(t, maxRegDefault, got.MaxRegInt) + require.Equal(t, maxRegDefault, got.MaxRegPortion) + require.Equal(t, maxRegDefault, got.MaxRegBool) +} + +func TestMaxRegShortSectionDefaultsTrailingToZero(t *testing.T) { + // writer knew only 2 banks: string=3, int=7 + var buf []byte + buf = appendFormatHeader(buf, "NUMB", 1) + buf = appendSection(buf, SectionMaxRegisters, []byte{3, 7}) + + got, err := DecodeProgram(buf) + require.NoError(t, err) + require.Equal(t, byte(3), got.MaxRegString) + require.Equal(t, byte(7), got.MaxRegInt) + require.Equal(t, byte(0), got.MaxRegPortion) // beyond the writer's banks -> 0 + require.Equal(t, byte(0), got.MaxRegBool) +} + +func TestMaxRegExtraTrailingBytesIgnored(t *testing.T) { + // writer knew a 5th bank; this reader ignores the extra bytes + var buf []byte + buf = appendFormatHeader(buf, "NUMB", 1) + buf = appendSection(buf, SectionMaxRegisters, []byte{1, 2, 3, 4, 99}) + + got, err := DecodeProgram(buf) + require.NoError(t, err) + require.Equal(t, byte(1), got.MaxRegString) + require.Equal(t, byte(2), got.MaxRegInt) + require.Equal(t, byte(3), got.MaxRegPortion) + require.Equal(t, byte(4), got.MaxRegBool) +} + +func TestDecodeMalformed(t *testing.T) { + u32 := func(v uint32) []byte { + b := make([]byte, 4) + le.PutUint32(b, v) + return b + } + oneSection := func(tag uint16, content []byte) []byte { + var b []byte + b = appendFormatHeader(b, "NUMB", 1) + return appendSection(b, tag, content) + } + badMagic := Program{}.Encode() + badMagic[0] = 'X' + + cases := map[string][]byte{ + "bad magic": badMagic, + "short buffer": {'N', 'U', 'M'}, + "instructions not mult 4": oneSection(SectionInstructions, []byte{1, 2, 3}), + "string count truncated": oneSection(SectionStringsPool, []byte{0, 0}), + "string count absurd": oneSection(SectionStringsPool, u32(0xFFFFFFFF)), + "string body oob": oneSection(SectionStringsPool, append(u32(1), u32(5)...)), + "int count absurd": oneSection(SectionIntsPool, u32(0xFFFFFFFF)), + "int magnitude oob": oneSection(SectionIntsPool, append(append(u32(1), 0), u32(5)...)), + "int invalid sign": oneSection(SectionIntsPool, append(append(u32(1), 2), u32(0)...)), + } + for name, buf := range cases { + t.Run(name, func(t *testing.T) { + _, err := DecodeProgram(buf) + require.Error(t, err) + }) + } +} + +func FuzzDecodeProgram(f *testing.F) { + f.Add(Program{}.Encode()) + f.Add(Program{ + Instructions: []Instruction{bc(Op_LoadInt, 0, 0)}, + StringsPool: []string{"x"}, + IntsPool: []big.Int{*big.NewInt(1)}, + }.Encode()) + f.Fuzz(func(t *testing.T, data []byte) { + _, _ = DecodeProgram(data) // must not panic on arbitrary input + }) +} + +func TestDecodeMissingPoolIsEmpty(t *testing.T) { + // a program with only an instructions section + var buf []byte + buf = appendFormatHeader(buf, "NUMB", 1) + buf = appendSection(buf, SectionInstructions, []byte{byte(Op_LoadInt), 0, 0, 0}) + + got, err := DecodeProgram(buf) + require.NoError(t, err) + require.Empty(t, got.StringsPool) + require.Empty(t, got.IntsPool) +} diff --git a/internal/vm/section.go b/internal/vm/section.go new file mode 100644 index 00000000..a8c48d0e --- /dev/null +++ b/internal/vm/section.go @@ -0,0 +1,136 @@ +package vm + +import "fmt" + +// BytecodeVersion identifies the bytecode wire format a program or vars blob +// was encoded with, as major.minor — two 16-bit header fields, major first. +// It versions the bytecode format only: the compiler and the library are +// versioned independently of it, and a new compiler release does not imply a +// new bytecode version. +// +// The split carries the compatibility rule (see CanRead). A minor bump is +// additive — new opcodes, new sections, new optional operands — and leaves the +// meaning of everything an older writer could produce untouched, so a 1.1 +// reader runs 1.0 bytecode. A 1.0 reader does not accept 1.1 bytecode: it may +// happen to know every opcode a given blob uses, but that is not assumed. A +// major bump changes the meaning of existing encodings, so a 2.0 reader +// accepts no 1.x blob at all. +type BytecodeVersion struct { + Major uint16 + Minor uint16 +} + +// CurrentBytecodeVersion is the version Encode writes and the newest one the +// decoders read. +var CurrentBytecodeVersion = BytecodeVersion{Major: 1, Minor: 0} + +// CanRead reports whether a reader at version v accepts a blob encoded with +// version encoded: the same major, and a minor no newer than the reader's. +func (v BytecodeVersion) CanRead(encoded BytecodeVersion) bool { + return encoded.Major == v.Major && encoded.Minor <= v.Minor +} + +func (v BytecodeVersion) String() string { + return fmt.Sprintf("%d.%d", v.Major, v.Minor) +} + +// UnsupportedBytecodeVersionError is what the decoders return for a blob this +// build cannot read: another major, or a minor newer than +// CurrentBytecodeVersion. +type UnsupportedBytecodeVersionError struct { + Encoded BytecodeVersion + Supported BytecodeVersion +} + +func (e UnsupportedBytecodeVersionError) Error() string { + return fmt.Sprintf("bytecode version %s is not readable by this build, which reads %d.0 through %s", + e.Encoded, e.Supported.Major, e.Supported) +} + +const ( + SectionInstructions uint16 = 0x01 // NUMB only + SectionStringsPool uint16 = 0x02 + SectionIntsPool uint16 = 0x03 + SectionMaxRegisters uint16 = 0x04 // NUMB only; optional, absent => every bank defaults to maxRegDefault +) + +// A section tag with this bit set must be understood by the decoder: an unknown +// such tag is a hard error rather than a skipped section. +const mustUnderstandBit uint16 = 0x8000 + +// magic(4) + major(2) + minor(2) + section count(2) +const formatHeaderLen = 4 + 2 + 2 + 2 + +func appendFormatHeader(buf []byte, magic string, sectionCount uint16) []byte { + buf = append(buf, magic...) + var h [6]byte + le.PutUint16(h[0:], CurrentBytecodeVersion.Major) + le.PutUint16(h[2:], CurrentBytecodeVersion.Minor) + le.PutUint16(h[4:], sectionCount) + return append(buf, h[:]...) +} + +func appendSection(buf []byte, tag uint16, content []byte) []byte { + var h [6]byte + le.PutUint16(h[0:], tag) + le.PutUint32(h[2:], uint32(len(content))) + buf = append(buf, h[:]...) + return append(buf, content...) +} + +// peekVersion validates the magic and returns the header's version without +// checking that this build can read it, so a caller can report which version +// a blob it cannot read was written with. +func peekVersion(magic string, buf []byte) (BytecodeVersion, error) { + if len(buf) < formatHeaderLen || string(buf[0:4]) != magic { + return BytecodeVersion{}, fmt.Errorf("bad magic (expected %q)", magic) + } + return BytecodeVersion{Major: le.Uint16(buf[4:]), Minor: le.Uint16(buf[6:])}, nil +} + +// decodeSections validates the magic and version, then walks the section list +// into a tag -> content map. Missing sections are simply absent (callers treat +// them as empty). Unknown tags are skipped unless they carry mustUnderstandBit. +// The encoded version is returned alongside, for callers that want to record +// which version a decoded value was written by. +func decodeSections(magic string, buf []byte, knownTags ...uint16) (map[uint16][]byte, BytecodeVersion, error) { + version, err := peekVersion(magic, buf) + if err != nil { + return nil, BytecodeVersion{}, err + } + if !CurrentBytecodeVersion.CanRead(version) { + return nil, BytecodeVersion{}, UnsupportedBytecodeVersionError{Encoded: version, Supported: CurrentBytecodeVersion} + } + + known := make(map[uint16]bool, len(knownTags)) + for _, t := range knownTags { + known[t] = true + } + + count := le.Uint16(buf[8:]) + idx := formatHeaderLen + sections := make(map[uint16][]byte, count) + for i := range count { + if idx+6 > len(buf) { + return nil, BytecodeVersion{}, fmt.Errorf("section %d: header truncated at offset %d", i, idx) + } + tag := le.Uint16(buf[idx:]) + length := le.Uint32(buf[idx+2:]) + idx += 6 + + end := uint64(idx) + uint64(length) + if end > uint64(len(buf)) { + return nil, BytecodeVersion{}, fmt.Errorf("section %d (tag 0x%x): content [%d:%d] exceeds buffer %d", i, tag, idx, end, len(buf)) + } + + if !known[tag] && tag&mustUnderstandBit != 0 { + return nil, BytecodeVersion{}, fmt.Errorf("unknown required section tag 0x%x", tag) + } + if _, dup := sections[tag]; dup { + return nil, BytecodeVersion{}, fmt.Errorf("duplicate section tag 0x%x", tag) + } + sections[tag] = buf[idx:end] + idx = int(end) + } + return sections, version, nil +} diff --git a/internal/vm/vars.go b/internal/vm/vars.go new file mode 100644 index 00000000..a50dd7a7 --- /dev/null +++ b/internal/vm/vars.go @@ -0,0 +1,56 @@ +package vm + +import ( + "math/big" +) + +type Vars struct { + StringsPool []string + IntsPool []big.Int + + // Version is the bytecode version these vars were built for, or (after + // DecodeVars) the version they were actually encoded with. Encode always + // writes CurrentBytecodeVersion regardless of this field. + Version BytecodeVersion +} + +// PeekVarsVersion checks that buf starts with a valid NVAR header and returns +// its bytecode version, without parsing the sections that follow and without +// checking that this build can read it. +func PeekVarsVersion(buf []byte) (BytecodeVersion, error) { + return peekVersion("NVAR", buf) +} + +func DecodeVars(buf []byte) (Vars, error) { + sections, version, err := decodeSections("NVAR", buf, SectionStringsPool, SectionIntsPool) + if err != nil { + return Vars{}, err + } + + stringsPool, err := parseStringsPool(sections[SectionStringsPool]) + if err != nil { + return Vars{}, err + } + + intsPool, err := parseIntsPool(sections[SectionIntsPool]) + if err != nil { + return Vars{}, err + } + + return Vars{ + StringsPool: stringsPool, + IntsPool: intsPool, + Version: version, + }, nil +} + +func (v Vars) Encode() []byte { + strs := encodeStringsPool(v.StringsPool) + ints := encodeIntsPool(v.IntsPool) + + buf := make([]byte, 0, formatHeaderLen+2*6+len(strs)+len(ints)) + buf = appendFormatHeader(buf, "NVAR", 2) + buf = appendSection(buf, SectionStringsPool, strs) + buf = appendSection(buf, SectionIntsPool, ints) + return buf +} diff --git a/internal/vm/vars_test.go b/internal/vm/vars_test.go new file mode 100644 index 00000000..1c439b57 --- /dev/null +++ b/internal/vm/vars_test.go @@ -0,0 +1,112 @@ +package vm + +import ( + "context" + "math/big" + "testing" + + "github.com/formancehq/numscript/internal/funds" + "github.com/stretchr/testify/require" +) + +func TestVarsRoundTrip(t *testing.T) { + in := Vars{ + StringsPool: []string{"alice", "USD/2"}, + IntsPool: []big.Int{*big.NewInt(1), *big.NewInt(4), *big.NewInt(-100)}, + Version: CurrentBytecodeVersion, + } + + out, err := DecodeVars(in.Encode()) + require.NoError(t, err) + require.Equal(t, in, out) +} + +func TestVarsRoundTripEdgeValues(t *testing.T) { + in := Vars{ + StringsPool: []string{"", "héllo", "x"}, + IntsPool: []big.Int{*big.NewInt(0), *new(big.Int).Lsh(big.NewInt(1), 300), *big.NewInt(-1)}, + } + out, err := DecodeVars(in.Encode()) + require.NoError(t, err) + require.Equal(t, in.StringsPool, out.StringsPool) + for i := range in.IntsPool { + require.Zero(t, out.IntsPool[i].Cmp(&in.IntsPool[i])) + } +} + +func TestDecodeVarsMalformed(t *testing.T) { + u32 := func(v uint32) []byte { + b := make([]byte, 4) + le.PutUint32(b, v) + return b + } + oneSection := func(tag uint16, content []byte) []byte { + var b []byte + b = appendFormatHeader(b, "NVAR", 1) + return appendSection(b, tag, content) + } + badMagic := Vars{}.Encode() + badMagic[0] = 'X' + + newerVersion := Vars{}.Encode() + le.PutUint16(newerVersion[6:], CurrentBytecodeVersion.Minor+1) // the minor field + + cases := map[string][]byte{ + "bad magic": badMagic, + "short buffer": {'N', 'V', 'A'}, + "newer version": newerVersion, + "string count truncated": oneSection(SectionStringsPool, []byte{0, 0}), + "string count absurd": oneSection(SectionStringsPool, u32(0xFFFFFFFF)), + "string body oob": oneSection(SectionStringsPool, append(u32(1), u32(5)...)), + "int count absurd": oneSection(SectionIntsPool, u32(0xFFFFFFFF)), + "int magnitude oob": oneSection(SectionIntsPool, append(append(u32(1), 0), u32(5)...)), + "int invalid sign": oneSection(SectionIntsPool, append(append(u32(1), 2), u32(0)...)), + } + for name, buf := range cases { + t.Run(name, func(t *testing.T) { + _, err := DecodeVars(buf) + require.Error(t, err) + }) + } +} + +func FuzzDecodeVars(f *testing.F) { + f.Add(Vars{}.Encode()) + f.Add(Vars{ + StringsPool: []string{"x"}, + IntsPool: []big.Int{*big.NewInt(1)}, + }.Encode()) + f.Fuzz(func(t *testing.T, data []byte) { + _, _ = DecodeVars(data) // must not panic on arbitrary input + }) +} + +func TestLoadVarOpcodes(t *testing.T) { + vars, err := DecodeVars(Vars{ + StringsPool: []string{"world", "dest"}, + IntsPool: []big.Int{*big.NewInt(42)}, + }.Encode()) + require.NoError(t, err) + + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadStr, sUSD, 0), // r_s0 = "USD/2" (current asset) + abc(Op_SetCurrentAsset, sUSD, 0, 0), + bc(Op_LoadVarStr, 1, 0), // r_s1 = var strings[0] = "world" + bc(Op_LoadVarStr, 2, 1), // r_s2 = var strings[1] = "dest" + bc(Op_LoadVarInt, 0, 0), // r_i0 = var ints[0] = 42 + abc(Op_PullAccount, 1, 1, 0), // r_i1 = pull(world, cap r_i0) + abc(0, nilReg, nilReg, nilReg), // ext: no overdraft, no color, no scope + abc(Op_SendToAccount, 2, nilReg, nilReg), // send to dest + }, + StringsPool: []string{"USD/2"}, + } + + res, execErr := Exec(context.Background(), newTestVm(prog), &vars, mockStore{}) + require.Nil(t, execErr) + + want := []funds.Posting{ + {Source: "world", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(42)}, + } + require.Equal(t, want, res.Postings) +} diff --git a/internal/vm/verify.go b/internal/vm/verify.go new file mode 100644 index 00000000..11b4bdcb --- /dev/null +++ b/internal/vm/verify.go @@ -0,0 +1,652 @@ +package vm + +import "fmt" + +// regBank identifies which of the VM's register banks an operand indexes. The +// bank is never encoded in the instruction — it is implied by the opcode — so the +// verifier has to reconstruct it from the same table Exec switches on. +type regBank int + +const ( + bankInt regBank = iota + bankStr + bankPortion + bankBool + // bankCurrentAsset is a pseudo-bank with a single slot, tracking whether the + // current asset has been set. It is never allocated as a register. + bankCurrentAsset +) + +func (b regBank) String() string { + switch b { + case bankInt: + return "int" + case bankStr: + return "string" + case bankPortion: + return "portion" + case bankBool: + return "bool" + case bankCurrentAsset: + return "current asset" + default: + return "?" + } +} + +var currentAssetRef = regRef{bankCurrentAsset, 0} + +type regRef struct { + bank regBank + index int +} + +func (r regRef) String() string { + if r.bank == bankCurrentAsset { + return "current asset" + } + return fmt.Sprintf("%s register %d", r.bank, r.index) +} + +// decoded describes the register/pool/jump operands an instruction touches, so +// the verifier can check every access the execution loop will make without the +// VM having to guard anything at run time. +type decoded struct { + reads []regRef + writes []regRef + constInt int // index into Program.IntsPool, or -1 + constStr int + varInt int // index into Vars.IntsPool, or -1 + varStr int + // jumpDelta is the unsigned forward offset from the instruction *after* this + // one, or -1 when the instruction does not jump + jumpDelta int + // noFallThrough marks an unconditional jump, whose successor is its target + // only. Giving it a fallthrough edge too would intersect the assigned-set + // with a path that cannot be taken, rejecting valid programs. + noFallThrough bool + // markDelta is +1 for a mark push, -1 for a mark end, 0 otherwise + markDelta int + // notInMark marks an op Exec refuses while a mark is open + notInMark bool +} + +type programInfo struct { + varIntsLen int + varStrsLen int +} + +// Verify statically checks that a program is safe to execute. It is opt-in: Exec +// assumes well-formed bytecode (the compiler's output always is) and does not +// call this. Run it on any program that did not come out of ir.Assemble in this +// process — anything decoded from a file or a wire. +// +// A nil result guarantees the execution loop cannot read out of bounds: no +// truncated instruction, no jump into the middle of one, no out-of-range pool or +// register index, and no read of a register that was not written on every path +// reaching it. +// +// It does not check that vars were supplied; for that see VerifyWithVars. +func Verify(p Program) error { + _, err := verify(p) + return err +} + +// VerifiedVarsInfo records, opaquely, the variable-pool sizes a Program was +// found to require by a successful VerifyWithVars call. A caller that keeps +// the Program around — e.g. one compiled artifact reused across many calls +// with (normally) the same Vars shape — can hold on to this instead of the +// Program's raw pool-size requirements, and later ask CheckVars whether a +// new Vars value would still satisfy VerifyWithVars, without paying for the +// static pass again. The fields are deliberately unexported: what "shape" +// means is this package's business, not a caller's. +type VerifiedVarsInfo struct { + varIntsLen int + varStrsLen int +} + +// CheckVars reports, in O(1) and without re-running verification, whether +// vars is guaranteed to satisfy the Program this VerifiedVarsInfo was +// obtained from — i.e. whether VerifyWithVars(thatProgram, vars) would +// succeed, without calling it. A false result is not itself a rejection: it +// only means the caller must call VerifyWithVars to get an authoritative +// answer (and, on failure, a precise error). +func (info VerifiedVarsInfo) CheckVars(vars *Vars) bool { + if vars == nil { + return info.varIntsLen == 0 && info.varStrsLen == 0 + } + return len(vars.IntsPool) >= info.varIntsLen && len(vars.StringsPool) >= info.varStrsLen +} + +// VerifyWithVars is Verify plus the check that vars carries every variable the +// program loads. A nil *Vars is only legal for a program that reads none. On +// success it also returns the VerifiedVarsInfo backing that check; see its +// doc for why a caller would want to keep it. +func VerifyWithVars(p Program, vars *Vars) (VerifiedVarsInfo, error) { + info, err := verify(p) + if err != nil { + return VerifiedVarsInfo{}, err + } + + result := VerifiedVarsInfo(info) + if result.CheckVars(vars) { + return result, nil + } + + if vars == nil { + return VerifiedVarsInfo{}, fmt.Errorf("program reads variables but none were provided") + } + if len(vars.IntsPool) < info.varIntsLen { + return VerifiedVarsInfo{}, fmt.Errorf("program reads int var %d but only %d were provided", info.varIntsLen-1, len(vars.IntsPool)) + } + return VerifiedVarsInfo{}, fmt.Errorf("program reads string var %d but only %d were provided", info.varStrsLen-1, len(vars.StringsPool)) +} + +// instrWords is how many 4-byte words an opcode occupies. The ones returning 2 +// carry an "ext" word whose opcode byte is ignored; Exec reads it with +// `instrs[pc]; pc++`, which is what makes a truncated tail or a jump landing on +// an ext word a crash rather than an error. +func instrWords(op byte) int { + switch Opcode(op) { + case Op_PullAccount, + Op_Save, + Op_SetAccountMeta, + Op_MetaStr, + Op_MetaInt, + Op_MetaPortion, + Op_MetaMonetary, + Op_Balance: + return 2 + default: + return 1 + } +} + +func (p Program) maxReg(bank regBank) int { + switch bank { + case bankInt: + return int(p.MaxRegInt) + case bankStr: + return int(p.MaxRegString) + case bankPortion: + return int(p.MaxRegPortion) + case bankBool: + return int(p.MaxRegBool) + default: + return 1 // the current-asset pseudo-bank has exactly one slot + } +} + +func verify(p Program) (programInfo, error) { + instrs := p.Instructions + n := len(instrs) + + // boundary[i] is true when i starts an instruction; the ext word of a + // two-word instruction is deliberately not a boundary + boundary := make([]bool, n+1) + boundary[n] = true // jumping past the last instruction halts, which is fine + + var steps []step + + for i := 0; i < n; { + op := instrs[i].Opcode + w := instrWords(op) + if i+w > n { + return programInfo{}, fmt.Errorf("truncated instruction at %d: opcode 0x%02X needs %d words", i, op, w) + } + var ext Instruction + if w == 2 { + ext = instrs[i+1] + } + d, err := decodeInstr(instrs[i], ext) + if err != nil { + return programInfo{}, fmt.Errorf("at instruction %d: %w", i, err) + } + boundary[i] = true + steps = append(steps, step{at: i, words: w, d: d}) + i += w + } + + info := programInfo{} + for _, st := range steps { + d := st.d + + for _, r := range append(append([]regRef{}, d.reads...), d.writes...) { + if r.bank == bankCurrentAsset { + continue + } + if r.index >= p.maxReg(r.bank) { + return programInfo{}, fmt.Errorf( + "at instruction %d: %s is beyond the program's declared %d %s registers", + st.at, r, p.maxReg(r.bank), r.bank) + } + } + + if d.constInt >= len(p.IntsPool) { + return programInfo{}, fmt.Errorf("at instruction %d: int constant %d out of range (pool size %d)", st.at, d.constInt, len(p.IntsPool)) + } + if d.constStr >= len(p.StringsPool) { + return programInfo{}, fmt.Errorf("at instruction %d: string constant %d out of range (pool size %d)", st.at, d.constStr, len(p.StringsPool)) + } + if d.varInt >= info.varIntsLen { + info.varIntsLen = d.varInt + 1 + } + if d.varStr >= info.varStrsLen { + info.varStrsLen = d.varStr + 1 + } + + // deltas are unsigned and relative to the following instruction, so a + // backward jump cannot be encoded; the hazard is landing on an ext word + if t := st.target(); t >= 0 && t < n && !boundary[t] { + return programInfo{}, fmt.Errorf("at instruction %d: jump to %d lands inside an instruction", st.at, t) + } + } + + g := buildCFG(steps) + if err := checkDefiniteAssignment(steps, g); err != nil { + return programInfo{}, err + } + if err := checkMarkBalance(steps, g); err != nil { + return programInfo{}, err + } + + return info, nil +} + +// step is one decoded instruction, tagged with where it sits in the stream. +type step struct { + at int + words int + d decoded +} + +// target is the absolute instruction index this step jumps to, or -1 when it +// does not jump. A target of len(instrs) or beyond halts, which is legal. +func (s step) target() int { + if s.d.jumpDelta < 0 { + return -1 + } + return s.at + s.words + s.d.jumpDelta +} + +// cfg is the control-flow graph over steps: edges are fallthroughs and jump +// targets, indexed by position in steps. +type cfg struct { + preds [][]int + reachable []bool + // exits[k] is true when step k can end the run: by falling off the last + // instruction or jumping to len(instrs) or beyond + exits []bool +} + +// Jumps are forward-only, so every predecessor of a step is earlier in the +// stream: one ordered pass over the steps sees each step after all of its +// predecessors. +func buildCFG(steps []step) cfg { + at2idx := make(map[int]int, len(steps)) + for k, st := range steps { + at2idx[st.at] = k + } + + g := cfg{ + preds: make([][]int, len(steps)), + reachable: make([]bool, len(steps)), + exits: make([]bool, len(steps)), + } + edge := func(k, at int) { + if j, ok := at2idx[at]; ok { + g.preds[j] = append(g.preds[j], k) + } else { + g.exits[k] = true + } + } + for k, st := range steps { + if !st.d.noFallThrough { + edge(k, st.at+st.words) + } + if t := st.target(); t >= 0 { + edge(k, t) + } + } + + // An unreachable step never executes, so its reads cannot crash and its + // (empty) assigned-set must not poison the intersection at a later join. + if len(steps) > 0 { + g.reachable[0] = true + } + for k := range steps { + for _, p := range g.preds[k] { + if g.reachable[p] { + g.reachable[k] = true + break + } + } + } + return g +} + +// checkDefiniteAssignment rejects a read of a register (or of the current asset) +// that was not written on every path reaching that instruction. +// +// This subsumes bytecode-level type confusion: a bank is part of a regRef, so a +// slot written as an int and later read as a string is a read of a register that +// was never written. +func checkDefiniteAssignment(steps []step, g cfg) error { + assignedOut := make([]map[regRef]bool, len(steps)) + for k, st := range steps { + if !g.reachable[k] { + continue + } + in := intersectAssigned(assignedOut, filterReachable(g.preds[k], g.reachable)) + for _, r := range st.d.reads { + if !in[r] { + return fmt.Errorf("at instruction %d: %s read before being assigned on all paths", st.at, r) + } + } + for _, r := range st.d.writes { + in[r] = true + } + assignedOut[k] = in + } + return nil +} + +// checkMarkBalance rejects a program in which the number of open marks is not +// the same on every path reaching an instruction, a mark end runs with no open +// mark, a send, save or set_current_asset runs with a mark open, or the run can +// end with a mark open. +func checkMarkBalance(steps []step, g cfg) error { + depthOut := make([]int, len(steps)) + for k, st := range steps { + if !g.reachable[k] { + continue + } + in := 0 + preds := filterReachable(g.preds[k], g.reachable) + for i, p := range preds { + if i == 0 { + in = depthOut[p] + } else if depthOut[p] != in { + return fmt.Errorf("at instruction %d: paths join with different numbers of open marks (%d and %d)", st.at, in, depthOut[p]) + } + } + if st.d.notInMark && in > 0 { + return fmt.Errorf("at instruction %d: send, save or set_current_asset while a mark is open", st.at) + } + out := in + st.d.markDelta + if out < 0 { + return fmt.Errorf("at instruction %d: mark end with no open mark", st.at) + } + if g.exits[k] && out != 0 { + return fmt.Errorf("at instruction %d: the run can end with %d open marks", st.at, out) + } + depthOut[k] = out + } + return nil +} + +func filterReachable(preds []int, reachable []bool) []int { + out := preds[:0:0] + for _, p := range preds { + if reachable[p] { + out = append(out, p) + } + } + return out +} + +func intersectAssigned(out []map[regRef]bool, preds []int) map[regRef]bool { + res := map[regRef]bool{} + if len(preds) == 0 { + return res + } + for r := range out[preds[0]] { + res[r] = true + } + for _, p := range preds[1:] { + for r := range res { + if !out[p][r] { + delete(res, r) + } + } + } + return res +} + +// decodeInstr mirrors, operand for operand, what the matching arm of Exec reads +// and writes. Every opcode in instruction.go must have a case here: an opcode +// missing from this switch is rejected as unknown, which is what keeps the two +// in step. +func decodeInstr(instr, ext Instruction) (decoded, error) { + d := decoded{constInt: -1, constStr: -1, varInt: -1, varStr: -1, jumpDelta: -1} + + var err error + // read/write name an operand Exec dereferences unconditionally, so nilReg + // there is a malformed instruction rather than an absent operand + read := func(bank regBank, idx byte) { + if idx == nilReg { + err = fmt.Errorf("%s operand is the nil register, but this operand is not optional", bank) + return + } + d.reads = append(d.reads, regRef{bank, int(idx)}) + } + write := func(bank regBank, idx byte) { + if idx == nilReg { + err = fmt.Errorf("%s destination is the nil register", bank) + return + } + d.writes = append(d.writes, regRef{bank, int(idx)}) + } + readOpt := func(bank regBank, idx byte) { + if idx != nilReg { + d.reads = append(d.reads, regRef{bank, int(idx)}) + } + } + flag := func(v byte) { + if v > 1 { + err = fmt.Errorf("flag operand is %d, expected 0 or 1", v) + } + } + + switch Opcode(instr.Opcode) { + // --- state & assertions + case Op_SetCurrentAsset: + read(bankStr, instr.A) + d.writes = append(d.writes, currentAssetRef) + d.notInMark = true + case Op_AssertSameAsset: + read(bankStr, instr.A) + read(bankStr, instr.B) + case Op_AssertValidAccount, Op_AssertValidColor, Op_AssertValidScope: + read(bankStr, instr.A) + case Op_AssertNonNegativeBalance: + read(bankInt, instr.A) + read(bankStr, instr.B) + case Op_AssertUnscoped: + read(bankStr, instr.A) + read(bankStr, instr.B) + case Op_AssertNonNegativeAmount: + read(bankInt, instr.A) + case Op_AssertNonNegativePortion: + read(bankPortion, instr.A) + case Op_AssertLeftover: + read(bankPortion, instr.A) + flag(instr.B) + case Op_CheckEnoughFunds: + read(bankInt, instr.A) + read(bankInt, instr.B) + // the asset names the MissingFundsError this may raise + d.reads = append(d.reads, currentAssetRef) + + // --- constants & variables + case Op_LoadInt: + d.constInt = int(instr.GetBC()) + write(bankInt, instr.A) + case Op_LoadStr: + d.constStr = int(instr.GetBC()) + write(bankStr, instr.A) + case Op_LoadVarInt: + d.varInt = int(instr.GetBC()) + write(bankInt, instr.A) + case Op_LoadVarStr: + d.varStr = int(instr.GetBC()) + write(bankStr, instr.A) + case Op_ConstTrue, Op_ConstFalse: + write(bankBool, instr.A) + + // --- metadata + case Op_SetTxMeta: + read(bankStr, instr.A) + read(bankStr, instr.B) + case Op_SetAccountMeta: + read(bankStr, instr.A) + read(bankStr, instr.B) + read(bankStr, instr.C) + readOpt(bankStr, ext.A) // scope + case Op_MetaStr: + read(bankStr, instr.B) + read(bankStr, instr.C) + readOpt(bankStr, ext.A) + write(bankStr, instr.A) + case Op_MetaInt: + read(bankStr, instr.B) + read(bankStr, instr.C) + readOpt(bankStr, ext.A) + write(bankInt, instr.A) + case Op_MetaPortion: + read(bankStr, instr.B) + read(bankStr, instr.C) + readOpt(bankStr, ext.A) + write(bankPortion, instr.A) + case Op_MetaMonetary: + read(bankStr, instr.B) + read(bankStr, instr.C) + readOpt(bankStr, ext.B) // scope + write(bankStr, instr.A) // asset + write(bankInt, ext.A) // amount + case Op_Balance: + read(bankStr, instr.B) + read(bankStr, instr.C) + readOpt(bankStr, ext.A) // scope + write(bankInt, instr.A) + + // --- arithmetic & constructors + case Op_AddInt, Op_SubInt: + read(bankInt, instr.B) + read(bankInt, instr.C) + write(bankInt, instr.A) + case Op_AddPortion, Op_SubPortion, Op_MulPortion: + read(bankPortion, instr.B) + read(bankPortion, instr.C) + write(bankPortion, instr.A) + case Op_MkPortion: + read(bankInt, instr.B) + read(bankInt, instr.C) + write(bankPortion, instr.A) + case Op_AddString: + read(bankStr, instr.B) + read(bankStr, instr.C) + write(bankStr, instr.A) + + // --- unary & conversions + case Op_IntCopy, Op_NegInt: + read(bankInt, instr.B) + write(bankInt, instr.A) + case Op_PortionCopy: + read(bankPortion, instr.B) + write(bankPortion, instr.A) + case Op_StrCopy: + read(bankStr, instr.B) + write(bankStr, instr.A) + case Op_BoolCopy: + read(bankBool, instr.B) + write(bankBool, instr.A) + case Op_IntToString: + read(bankInt, instr.B) + write(bankStr, instr.A) + case Op_PortionToString: + read(bankPortion, instr.B) + write(bankStr, instr.A) + case Op_MonetaryToString: + read(bankStr, instr.B) + read(bankInt, instr.C) + write(bankStr, instr.A) + case Op_IntToPortion: + read(bankInt, instr.B) + write(bankPortion, instr.A) + case Op_PortionToInt: + read(bankPortion, instr.B) + write(bankInt, instr.A) + + // --- funds & postings + case Op_PullAccount: + read(bankStr, instr.B) // account + readOpt(bankInt, instr.C) // cap + readOpt(bankInt, ext.A) // overdraft + readOpt(bankStr, ext.B) // color + readOpt(bankStr, ext.C) // scope + d.reads = append(d.reads, currentAssetRef) + write(bankInt, instr.A) + case Op_SendToAccount: + readOpt(bankStr, instr.A) // destination + readOpt(bankInt, instr.B) // cap + readOpt(bankStr, instr.C) // scope + d.reads = append(d.reads, currentAssetRef) + d.notInMark = true + case Op_Save: + read(bankStr, instr.A) // account + read(bankStr, instr.B) // asset + readOpt(bankInt, instr.C) // amount, nil = save all + readOpt(bankStr, ext.A) // scope + d.notInMark = true + + // --- marks + case Op_MarkPush: + d.markDelta = 1 + case Op_MarkEnd: + flag(instr.A) + d.markDelta = -1 + // a rewind repays queued sources into the current asset's balance, but it + // is not listed as a reader: the repay loop only runs over sources queued + // inside the region, and queueing one takes an Op_PullAccount, which + // already requires the asset. A region that pulled nothing reads nothing. + + // --- comparisons + case Op_LtInt, Op_EqInt: + read(bankInt, instr.B) + read(bankInt, instr.C) + write(bankBool, instr.A) + case Op_LtPortion, Op_EqPortion: + read(bankPortion, instr.B) + read(bankPortion, instr.C) + write(bankBool, instr.A) + case Op_StrEq: + read(bankStr, instr.B) + read(bankStr, instr.C) + write(bankBool, instr.A) + case Op_IsZero: + read(bankInt, instr.B) + write(bankBool, instr.A) + + // --- bool ops + case Op_Not: + read(bankBool, instr.B) + write(bankBool, instr.A) + + // --- control flow + case Op_JmpIfFalse, Op_JmpIfTrue: + read(bankBool, instr.A) + d.jumpDelta = int(instr.GetBC()) + case Op_Jmp: + d.jumpDelta = int(instr.GetBC()) + d.noFallThrough = true + + default: + return decoded{}, fmt.Errorf("unknown opcode 0x%02X", instr.Opcode) + } + + if err != nil { + return decoded{}, err + } + return d, nil +} diff --git a/internal/vm/verify_test.go b/internal/vm/verify_test.go new file mode 100644 index 00000000..b5f77d37 --- /dev/null +++ b/internal/vm/verify_test.go @@ -0,0 +1,260 @@ +package vm + +import ( + "math/big" + "testing" + + "github.com/stretchr/testify/require" +) + +func mustReject(t *testing.T, p Program) { + t.Helper() + require.Error(t, Verify(fullBanks(p))) +} + +func mustAccept(t *testing.T, p Program) { + t.Helper() + require.NoError(t, Verify(fullBanks(p))) +} + +func TestVerifyUnknownOpcode(t *testing.T) { + mustReject(t, Program{Instructions: []Instruction{abc(0xFE, 0, 0, 0)}}) +} + +func TestVerifyTruncatedMultiWord(t *testing.T) { + // each of these needs an ext word the stream doesn't have + for _, op := range []Opcode{ + Op_PullAccount, + Op_Save, + Op_SetAccountMeta, + Op_MetaStr, + Op_MetaInt, + Op_MetaPortion, + Op_MetaMonetary, + Op_Balance, + } { + mustReject(t, Program{ + Instructions: []Instruction{bc(Op_LoadStr, 0, 0), abc(op, 0, 0, 0)}, + StringsPool: []string{"acc"}, + }) + } +} + +func TestVerifyConstIndexOutOfRange(t *testing.T) { + mustReject(t, Program{Instructions: []Instruction{bc(Op_LoadInt, 0, 3)}}) + mustReject(t, Program{Instructions: []Instruction{bc(Op_LoadStr, 0, 3)}}) +} + +func TestVerifyJumpPastEndIsFine(t *testing.T) { + // pc runs off the end and the loop stops; nothing to guard + mustAccept(t, Program{Instructions: []Instruction{ + abc(Op_ConstTrue, 0, nilReg, nilReg), + bc(Op_JmpIfTrue, 0, 99), + }}) +} + +func TestVerifyJumpIntoExtWord(t *testing.T) { + // instruction 2 is Op_Balance's ext word, not an instruction boundary + mustReject(t, Program{ + Instructions: []Instruction{ + abc(Op_ConstTrue, 0, nilReg, nilReg), + bc(Op_JmpIfTrue, 0, 1), + abc(Op_Balance, 0, 0, 0), + abc(0, nilReg, nilReg, nilReg), // ext + }, + StringsPool: []string{"acc"}, + }) +} + +func TestVerifyReadNotAssignedOnAllPaths(t *testing.T) { + // int reg 1 is written only on the fall-through path, then read at the join + mustReject(t, Program{ + Instructions: []Instruction{ + abc(Op_ConstTrue, 0, nilReg, nilReg), // 0: b0 = true + bc(Op_JmpIfTrue, 0, 1), // 1: skip instruction 2 + bc(Op_LoadInt, 1, 0), // 2: i1 = 0 (skipped when jumping) + abc(Op_NegInt, 2, 1, nilReg), // 3: i2 = -i1 (i1 maybe unassigned) + }, + IntsPool: []big.Int{*big.NewInt(0)}, + }) +} + +func TestVerifyAssignedOnBothBranchesIsFine(t *testing.T) { + // the same register written on either side of a diamond is assigned at the join + mustAccept(t, Program{ + Instructions: []Instruction{ + abc(Op_ConstTrue, 0, nilReg, nilReg), // 0: b0 = true + bc(Op_JmpIfTrue, 0, 2), // 1: -> 4 + bc(Op_LoadInt, 1, 0), // 2: i1 = 0 + bc(Op_Jmp, 0, 1), // 3: -> 5 + bc(Op_LoadInt, 1, 0), // 4: i1 = 0 + abc(Op_NegInt, 2, 1, nilReg), // 5: i2 = -i1 + }, + IntsPool: []big.Int{*big.NewInt(0)}, + }) +} + +// A type confusion at the bytecode level shows up as a read of a register that +// was never written: the bank is part of a register's identity, so string 0 and +// int 0 are different registers. +func TestVerifyTypeConfusionIsAnUnassignedRead(t *testing.T) { + mustReject(t, Program{ + Instructions: []Instruction{ + bc(Op_LoadStr, 0, 0), // strings[0] = "x" + abc(Op_NegInt, 1, 0, nilReg), // reads ints[0], never written + }, + StringsPool: []string{"x"}, + }) +} + +func TestVerifyCurrentAssetNotSet(t *testing.T) { + mustReject(t, Program{ + Instructions: []Instruction{ + bc(Op_LoadStr, 0, 0), + abc(Op_SendToAccount, 0, nilReg, nilReg), + }, + StringsPool: []string{"dest"}, + }) +} + +func TestVerifyRegisterBeyondDeclaredMax(t *testing.T) { + p := Program{ + Instructions: []Instruction{bc(Op_LoadInt, 5, 0)}, + IntsPool: []big.Int{*big.NewInt(1)}, + } + + p.MaxRegInt = 5 // registers 0..4, so 5 is out + require.Error(t, Verify(p)) + + p.MaxRegInt = 6 + require.NoError(t, Verify(p)) +} + +func TestVerifyNilRegInNonOptionalOperand(t *testing.T) { + // Op_SetCurrentAsset always dereferences A, so nilReg there is malformed + mustReject(t, Program{Instructions: []Instruction{ + abc(Op_SetCurrentAsset, nilReg, nilReg, nilReg), + }}) +} + +func TestVerifyNilRegInOptionalOperandIsFine(t *testing.T) { + // Op_Save's amount is optional: nilReg means "save all" + mustAccept(t, Program{ + Instructions: []Instruction{ + bc(Op_LoadStr, 0, 0), + bc(Op_LoadStr, 1, 1), + abc(Op_Save, 0, 1, nilReg), + abc(0, nilReg, nilReg, nilReg), // ext: no scope + }, + StringsPool: []string{"acc", "USD/2"}, + }) +} + +func TestVerifyFlagOperand(t *testing.T) { + // Exec tests `instr.A == 1`, so a 2 would silently commit a region meant to + // rewind + mustReject(t, Program{Instructions: []Instruction{ + abc(Op_MarkPush, nilReg, nilReg, nilReg), + abc(Op_MarkEnd, 2, nilReg, nilReg), + }}) + mustAccept(t, Program{Instructions: []Instruction{ + abc(Op_MarkPush, nilReg, nilReg, nilReg), + abc(Op_MarkEnd, 1, nilReg, nilReg), + }}) +} + +func requireRejectedWith(t *testing.T, p Program, msg string) { + t.Helper() + require.ErrorContains(t, Verify(fullBanks(p)), msg) +} + +func TestVerifyMarkBalance(t *testing.T) { + push := abc(Op_MarkPush, nilReg, nilReg, nilReg) + commit := abc(Op_MarkEnd, 0, nilReg, nilReg) + rewind := abc(Op_MarkEnd, 1, nilReg, nilReg) + + t.Run("end with no open mark", func(t *testing.T) { + requireRejectedWith(t, Program{Instructions: []Instruction{commit}}, "mark end with no open mark") + requireRejectedWith(t, Program{Instructions: []Instruction{push, rewind, commit}}, "mark end with no open mark") + }) + + t.Run("mark open at the end", func(t *testing.T) { + requireRejectedWith(t, Program{Instructions: []Instruction{push}}, "end with 1 open marks") + }) + + t.Run("mark open when jumping past the end", func(t *testing.T) { + requireRejectedWith(t, Program{Instructions: []Instruction{ + abc(Op_ConstTrue, 0, nilReg, nilReg), // 0 + push, // 1 + bc(Op_JmpIfTrue, 0, 1), // 2: -> 4, past the end + commit, // 3 + }}, "end with 1 open marks") + }) + + t.Run("paths join with different depths", func(t *testing.T) { + requireRejectedWith(t, Program{Instructions: []Instruction{ + abc(Op_ConstTrue, 0, nilReg, nilReg), // 0 + bc(Op_JmpIfTrue, 0, 1), // 1: -> 3, skipping the push + push, // 2 + commit, // 3 + }}, "different numbers of open marks") + }) + + // the oneof shape: a branch that covers the amount jumps out with its mark + // still open, and the commit after the join closes it + t.Run("oneof shape is fine", func(t *testing.T) { + mustAccept(t, Program{Instructions: []Instruction{ + abc(Op_ConstTrue, 0, nilReg, nilReg), // 0 + push, // 1 + bc(Op_JmpIfTrue, 0, 2), // 2: -> 5 + rewind, // 3 + push, // 4 + commit, // 5 + }}) + }) + + t.Run("unreachable mark end is ignored", func(t *testing.T) { + mustAccept(t, Program{Instructions: []Instruction{ + bc(Op_Jmp, 0, 1), // 0: -> 2 + commit, // 1 + }}) + }) +} + +func TestVerifyWithVars(t *testing.T) { + p := fullBanks(Program{Instructions: []Instruction{ + bc(Op_LoadVarInt, 0, 1), // reads int var 1, so 2 are needed + }}) + + require.NoError(t, Verify(p), "Verify alone says nothing about vars") + + _, err := VerifyWithVars(p, nil) + require.Error(t, err) + + _, err = VerifyWithVars(p, &Vars{IntsPool: []big.Int{*big.NewInt(0)}}) + require.Error(t, err) + + info, err := VerifyWithVars(p, &Vars{IntsPool: []big.Int{*big.NewInt(0), *big.NewInt(1)}}) + require.NoError(t, err) + require.True(t, info.CheckVars(&Vars{IntsPool: []big.Int{*big.NewInt(0), *big.NewInt(1)}}), + "the exact vars just verified must be reported sufficient") + require.True(t, info.CheckVars(&Vars{IntsPool: []big.Int{*big.NewInt(0), *big.NewInt(1), *big.NewInt(2)}}), + "a larger pool still satisfies the same requirement") + require.False(t, info.CheckVars(&Vars{IntsPool: []big.Int{*big.NewInt(0)}}), + "a smaller pool no longer satisfies the requirement") + require.False(t, info.CheckVars(nil), + "nil is only sufficient when the program reads no vars") +} + +func TestVerifyWithVarsAllowsNilWhenNoneAreRead(t *testing.T) { + info, err := VerifyWithVars(fullBanks(Program{ + Instructions: []Instruction{bc(Op_LoadInt, 0, 0)}, + IntsPool: []big.Int{*big.NewInt(1)}, + }), nil) + require.NoError(t, err) + require.True(t, info.CheckVars(nil)) +} + +func TestVerifyEmptyProgram(t *testing.T) { + mustAccept(t, Program{}) +} diff --git a/internal/vm/version_test.go b/internal/vm/version_test.go new file mode 100644 index 00000000..aa41ff23 --- /dev/null +++ b/internal/vm/version_test.go @@ -0,0 +1,160 @@ +package vm + +import ( + "bytes" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestBytecodeVersionCanRead(t *testing.T) { + cases := []struct { + reader, encoded BytecodeVersion + ok bool + }{ + {BytecodeVersion{1, 0}, BytecodeVersion{1, 0}, true}, + {BytecodeVersion{1, 1}, BytecodeVersion{1, 0}, true}, // newer reader runs older minor + {BytecodeVersion{1, 0}, BytecodeVersion{1, 1}, false}, // older reader does not assume it knows the new opcodes + {BytecodeVersion{2, 0}, BytecodeVersion{1, 0}, false}, // major bump: meanings changed + {BytecodeVersion{2, 0}, BytecodeVersion{1, 9}, false}, + {BytecodeVersion{1, 9}, BytecodeVersion{2, 0}, false}, + {BytecodeVersion{2, 0}, BytecodeVersion{2, 0}, true}, + {BytecodeVersion{1, 300}, BytecodeVersion{1, 299}, true}, // minor is a full u16, not a byte + } + for _, tc := range cases { + t.Run(tc.reader.String()+" reads "+tc.encoded.String(), func(t *testing.T) { + require.Equal(t, tc.ok, tc.reader.CanRead(tc.encoded)) + }) + } +} + +// The header layout is pinned: the 4-byte magic, then major and minor as two +// little-endian u16 fields, then the u16 section count — 10 bytes in all. +func TestHeaderLayout(t *testing.T) { + require.Equal(t, 10, formatHeaderLen) + + program := Program{}.Encode() + require.Equal(t, "NUMB", string(program[:4])) + require.Equal(t, CurrentBytecodeVersion.Major, le.Uint16(program[4:])) + require.Equal(t, CurrentBytecodeVersion.Minor, le.Uint16(program[6:])) + require.Equal(t, uint16(4), le.Uint16(program[8:]), "Program.Encode writes four sections") + + vars := Vars{}.Encode() + require.Equal(t, "NVAR", string(vars[:4])) + require.Equal(t, CurrentBytecodeVersion.Major, le.Uint16(vars[4:])) + require.Equal(t, CurrentBytecodeVersion.Minor, le.Uint16(vars[6:])) + require.Equal(t, uint16(2), le.Uint16(vars[8:]), "Vars.Encode writes two sections") + + require.Equal(t, "2.1", BytecodeVersion{2, 1}.String()) +} + +// Encode stamps CurrentBytecodeVersion whatever the struct's own Version says, +// and both the peek and the decode read it back. +func TestEncodeWritesCurrentVersion(t *testing.T) { + stale := BytecodeVersion{9, 9} + + programBytes := Program{Version: stale}.Encode() + peeked, err := PeekProgramVersion(programBytes) + require.NoError(t, err) + require.Equal(t, CurrentBytecodeVersion, peeked) + + program, err := DecodeProgram(programBytes) + require.NoError(t, err) + require.Equal(t, CurrentBytecodeVersion, program.Version) + + varsBytes := Vars{Version: stale}.Encode() + peeked, err = PeekVarsVersion(varsBytes) + require.NoError(t, err) + require.Equal(t, CurrentBytecodeVersion, peeked) + + vars, err := DecodeVars(varsBytes) + require.NoError(t, err) + require.Equal(t, CurrentBytecodeVersion, vars.Version) +} + +func TestPeekVersionChecksMagicOnly(t *testing.T) { + _, err := PeekProgramVersion(Vars{}.Encode()) + require.Error(t, err, "a vars blob is not a program") + + _, err = PeekVarsVersion(Program{}.Encode()) + require.Error(t, err, "a program blob is not vars") + + _, err = PeekProgramVersion([]byte("NUM")) + require.Error(t, err) + + _, err = PeekProgramVersion(Program{}.Encode()[:formatHeaderLen-1]) + require.Error(t, err, "a header short of the section count is not a header") +} + +// The decoders apply CanRead against CurrentBytecodeVersion and report an +// unreadable blob with the typed error carrying both versions, while the peek +// still returns the encoded version so a caller can say what it was. +func TestDecodeAppliesVersionRule(t *testing.T) { + withVersion := func(buf []byte, v BytecodeVersion) []byte { + out := bytes.Clone(buf) + le.PutUint16(out[4:], v.Major) + le.PutUint16(out[6:], v.Minor) + return out + } + + current := CurrentBytecodeVersion + require.Positive(t, current.Major, "a 0.x current version would leave no older major to test against") + + versions := map[string]struct { + v BytecodeVersion + ok bool + }{ + "current": {current, true}, + "newer minor": {BytecodeVersion{current.Major, current.Minor + 1}, false}, + "far newer minor": {BytecodeVersion{current.Major, 0xFFFF}, false}, + "newer major": {BytecodeVersion{current.Major + 1, 0}, false}, + "older major": {BytecodeVersion{current.Major - 1, current.Minor}, false}, + } + if current.Minor > 0 { + versions["older minor"] = struct { + v BytecodeVersion + ok bool + }{BytecodeVersion{current.Major, current.Minor - 1}, true} + } + + blobs := []struct { + name string + magic string + buf []byte + decode func([]byte) (BytecodeVersion, error) + }{ + {"program", "NUMB", Program{}.Encode(), func(b []byte) (BytecodeVersion, error) { + p, err := DecodeProgram(b) + return p.Version, err + }}, + {"vars", "NVAR", Vars{}.Encode(), func(b []byte) (BytecodeVersion, error) { + v, err := DecodeVars(b) + return v.Version, err + }}, + } + + for name, tc := range versions { + for _, blob := range blobs { + t.Run(name+"/"+blob.name, func(t *testing.T) { + buf := withVersion(blob.buf, tc.v) + + peeked, err := peekVersion(blob.magic, buf) + require.NoError(t, err) + require.Equal(t, tc.v, peeked) + + got, err := blob.decode(buf) + if tc.ok { + require.NoError(t, err) + require.Equal(t, tc.v, got) + return + } + + var unsupported UnsupportedBytecodeVersionError + require.ErrorAs(t, err, &unsupported) + require.Equal(t, tc.v, unsupported.Encoded) + require.Equal(t, current, unsupported.Supported) + require.Contains(t, err.Error(), tc.v.String()) + }) + } + } +} diff --git a/internal/vm/vm.go b/internal/vm/vm.go new file mode 100644 index 00000000..95318aea --- /dev/null +++ b/internal/vm/vm.go @@ -0,0 +1,645 @@ +package vm + +import ( + "context" + "errors" + "fmt" + "math/big" + "sort" + + "github.com/formancehq/numscript/internal/funds" +) + +const nilReg byte = 0xFF + +// accountMetaKey identifies one set_account_meta slot during a run, so repeated +// writes to the same (account, scope, key) upsert rather than accumulating +// duplicate rows. +type accountMetaKey struct { + account string + scope string + key string +} + +// The three ops a mark cannot survive, rejected by the arms below. save is on the +// list permanently, unlike the other two: it floors the balance at zero, and the +// clamp destroys the information needed to invert it. +var ( + errSendWhileMarkOpen = errors.New("send while a mark is open") + errSetAssetWhileMarkOpen = errors.New("set_current_asset while a mark is open") + errSaveWhileMarkOpen = errors.New("save while a mark is open") +) + +type Vm struct { + Program Program + runstate *funds.RunState + + // a monetary is not a bank of its own: it travels as a (str asset, int amount) + // register pair + stringsRegs []string // asset,string,account + intsRegs []big.Int + portionsRegs []big.Rat + boolsRegs []bool +} + +// NewVm sizes each register bank from the count the program declares. MaxRegX is +// a count and 0xFF is the nil-register sentinel, so the real indices are +// 0..MaxRegX-1 and this is exact rather than an upper bound. +// +// Nothing here re-derives those counts: a program whose instructions name a +// register beyond its own declaration is malformed, and Exec is entitled to +// assume it isn't — see Verify, which is what checks it. +func NewVm( + program Program, +) *Vm { + return &Vm{ + Program: program, + stringsRegs: make([]string, program.MaxRegString), + intsRegs: make([]big.Int, program.MaxRegInt), + portionsRegs: make([]big.Rat, program.MaxRegPortion), + boolsRegs: make([]bool, program.MaxRegBool), + } +} + +type Store interface { + GetBalance( + ctx context.Context, + account string, + scope string, + asset string, + color string, + ) (*big.Int, error) + + GetMetadata( + ctx context.Context, + account, + scope, + key string, + ) (string, bool, error) +} + +func lookupMeta(ctx context.Context, store Store, account, scope, key string) (string, ExecutionError) { + v, ok, err := store.GetMetadata(ctx, account, scope, key) + if err != nil { + return "", StoreError{Wrapped: err} + } + if !ok { + return "", MetadataNotFoundError{Account: account, Key: key} + } + return v, nil +} + +type fundsStoreAdapter struct { + ctx context.Context + store Store +} + +func (s fundsStoreAdapter) GetBalance( + account string, + scope string, + asset string, + color string, +) (*big.Int, error) { + return s.store.GetBalance(s.ctx, account, scope, asset, color) +} + +func Exec[S Store]( + ctx context.Context, + vm *Vm, + vars *Vars, + store S, // a generic S should allow monomorphisation of the Store +) (funds.ExecutionResult, ExecutionError) { + fundsStore := fundsStoreAdapter{ + ctx: ctx, + store: store, + } + // RunState fetches balances lazily through this store; a fetch error surfaces + // from the RunState call that triggered it, wrapped in StoreError below. + if vm.runstate == nil { + vm.runstate = funds.New(fundsStore) + } else { + vm.runstate.Reset(fundsStore) + } + runstate := vm.runstate + + var txMeta map[string]string + // accountsMeta accumulates with upsert semantics (last write to a given + // (account, scope, key) wins), so it is keyed during the run and only + // flattened into the row-based funds.AccountsMetadata at the very end. + var accountsMeta map[accountMetaKey]string + + // Hoist register banks and constant pools into locals so the hot loop indexes + // them directly instead of reloading the header off *vm / vm.program on every + // access. + intsRegs := vm.intsRegs + stringsRegs := vm.stringsRegs + portionsRegs := vm.portionsRegs + boolsRegs := vm.boolsRegs + intsPool := vm.Program.IntsPool + stringsPool := vm.Program.StringsPool + + instrs := vm.Program.Instructions + instructionsLen := len(instrs) + + var currentAsset string + pc := 0 + + for pc < instructionsLen { + instr := instrs[pc] + pc++ + + switch Opcode(instr.Opcode) { + // --- Domain-specific ops + case Op_PullAccount: + // TODO crashes if this is the last instruction (the ext word is + // missing): instrs[pc] reads past the end. e.g. a program ending in a + // lone Op_PullAccount word. + instrExt := instrs[pc] + pc++ + + account := stringsRegs[instr.B] + + var cap *big.Int + if instr.C != nilReg { + cap = &intsRegs[instr.C] + } + + var overdraft *big.Int + if instrExt.A != nilReg { + overdraft = &intsRegs[instrExt.A] + } + + var color string + if instrExt.B != nilReg { + color = stringsRegs[instrExt.B] + } + + var scope string + if instrExt.C != nilReg { + scope = stringsRegs[instrExt.C] + } + + out := &intsRegs[instr.A] + switch { + case cap != nil: + if err := runstate.Pull(out, account, scope, cap, overdraft, color); err != nil { + return funds.ExecutionResult{}, StoreError{Wrapped: err} + } + case overdraft != nil: + if err := runstate.PullUncapped(out, account, scope, overdraft, color); err != nil { + return funds.ExecutionResult{}, StoreError{Wrapped: err} + } + default: + return funds.ExecutionResult{}, InvalidUncappedSource{Account: account} + } + + case Op_SendToAccount: + // a send while a mark is open would consume the queue from the front and + // leave that mark pointing at the wrong boundary. Compiled numscript never + // emits it, since sources only pull. + if runstate.HasOpenMark() { + return funds.ExecutionResult{}, InternalError{Err: errSendWhileMarkOpen} + } + + var dest *string + if instr.A != nilReg { + s := stringsRegs[instr.A] + dest = &s + } + + var cap *big.Int + if instr.B != nilReg { + cap = &intsRegs[instr.B] + } + + var scope string + if instr.C != nilReg { + scope = stringsRegs[instr.C] + } + + if cap == nil { + if err := runstate.SendUncapped(dest, scope, nil); err != nil { + return funds.ExecutionResult{}, StoreError{Wrapped: err} + } + } else { + if err := runstate.Send(dest, scope, cap, nil); err != nil { + return funds.ExecutionResult{}, StoreError{Wrapped: err} + } + } + + case Op_CheckEnoughFunds: + got := &intsRegs[instr.A] + needed := &intsRegs[instr.B] + // exact, like the interpreter's tryTakingExact + if got.Cmp(needed) != 0 { + return funds.ExecutionResult{}, MissingFundsError{ + Asset: currentAsset, + Got: new(big.Int).Set(got), + Needed: new(big.Int).Set(needed), + } + } + + case Op_Save: + // a save while a mark is open survives the rewind, which only repays queued + // sources and reverses postings; its floor at zero is not invertible at all. + if runstate.HasOpenMark() { + return funds.ExecutionResult{}, InternalError{Err: errSaveWhileMarkOpen} + } + + instrExt := instrs[pc] + pc++ + + account := stringsRegs[instr.A] + asset := stringsRegs[instr.B] + var amount *big.Int + if instr.C != nilReg { + amount = &intsRegs[instr.C] + } + var scope string + if instrExt.A != nilReg { + scope = stringsRegs[instrExt.A] + } + if err := runstate.Save(account, scope, asset, "", amount); err != nil { + return funds.ExecutionResult{}, StoreError{Wrapped: err} + } + + // the mark ops take no register: the mark is a depth on a LIFO the run-state + // owns, so nothing can name a depth it never marked + case Op_MarkPush: + runstate.MarkPush() + + case Op_MarkEnd: + if err := runstate.MarkEnd(instr.A == 1); err != nil { + return funds.ExecutionResult{}, InternalError{Err: err} + } + + case Op_AssertLeftover: + leftover := &portionsRegs[instr.A] + sign := leftover.Sign() + if sign < 0 || (instr.B == 1 && sign != 0) { + sum := new(big.Rat).Sub(big.NewRat(1, 1), leftover) + return funds.ExecutionResult{}, InvalidAllotmentSum{ActualSum: *new(big.Rat).Set(sum)} + } + + case Op_SetCurrentAsset: + // a rewind repays queued funds into the current asset's balance, so + // changing the asset mid-region would repay the wrong one + if runstate.HasOpenMark() { + return funds.ExecutionResult{}, InternalError{Err: errSetAssetWhileMarkOpen} + } + currentAsset = stringsRegs[instr.A] + runstate.SetCurrentAsset(currentAsset) + + case Op_AssertSameAsset: + left := stringsRegs[instr.A] + right := stringsRegs[instr.B] + if left != right { + return funds.ExecutionResult{}, AssetMismatchError{ + Expected: left, + Got: right, + } + } + + case Op_AssertValidAccount: + account := stringsRegs[instr.A] + if !funds.ValidateAccount(account) { + return funds.ExecutionResult{}, InvalidAccountName{Name: account} + } + + case Op_AssertValidColor: + color := stringsRegs[instr.A] + if !funds.ValidateColor(color) { + return funds.ExecutionResult{}, InvalidColor{Color: color} + } + + case Op_AssertValidScope: + scope := stringsRegs[instr.A] + if !funds.ValidateScope(scope) { + return funds.ExecutionResult{}, InvalidScope{Scope: scope} + } + + case Op_AssertUnscoped: + if scope := stringsRegs[instr.A]; scope != "" { + return funds.ExecutionResult{}, CannotCastScopedAccountToString{Account: stringsRegs[instr.B], Scope: scope} + } + + case Op_AssertNonNegativeBalance: + amount := &intsRegs[instr.A] + if amount.Sign() < 0 { + return funds.ExecutionResult{}, NegativeBalanceError{ + Account: stringsRegs[instr.B], + Amount: *new(big.Int).Set(amount), + } + } + + case Op_AssertNonNegativeAmount: + amount := &intsRegs[instr.A] + if amount.Sign() < 0 { + return funds.ExecutionResult{}, NegativeAmountError{Amount: *new(big.Int).Set(amount)} + } + + case Op_AssertNonNegativePortion: + portion := &portionsRegs[instr.A] + if portion.Sign() < 0 { + return funds.ExecutionResult{}, NegativePortionError{Portion: *new(big.Rat).Set(portion)} + } + + case Op_SetTxMeta: + if txMeta == nil { + txMeta = map[string]string{} + } + txMeta[stringsRegs[instr.A]] = stringsRegs[instr.B] + + case Op_SetAccountMeta: + instrExt := instrs[pc] + pc++ + + if accountsMeta == nil { + accountsMeta = map[accountMetaKey]string{} + } + var scope string + if instrExt.A != nilReg { + scope = stringsRegs[instrExt.A] + } + key := accountMetaKey{ + account: stringsRegs[instr.A], + scope: scope, + key: stringsRegs[instr.B], + } + accountsMeta[key] = stringsRegs[instr.C] + + case Op_MetaStr: + instrExt := instrs[pc] + pc++ + + var scope string + if instrExt.A != nilReg { + scope = stringsRegs[instrExt.A] + } + v, err := lookupMeta(ctx, store, stringsRegs[instr.B], scope, stringsRegs[instr.C]) + if err != nil { + return funds.ExecutionResult{}, err + } + stringsRegs[instr.A] = v + + case Op_MetaInt: + instrExt := instrs[pc] + pc++ + + account, key := stringsRegs[instr.B], stringsRegs[instr.C] + var scope string + if instrExt.A != nilReg { + scope = stringsRegs[instrExt.A] + } + v, err := lookupMeta(ctx, store, account, scope, key) + if err != nil { + return funds.ExecutionResult{}, err + } + n, ok := funds.ParseNumber(v) + if !ok { + return funds.ExecutionResult{}, BadMetaValueError{Account: account, Key: key, Raw: v} + } + intsRegs[instr.A].Set(n) + + case Op_MetaPortion: + instrExt := instrs[pc] + pc++ + + account, key := stringsRegs[instr.B], stringsRegs[instr.C] + var scope string + if instrExt.A != nilReg { + scope = stringsRegs[instrExt.A] + } + v, err := lookupMeta(ctx, store, account, scope, key) + if err != nil { + return funds.ExecutionResult{}, err + } + r, perr := funds.ParsePortion(v) + if perr != nil { + return funds.ExecutionResult{}, BadMetaValueError{Account: account, Key: key, Raw: v} + } + portionsRegs[instr.A].Set(r) + + case Op_MetaMonetary: + // TODO crashes if this is the last instruction (the ext word carrying the + // amount destination is missing), same as Op_PullAccount. + instrExt := instrs[pc] + pc++ + + account, key := stringsRegs[instr.B], stringsRegs[instr.C] + var scope string + if instrExt.B != nilReg { + scope = stringsRegs[instrExt.B] + } + v, err := lookupMeta(ctx, store, account, scope, key) + if err != nil { + return funds.ExecutionResult{}, err + } + asset, amount, merr := funds.ParseMonetary(v) + if merr != nil { + return funds.ExecutionResult{}, BadMetaValueError{Account: account, Key: key, Raw: v} + } + stringsRegs[instr.A] = asset + intsRegs[instrExt.A].Set(amount) + + // --- Vars + // TODO both crash if vars is nil (Exec called with no vars for a + // program that reads them), or if GetBC() >= len(vars pool) (caller + // passed fewer vars than the program declares). + case Op_LoadVarInt: + intsRegs[instr.A].Set(&vars.IntsPool[instr.GetBC()]) + + case Op_LoadVarStr: + stringsRegs[instr.A] = vars.StringsPool[instr.GetBC()] + + // --- Jumps + case Op_JmpIfFalse: + if !boolsRegs[instr.A] { + pc += int(instr.GetBC()) + } + + case Op_JmpIfTrue: + if boolsRegs[instr.A] { + pc += int(instr.GetBC()) + } + + case Op_Jmp: + pc += int(instr.GetBC()) + + // --- consts + // TODO both crash if GetBC() >= len(pool), e.g. an Op_LoadInt referring to + // pool index 5 in a program whose ints pool has 3 entries. + case Op_LoadInt: + const_ := &intsPool[instr.GetBC()] + intsRegs[instr.A].Set(const_) + + case Op_LoadStr: + const_ := stringsPool[instr.GetBC()] + stringsRegs[instr.A] = const_ + + case Op_ConstTrue: + boolsRegs[instr.A] = true + + case Op_ConstFalse: + boolsRegs[instr.A] = false + + // --- Binary ops + case Op_LtInt: + boolsRegs[instr.A] = intsRegs[instr.B].Cmp(&intsRegs[instr.C]) < 0 + + case Op_EqInt: + boolsRegs[instr.A] = intsRegs[instr.B].Cmp(&intsRegs[instr.C]) == 0 + + case Op_AddInt: + left := &intsRegs[instr.B] + right := &intsRegs[instr.C] + intsRegs[instr.A].Add(left, right) + + case Op_SubInt: + left := &intsRegs[instr.B] + right := &intsRegs[instr.C] + intsRegs[instr.A].Sub(left, right) + + case Op_AddString: + stringsRegs[instr.A] = stringsRegs[instr.B] + stringsRegs[instr.C] + + case Op_StrEq: + boolsRegs[instr.A] = stringsRegs[instr.B] == stringsRegs[instr.C] + + // portion comparison is *value* comparison: big.Rat normalises on + // construction, so 1/2 and 2/4 are the same rational and compare equal. + // Comparing numerator/denominator pairs separately would be wrong. + case Op_LtPortion: + boolsRegs[instr.A] = portionsRegs[instr.B].Cmp(&portionsRegs[instr.C]) < 0 + + case Op_EqPortion: + boolsRegs[instr.A] = portionsRegs[instr.B].Cmp(&portionsRegs[instr.C]) == 0 + + case Op_AddPortion: + left := &portionsRegs[instr.B] + right := &portionsRegs[instr.C] + portionsRegs[instr.A].Add(left, right) + + case Op_SubPortion: + left := &portionsRegs[instr.B] + right := &portionsRegs[instr.C] + portionsRegs[instr.A].Sub(left, right) + + case Op_MulPortion: + left := &portionsRegs[instr.B] + right := &portionsRegs[instr.C] + portionsRegs[instr.A].Mul(left, right) + + case Op_IntToPortion: + portionsRegs[instr.A].SetInt(&intsRegs[instr.B]) + + // floor: big.Rat's denominator is always positive, so Div (Euclidean) is + // the floor for negatives too + case Op_PortionToInt: + p := &portionsRegs[instr.B] + intsRegs[instr.A].Div(p.Num(), p.Denom()) + + case Op_MkPortion: + num := &intsRegs[instr.B] + den := &intsRegs[instr.C] + if den.Sign() == 0 { + return funds.ExecutionResult{}, DivideByZeroError{Numerator: *new(big.Int).Set(num)} + } + portionsRegs[instr.A].SetFrac(num, den) + + case Op_Balance: + instrExt := instrs[pc] + pc++ + + account := stringsRegs[instr.B] + asset := stringsRegs[instr.C] + var scope string + if instrExt.A != nilReg { + scope = stringsRegs[instrExt.A] + } + + bal, err := runstate.GetAccountBalance(account, scope, asset, "") + if err != nil { + return funds.ExecutionResult{}, StoreError{Wrapped: err} + } + // only the amount: the asset of the result is the asset operand, which + // the caller already holds in reg C + intsRegs[instr.A].Set(bal) + + // --- Unary ops + case Op_IntCopy: + arg := &intsRegs[instr.B] + intsRegs[instr.A].Set(arg) + + case Op_PortionCopy: + arg := &portionsRegs[instr.B] + portionsRegs[instr.A].Set(arg) + + case Op_StrCopy: + stringsRegs[instr.A] = stringsRegs[instr.B] + + case Op_BoolCopy: + boolsRegs[instr.A] = boolsRegs[instr.B] + + case Op_NegInt: + arg := &intsRegs[instr.B] + intsRegs[instr.A].Neg(arg) + + case Op_IntToString: + stringsRegs[instr.A] = intsRegs[instr.B].String() + + case Op_PortionToString: + stringsRegs[instr.A] = portionsRegs[instr.B].String() + + case Op_MonetaryToString: + stringsRegs[instr.A] = stringsRegs[instr.B] + " " + intsRegs[instr.C].String() + + case Op_IsZero: + boolsRegs[instr.A] = intsRegs[instr.B].Sign() == 0 + + case Op_Not: + boolsRegs[instr.A] = !boolsRegs[instr.B] + + default: + return funds.ExecutionResult{}, InternalError{Err: fmt.Errorf("unknown opcode %d", instr.Opcode)} + } + } + + var accountsMetaRows funds.AccountsMetadata + if len(accountsMeta) != 0 { + accountsMetaRows = make(funds.AccountsMetadata, 0, len(accountsMeta)) + for k, v := range accountsMeta { + accountsMetaRows = append(accountsMetaRows, funds.AccountMetadataEntry{ + Account: k.account, + Scope: k.scope, + Key: k.key, + Value: v, + }) + } + // deterministic output: accountsMeta was built from a map, whose iteration + // order is random + sort.Slice(accountsMetaRows, func(i, j int) bool { + a, b := accountsMetaRows[i], accountsMetaRows[j] + if a.Account != b.Account { + return a.Account < b.Account + } + if a.Scope != b.Scope { + return a.Scope < b.Scope + } + return a.Key < b.Key + }) + } + + postings := runstate.GetPostings() + for _, p := range postings { + if !funds.ValidatePosting(p) { + return funds.ExecutionResult{}, InternalError{Err: InvalidPostingError{Posting: p}} + } + } + + return funds.ExecutionResult{ + Postings: postings, + Metadata: txMeta, + AccountsMetadata: accountsMetaRows, + }, nil +} diff --git a/internal/vm/vm_test.go b/internal/vm/vm_test.go new file mode 100644 index 00000000..6d8922e7 --- /dev/null +++ b/internal/vm/vm_test.go @@ -0,0 +1,539 @@ +package vm + +// White-box tests (package vm) that build a Program from struct literals, so +// they can reach encodings the compiler doesn't emit. Behavioural VM tests are +// written in the IR textual format instead — see ir_test.go. + +import ( + "context" + "errors" + "math/big" + "testing" + + "github.com/formancehq/numscript/internal/funds" + "github.com/stretchr/testify/require" +) + +// --- register allocation: one $rN namespace -> typed banks ---------------- +// +// $r0 "USD/2" -> strings[0] (sUSD) $r6 remaining -> ints[3] (iRem) +// $r1 10 -> ints[0] (iTen) $r7 "s1" -> strings[2] (sS1) +// $r3 asset -> strings[1] (sAsset) $r8 pulled1 -> ints[4] (iPulled1) +// $r4 amount -> ints[1] (iAmount) $r9 "s2" -> strings[3] (sS2) +// $r5 sum=0 -> ints[2] (iSum) $r10 pulled2 -> ints[5] (iPulled2) +// $r11 "dest" -> strings[4] (sDest) +// (added) zero overdraft bound -> ints[6] (iZero) -- gives BoundedZero +const ( + sUSD, sAsset, sS1, sS2, sDest = 0, 1, 2, 3, 4 + iTen, iAmount, iSum, iRem, iPulled1 = 0, 1, 2, 3, 4 + iPulled2, iZero = 5, 6 +) + +func abc(op Opcode, a, b, c byte) Instruction { + return Instruction{Opcode: byte(op), A: a, B: b, C: c} +} + +// fullBanks declares every register bank at its maximum. NewVm sizes the banks +// from these counts, and the Program literals in these files don't set them — +// counting registers by hand would be noise in tests that aren't about +// allocation. 255 is what the fixed [256] banks gave before sizing became a +// function of the program. +func fullBanks(p Program) Program { + p.MaxRegString, p.MaxRegInt, p.MaxRegPortion, p.MaxRegBool = 255, 255, 255, 255 + return p +} + +func newTestVm(p Program) *Vm { return NewVm(fullBanks(p)) } + +func bc(op Opcode, a byte, v uint16) Instruction { + return Instruction{Opcode: byte(op), A: a, B: byte(v), C: byte(v >> 8)} +} + +// --- mock store ----------------------------------------------------------- + +type mockStore struct { + bal map[funds.PairKey]int64 + meta map[string]map[string]string +} + +func (m mockStore) GetBalance(ctx context.Context, account, scope, asset string, color string) (*big.Int, error) { + return big.NewInt(m.bal[funds.PairKey{Account: account, Scope: scope, Asset: asset}]), nil +} + +func (m mockStore) GetMetadata(ctx context.Context, account, scope, key string) (string, bool, error) { + v, ok := m.meta[account][key] + return v, ok, nil +} + +var _ Store = (*mockStore)(nil) + +// --- the test ------------------------------------------------------------- + +func assertValidAccountProgram(name string) Program { + return Program{ + Instructions: []Instruction{ + bc(Op_LoadStr, 0, 0), + abc(Op_AssertValidAccount, 0, nilReg, nilReg), + }, + StringsPool: []string{name}, + } +} + +func balanceNonNegativeProgram() Program { + return Program{ + Instructions: []Instruction{ + bc(Op_LoadStr, 0, 0), + bc(Op_LoadStr, 1, 1), + abc(Op_Balance, 0, 0, 1), + abc(0, nilReg, nilReg, nilReg), // ext: no scope + abc(Op_AssertNonNegativeBalance, 0, 0, nilReg), + }, + StringsPool: []string{"acc", "USD/2"}, + } +} + +func TestAssertNonNegativeBalance(t *testing.T) { + store := mockStore{bal: map[funds.PairKey]int64{{Account: "acc", Asset: "USD/2"}: 50}} + if _, err := Exec(context.Background(), newTestVm(balanceNonNegativeProgram()), nil, store); err != nil { + t.Fatalf("non-negative balance rejected: %v", err) + } + + store = mockStore{bal: map[funds.PairKey]int64{{Account: "acc", Asset: "USD/2"}: -50}} + _, err := Exec(context.Background(), newTestVm(balanceNonNegativeProgram()), nil, store) + if _, ok := err.(NegativeBalanceError); !ok { + t.Fatalf("expected NegativeBalanceError, got %v", err) + } +} + +func TestUnknownOpcode(t *testing.T) { + prog := Program{Instructions: []Instruction{abc(0xFE, 0, 0, 0)}} + _, err := Exec(context.Background(), newTestVm(prog), nil, mockStore{}) + if _, ok := err.(InternalError); !ok { + t.Fatalf("expected InternalError, got %v", err) + } +} + +func TestMkPortionDivideByZero(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadInt, 0, 0), + bc(Op_LoadInt, 1, 1), + abc(Op_MkPortion, 0, 0, 1), + }, + IntsPool: []big.Int{*big.NewInt(1), *big.NewInt(0)}, + } + _, err := Exec(context.Background(), newTestVm(prog), nil, mockStore{}) + if _, ok := err.(DivideByZeroError); !ok { + t.Fatalf("expected DivideByZeroError, got %v", err) + } +} + +func TestAssertValidAccount(t *testing.T) { + _, err := Exec(context.Background(), newTestVm(assertValidAccountProgram("users:001:wallet")), nil, mockStore{}) + if err != nil { + t.Fatalf("valid account rejected: %v", err) + } + + _, err = Exec(context.Background(), newTestVm(assertValidAccountProgram("bad name!")), nil, mockStore{}) + if _, ok := err.(InvalidAccountName); !ok { + t.Fatalf("expected InvalidAccountName, got %v", err) + } +} + +// Nothing reads a bool yet, so the bank itself is the only observable effect. +func TestConstBool(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + abc(Op_ConstTrue, 0, nilReg, nilReg), + abc(Op_ConstFalse, 1, nilReg, nilReg), + // a register written twice keeps the last value + abc(Op_ConstTrue, 2, nilReg, nilReg), + abc(Op_ConstFalse, 2, nilReg, nilReg), + }, + } + + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + + require.True(t, vm.boolsRegs[0]) + require.False(t, vm.boolsRegs[1]) + require.False(t, vm.boolsRegs[2]) + require.False(t, vm.boolsRegs[3], "untouched registers stay false") +} + +// is_zero is the only projection from a quantity to a condition, so it has to +// agree with what Op_JmpIfZero used to test: sign, not magnitude. +func TestIsZero(t *testing.T) { + testCases := []struct { + name string + value int64 + want bool + }{ + {"zero", 0, true}, + {"positive", 7, false}, + {"negative", -7, false}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadInt, 0, 0), + abc(Op_IsZero, 0, 0, nilReg), + }, + IntsPool: []big.Int{*big.NewInt(tc.value)}, + } + + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + require.Equal(t, tc.want, vm.boolsRegs[0]) + }) + } +} + +// Portion addition and subtraction, over unequal denominators so the result has +// to be a real rational sum rather than a numerator-wise one. +func TestPortionArithmetic(t *testing.T) { + testCases := []struct { + name string + op Opcode + numL, denL int64 + numR, denR int64 + wantNum, wantDen int64 + }{ + {"add with equal denominators", Op_AddPortion, 1, 4, 1, 4, 1, 2}, + {"add with unequal denominators", Op_AddPortion, 1, 6, 1, 3, 1, 2}, + {"add to a whole", Op_AddPortion, 1, 3, 2, 3, 1, 1}, + {"add past a whole", Op_AddPortion, 3, 4, 1, 2, 5, 4}, + {"sub with unequal denominators", Op_SubPortion, 1, 2, 1, 6, 1, 3}, + {"sub to zero", Op_SubPortion, 1, 3, 1, 3, 0, 1}, + {"sub below zero", Op_SubPortion, 1, 4, 1, 2, -1, 4}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadInt, 0, 0), bc(Op_LoadInt, 1, 1), + bc(Op_LoadInt, 2, 2), bc(Op_LoadInt, 3, 3), + abc(Op_MkPortion, 0, 0, 1), + abc(Op_MkPortion, 1, 2, 3), + abc(tc.op, 2, 0, 1), + }, + IntsPool: []big.Int{ + *big.NewInt(tc.numL), *big.NewInt(tc.denL), + *big.NewInt(tc.numR), *big.NewInt(tc.denR), + }, + } + + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + want := big.NewRat(tc.wantNum, tc.wantDen) + require.Zero(t, vm.portionsRegs[2].Cmp(want), + "got %s, want %s", vm.portionsRegs[2].RatString(), want.RatString()) + }) + } +} + +// One copy per bank. Each case writes a distinct value into reg 1, copies reg 1 +// into reg 0, and checks reg 0 took it — so a copy wired to the wrong bank, or a +// no-op, fails. +func TestBankCopies(t *testing.T) { + t.Run("int", func(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadInt, 1, 0), + abc(Op_IntCopy, 0, 1, nilReg), + }, + IntsPool: []big.Int{*big.NewInt(-42)}, + } + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + require.Zero(t, vm.intsRegs[0].Cmp(big.NewInt(-42))) + }) + + t.Run("portion", func(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadInt, 0, 0), + bc(Op_LoadInt, 1, 1), + abc(Op_MkPortion, 1, 0, 1), + abc(Op_PortionCopy, 0, 1, nilReg), + }, + IntsPool: []big.Int{*big.NewInt(1), *big.NewInt(3)}, + } + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + require.Zero(t, vm.portionsRegs[0].Cmp(big.NewRat(1, 3))) + }) + + t.Run("str", func(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadStr, 1, 0), + abc(Op_StrCopy, 0, 1, nilReg), + }, + StringsPool: []string{"USD/2"}, + } + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + require.Equal(t, "USD/2", vm.stringsRegs[0]) + }) + + t.Run("bool", func(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + abc(Op_ConstTrue, 1, nilReg, nilReg), + abc(Op_BoolCopy, 0, 1, nilReg), + // and the false direction, over a register that already held true + abc(Op_ConstTrue, 2, nilReg, nilReg), + abc(Op_ConstFalse, 3, nilReg, nilReg), + abc(Op_BoolCopy, 2, 3, nilReg), + }, + } + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + require.True(t, vm.boolsRegs[0]) + require.False(t, vm.boolsRegs[2], "copying false over true") + }) +} + +// A copy is a copy, not an alias: overwriting the source must not disturb the +// destination. Only the int and portion banks can get this wrong, since those two +// hold big values that are Set() into place rather than assigned. +func TestCopiesAreNotAliases(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadInt, 1, 0), // $1 = 7 + abc(Op_IntCopy, 0, 1, nilReg), // $0 = copy $1 + bc(Op_LoadInt, 1, 1), // $1 = 9 + }, + IntsPool: []big.Int{*big.NewInt(7), *big.NewInt(9)}, + } + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + require.Zero(t, vm.intsRegs[0].Cmp(big.NewInt(7)), "the copy tracked its source") + require.Zero(t, vm.intsRegs[1].Cmp(big.NewInt(9))) +} + +// The two int comparisons, over both signs. Each case also asserts the negation, +// so `!=` — which has no opcode — is covered wherever `==` is. +func TestIntComparisons(t *testing.T) { + testCases := []struct { + name string + op Opcode + left, right int64 + want bool + }{ + {"lt when less", Op_LtInt, 3, 7, true}, + {"lt when equal", Op_LtInt, 7, 7, false}, + {"lt when greater", Op_LtInt, 7, 3, false}, + {"lt across zero", Op_LtInt, -7, 3, true}, + {"lt on negatives", Op_LtInt, -7, -3, true}, + + {"eq when equal", Op_EqInt, 7, 7, true}, + {"eq when different", Op_EqInt, 7, 3, false}, + {"eq on negatives", Op_EqInt, -7, -7, true}, + {"eq distinguishes sign", Op_EqInt, -7, 7, false}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadInt, 0, 0), + bc(Op_LoadInt, 1, 1), + abc(tc.op, 0, 0, 1), + abc(Op_Not, 1, 0, nilReg), + }, + IntsPool: []big.Int{*big.NewInt(tc.left), *big.NewInt(tc.right)}, + } + + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + require.Equal(t, tc.want, vm.boolsRegs[0]) + require.Equal(t, !tc.want, vm.boolsRegs[1], "not") + }) + } +} + +// The four derived operators have no opcodes: the front end normalises them onto +// Lt / Eq / Not. This checks each lowering against the operator it stands for, +// which is the property that makes leaving them out safe. +func TestDerivedComparisonLowerings(t *testing.T) { + // each lowering as it would be emitted, over a grid that covers <, == and > + values := []int64{-7, -1, 0, 1, 7} + + testCases := []struct { + name string + emit []Instruction // leaves the answer in bool reg 0 + want func(l, r int64) bool + }{ + { + name: "a > b -> Lt(b, a)", + emit: []Instruction{abc(Op_LtInt, 0, 1, 0)}, + want: func(l, r int64) bool { return l > r }, + }, + { + name: "a <= b -> Not(Lt(b, a))", + emit: []Instruction{abc(Op_LtInt, 1, 1, 0), abc(Op_Not, 0, 1, nilReg)}, + want: func(l, r int64) bool { return l <= r }, + }, + { + name: "a >= b -> Not(Lt(a, b))", + emit: []Instruction{abc(Op_LtInt, 1, 0, 1), abc(Op_Not, 0, 1, nilReg)}, + want: func(l, r int64) bool { return l >= r }, + }, + { + name: "a != b -> Not(Eq(a, b))", + emit: []Instruction{abc(Op_EqInt, 1, 0, 1), abc(Op_Not, 0, 1, nilReg)}, + want: func(l, r int64) bool { return l != r }, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + for _, l := range values { + for _, r := range values { + instrs := []Instruction{bc(Op_LoadInt, 0, 0), bc(Op_LoadInt, 1, 1)} + prog := Program{ + Instructions: append(instrs, tc.emit...), + IntsPool: []big.Int{*big.NewInt(l), *big.NewInt(r)}, + } + + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + require.Equal(t, tc.want(l, r), vm.boolsRegs[0], "l=%d r=%d", l, r) + } + } + }) + } +} + +// Portion comparison is value comparison: big.Rat normalises on construction, so +// equal rationals with different spellings must compare equal. +func TestPortionComparisons(t *testing.T) { + // builds two portions from (numL/denL, numR/denR) and compares them + run := func(t *testing.T, op Opcode, numL, denL, numR, denR int64) bool { + t.Helper() + prog := Program{ + Instructions: []Instruction{ + bc(Op_LoadInt, 0, 0), bc(Op_LoadInt, 1, 1), + bc(Op_LoadInt, 2, 2), bc(Op_LoadInt, 3, 3), + abc(Op_MkPortion, 0, 0, 1), // portion 0 = numL/denL + abc(Op_MkPortion, 1, 2, 3), // portion 1 = numR/denR + abc(op, 0, 0, 1), + }, + IntsPool: []big.Int{ + *big.NewInt(numL), *big.NewInt(denL), + *big.NewInt(numR), *big.NewInt(denR), + }, + } + vm := newTestVm(prog) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + return vm.boolsRegs[0] + } + + t.Run("equality is by value, not by numerator/denominator", func(t *testing.T) { + require.True(t, run(t, Op_EqPortion, 1, 2, 2, 4), "1/2 == 2/4") + require.True(t, run(t, Op_EqPortion, 3, 9, 1, 3), "3/9 == 1/3") + require.False(t, run(t, Op_EqPortion, 1, 2, 1, 3), "1/2 != 1/3") + }) + + t.Run("ordering", func(t *testing.T) { + require.True(t, run(t, Op_LtPortion, 1, 3, 1, 2), "1/3 < 1/2") + require.False(t, run(t, Op_LtPortion, 1, 2, 1, 3), "1/2 not < 1/3") + require.False(t, run(t, Op_LtPortion, 1, 2, 2, 4), "equal values are not <") + }) +} + +// The two conditional jumps are duals: each takes the edge the other doesn't. +func TestConditionalJumps(t *testing.T) { + // jump over a CONST_TRUE writing bool reg 1, so reg 1 reports whether the + // jump was taken + prog := func(jmp Opcode, cond Opcode) Program { + return Program{ + Instructions: []Instruction{ + abc(cond, 0, nilReg, nilReg), + bc(jmp, 0, 1), + abc(Op_ConstTrue, 1, nilReg, nilReg), + }, + } + } + + testCases := []struct { + name string + jmp Opcode + cond Opcode + taken bool + }{ + {"jmp_if_false on false", Op_JmpIfFalse, Op_ConstFalse, true}, + {"jmp_if_false on true", Op_JmpIfFalse, Op_ConstTrue, false}, + {"jmp_if_true on true", Op_JmpIfTrue, Op_ConstTrue, true}, + {"jmp_if_true on false", Op_JmpIfTrue, Op_ConstFalse, false}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + vm := newTestVm(prog(tc.jmp, tc.cond)) + _, err := Exec(context.Background(), vm, nil, mockStore{}) + require.Nil(t, err) + require.Equal(t, tc.taken, !vm.boolsRegs[1], "jump taken") + }) + } +} + +func TestExecutionErrorMessages(t *testing.T) { + testCases := []struct { + name string + err ExecutionError + msg string + }{ + {"MissingFundsError", MissingFundsError{Asset: "USD/2", Needed: big.NewInt(10), Got: big.NewInt(4)}, + "missing funds for asset USD/2: needed 10, got 4"}, + {"AssetMismatchError", AssetMismatchError{Expected: "USD/2", Got: "EUR/2"}, + "asset mismatch: expected USD/2, got EUR/2"}, + {"InvalidUncappedSource", InvalidUncappedSource{Account: "src"}, + "unbounded source is not allowed here: @src"}, + {"InvalidAllotmentSum", InvalidAllotmentSum{ActualSum: *big.NewRat(3, 2)}, + "invalid allotment: portions must sum to 1, got 3/2"}, + {"MetadataNotFoundError", MetadataNotFoundError{Account: "acc", Key: "k"}, + `metadata not found: acc["k"]`}, + {"BadMetaValueError", BadMetaValueError{Account: "acc", Key: "k", Raw: "oops"}, + `invalid metadata value for acc["k"]: "oops"`}, + {"InvalidAccountName", InvalidAccountName{Name: "not an account"}, + `invalid account name: "not an account"`}, + {"InvalidColor", InvalidColor{Color: "red"}, + `invalid color name: "red"`}, + {"NegativeBalanceError", NegativeBalanceError{Account: "src", Amount: *big.NewInt(-1)}, + "cannot fetch negative balance from account @src"}, + {"DivideByZeroError", DivideByZeroError{Numerator: *big.NewInt(7)}, + "cannot divide by zero (in 7/0)"}, + {"InternalError", InternalError{Err: errors.New("boom")}, "internal error: boom"}, + {"StoreError", StoreError{Wrapped: errors.New("store is down")}, "store error: store is down"}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + require.Equal(t, tc.msg, tc.err.Error()) + }) + } +} + +// The two wrapping errors must stay unwrappable, so hosts can inspect the cause. +func TestExecutionErrorsUnwrap(t *testing.T) { + cause := errors.New("cause") + require.ErrorIs(t, InternalError{Err: cause}, cause) + require.ErrorIs(t, StoreError{Wrapped: cause}, cause) +} diff --git a/ir-textual-format.md b/ir-textual-format.md new file mode 100644 index 00000000..b7409167 --- /dev/null +++ b/ir-textual-format.md @@ -0,0 +1,373 @@ +# The IR textual format + +The compiler doesn't emit `vm.Instruction` directly: it emits a `[]ir.Instr` stream first (see [compiler-architecture.md](compiler-architecture.md)). This document specifies the **textual notation** for that stream — the thing you get when you dump a compiled program, and the thing you can write by hand to feed the assembler. + +It is a real format, not just a pretty-printing convention: it has a grammar, a parser, and a round-trip guarantee. + +The whole IR layer lives in [internal/ir/](internal/ir/), and that package is the entire API: + +| what | where | +| --- | --- | +| grammar | [IR.g4](IR.g4) (ANTLR; generated into `internal/ir/internal/syntax/antlrParser/` by `just generate`) | +| text → `[]ir.Instr` | `ir.Parse` in [internal/ir/parse.go](internal/ir/parse.go) | +| `[]ir.Instr` → text | `ir.Dump` in [internal/ir/dump.go](internal/ir/dump.go) | +| `[]ir.Instr` → `vm.Program` | `ir.Assemble` in [internal/ir/assemble.go](internal/ir/assemble.go) | +| register typing | `ir.Typecheck` in [internal/ir/typecheck.go](internal/ir/typecheck.go) | +| round-trip tests | `TestRoundtripAllInstructions` in [internal/ir/parse_test.go](internal/ir/parse_test.go) | + +`ir.Parse` is the only way in: the grammar's AST lives under `internal/ir/internal/syntax`, which Go's import rules make unreachable from anywhere outside `internal/ir`. Callers see instructions, never parse trees. + +**Round-trip property:** for every instruction, `ir.Dump` of what `ir.Parse` returns is the text it was given. This is what makes the format usable for snapshot tests and for hand-writing IR fixtures. See [Round-trip caveats](#round-trip-caveats) for the (few) inputs that don't survive it. + +## Lexical structure + +``` +REG '$' [a-zA-Z_] [a-zA-Z0-9_]* $r0, $r12, $pulled +LABEL '#' [a-zA-Z_] [a-zA-Z0-9_]* #inorder_end_0 +INT [0-9]+ 42 (no separators; a constant may be preceded by `-`) +STRING '"' ('\"' | ~["\r\n])* '"' "USD/2", "a\"b" +IDENTIFIER [a-z] [a-z0-9_]* mk_portion, account +TYPE_KEYWORD 'int' | 'str' | 'portion' | 'monetary' +BOOL 'true' | 'false' +``` + +* Spaces, tabs and newlines are **skipped**, not significant. Statements are delimited by the grammar, not by line breaks — `$r0 = 1 $r1 = 2` is two valid instructions. Newlines are pure convention (a very useful one). +* **There are no comments.** Any `//` or `#`-style comment is a syntax error (`#foo` lexes as a label). +* `BOOL` is likewise matched before `IDENTIFIER`, so `true` and `false` are reserved too. +* `TYPE_KEYWORD` is matched before `IDENTIFIER`, so `int`, `str`, `portion` and `monetary` are reserved: they cannot be used as an instruction name or an argument label. Register names are unaffected (`$int` is fine, the `$` starts a `REG`). `monetary` is vestigial — no instruction takes it as a type parameter any more (see below) — but it stays reserved until the grammar is regenerated. +* Instruction names and argument labels are lowercase-only (`IDENTIFIER` starts with `[a-z]`). + +## Statements + +A program is a flat sequence of two kinds of line: **label markers** and **instructions**. + +``` +#some_label ← label marker, flush left + set_current_asset($r3) ← instruction, indented 2 spaces +``` + +Indentation is cosmetic, but `ir.Dump` always emits labels at column 0 and instructions indented by two spaces. + +Instructions come in five shapes: + +``` + dest = name(args) instruction call with a destination + name(args) instruction call with no destination + dest = constant load + dest = $l + $r infix int arithmetic (only + and -) + $d += $r compound assign (only += and -=) +``` + +### Destinations + +``` +$r0 single register +[$r0, $r1] register list (meta_monetary only) +_ discard +``` + +A call that writes a register may also have no destination at all (`load_var(0)`): that is the same as writing `_ =` in front of it. A call that writes nothing (`set_current_asset($r0)`) takes no destination, not even `_`, and a register list is only for `meta_monetary`. + +`_` discards the result, and exists **only in the text**: there is no discard at the `ir.Instr` level. `ir.Parse` desugars each occurrence to a fresh register — allocated from the same counter as named ones, but bound to no name, so nothing can refer to it and each `_` gets its own (two discards that aliased would be forced to share a type). The write is still a write: the assembler gives that register a slot in its bank, so a discard costs a register even though nothing reads it. + +Because the desugaring happens on the way in, `_` doesn't survive a dump: `_ = int_copy($r0)` comes back as `$r1 = int_copy($r0)`. + +### Arguments + +Arguments are comma-separated and either **positional** or **labeled**: + +``` + check_enough_funds($r7, $r4) positional + $r8 = pull_account(account: $r5, cap: $r4) labeled +``` + +Which form an argument takes is fixed per instruction (see the reference below) — it is not a free choice. An argument value is one of: + +``` +$r0 register +#my_label label reference (the jumps only) +42 int literal (load_var index only) +``` + +A register list is a destination form only — no instruction takes one as an argument. + +Labeled arguments are looked up **by name**, so their order is free: `pull_account(cap: $c, account: $a)` is the same instruction as `pull_account(account: $a, cap: $c)`. `ir.Dump` always emits them in the canonical order given below. + +### Registers + +Registers in the IR are "logical": an unbounded stream of unsigned indices (`ir.Reg` is a `uint`), later mapped onto the VM's 256-per-bank physical registers by the assembler's allocator. Each register has exactly one type for its whole lifetime (`int`, `str`, `portion` or `bool`), checked by `ir.Typecheck` — the type is never written in the text, it is inferred from the instruction that writes the register. + +There is no monetary register. A monetary is a **pair** of registers, a `str` asset and an `int` amount, which the instructions below take and return separately. Nothing constructs or projects one, so `[USD/2 10]` is just the two registers holding `"USD/2"` and `10`. + +A register name is just a name: `ir.Parse` keeps a symbol table and allocates registers in order of **first appearance**, reusing the same one every later time a name shows up. `$r` is a convention, not an index — `$r7` is no more meaningful than `$src`. + +``` +$src = "acc" dumps back as $r0 = "acc" +$pulled = pull_account(account: $src) $r1 = pull_account(account: $r0) +``` + +This is why a dump round-trips: `ir.Dump` numbers registers `$r0`, `$r1`, … in the order they first appear, so re-parsing binds each name to the register it already had. Names of your own choosing are fine to write, they just come back as `$r` in first-appearance order. + +The compiler holds up its end by allocating registers in the order it emits them — `getCompiledOutput` in the compiler tests asserts the round-trip on every snapshot, so a change that breaks the ordering fails there. + +## Instruction reference + +Types are the register types of each operand; `?` marks an optional labeled argument. + +### Constants and variables + +| syntax | types | +| --- | --- | +| `$d = 42` | `int` | +| `$d = "USD/2"` | `str` | +| `$d = true` | `bool` | +| `$d = false` | `bool` | +| `$d = load_var(0)` | `int`; the index is a literal in `0..65535` | +| `$d = load_var(1)` | `str` | + +`true` and `false` are constants like the other two, but they need no pool entry: the value is in the opcode (`CONST_TRUE` / `CONST_FALSE`). They are only ever the right-hand side of a const assignment — no instruction takes a bool *operand*, so `set_current_asset(true)` doesn't parse. There is no `load_var` either: numscript has no bool of its own, so a bool register can only come from an instruction inside the program. + +`load_var` reads from the encoded `vm.Vars` pool at that index. There is no `load_var` / `load_var`: composite vars are encoded as their int/str components. A monetary var is two `load_var`s — `load_var` for the asset, `load_var` for the amount — and that pair *is* the value. + +### Pure arithmetic and constructors + +| syntax | signature | +| --- | --- | +| `$d = add_int($l, $r)` | `(int, int) -> int` | +| `$d = sub_int($l, $r)` | `(int, int) -> int` | +| `$d = add_string($l, $r)` | `(str, str) -> str` | +| `$d = add_portion($l, $r)` | `(portion, portion) -> portion` | +| `$d = sub_portion($l, $r)` | `(portion, portion) -> portion` | +| `$d = mul_portion($l, $r)` | `(portion, portion) -> portion` | +| `$d = mk_portion($num, $den)` | `(int, int) -> portion` | +| `$d = monetary_to_string($asset, $amt)` | `(str, int) -> str` — the `"ASSET AMOUNT"` form | + +`add_int` and `sub_int` have infix sugar, which is what `ir.Dump` always prints: + +``` + $r2 = $r0 + $r1 add_int + $r2 = $r0 - $r1 sub_int + $r0 += $r1 add_int where dest == left + $r0 -= $r1 sub_int where dest == left +``` + +So `add_int($a, $b)` parses fine, but a dump never contains it. No other operator has infix syntax. + +### Unary ops + +| syntax | signature | +| --- | --- | +| `$d = int_copy($a)` | `int -> int` | +| `$d = portion_copy($a)` | `portion -> portion` | +| `$d = str_copy($a)` | `str -> str` | +| `$d = bool_copy($a)` | `bool -> bool` | +| `$d = neg_int($a)` | `int -> int` | +| `$d = int_to_string($a)` | `int -> str` | +| `$d = portion_to_string($a)` | `portion -> str` | +| `$d = int_to_portion($a)` | `int -> portion` — exact | +| `$d = portion_to_int($a)` | `portion -> int` — **floors** | + +`int_to_portion` and `portion_to_int` are the only numeric crossings between the int and portion banks. `portion_to_int` truncates towards negative infinity (a `big.Rat` denominator is always positive, so this is `Div`, not a rounding), which is what makes an allotment share exact. + +There is no register-to-register move: `$r0 = $r1` is not valid syntax. Use the copy for the bank instead — there is exactly one per bank, and none crosses banks. A monetary has no copy of its own, since it is a `(str, int)` pair: copy the two halves. + +There is no `get_asset` / `get_amount` either: projecting a monetary means naming one of its two registers, which costs no instruction. `monetary_to_string` is listed above with the other constructors, since it takes the pair. + +### Comparisons and `not` + +Every instruction that produces a `bool`, other than the `true`/`false` constants: + +| syntax | signature | +| --- | --- | +| `$d = lt_int($l, $r)` | `(int, int) -> bool` — strict | +| `$d = eq_int($l, $r)` | `(int, int) -> bool` | +| `$d = str_eq($l, $r)` | `(str, str) -> bool` | +| `$d = is_zero($a)` | `int -> bool` — tests the *sign*, so a negative amount is not zero | +| `$d = lt_portion($l, $r)` | `(portion, portion) -> bool` — strict | +| `$d = eq_portion($l, $r)` | `(portion, portion) -> bool` | +| `$d = not($a)` | `bool -> bool` | + +Only `<` and `==` exist per type. The other four operators are **front-end normalisations**, so the IR never sees them and there is no `gt_*`, `lte_*`, `gte_*` or `neq_*`: + +``` +a < b -> lt_int($a, $b) +a > b -> lt_int($b, $a) operands swapped +a <= b -> $t = lt_int($b, $a) ; not($t) +a >= b -> $t = lt_int($a, $b) ; not($t) +a == b -> eq_int($a, $b) +a != b -> $t = eq_int($a, $b) ; not($t) +``` + +12 surface operators over 5 instructions. The reason is that every extra predicate is another case in the SMT encoder and in any formal model of the VM, so its cost is paid three times over; LLVM canonicalises the same way. `is_zero` is kept next to `eq_int` because it needs no materialised zero and sits on every quantity branch. + +`eq_portion` is **value** equality: `1/2 == 2/4` is true, since a portion register holds a normalised rational. + +`str` gets equality only, never ordering. Bool equality and structural comparison of tuples/arrays would also be front-end expansions rather than instructions. + +### Run-state reads (impure) + +| syntax | signature | +| --- | --- | +| `$d = balance($account, $asset, scope: $s)` | `(str, str) -> int` — the amount only; the monetary's asset is the `$asset` operand you already hold | +| `$d = meta($account, $key, scope: $s)` | `(str, str) -> str` | +| `$d = meta($account, $key, scope: $s)` | `(str, str) -> int` | +| `$d = meta($account, $key, scope: $s)` | `(str, str) -> portion` | +| `[$asset, $amt] = meta_monetary($account, $key, scope: $s)` | `(str, str) -> (str, int)` | + +`meta_monetary` is not `meta`: one store read yields both halves, so it is the only instruction that writes a **dest list**, and its list must be exactly two registers (asset then amount). + +`scope` is optional on all five (omitted means unscoped) and, like `pull_account`'s `color`, is a labeled argument — `balance($account, $asset)` and `balance($account, $asset, scope: $s)`. + +### Funds movement + +``` + $pulled = pull_account(account: $a, cap: $c, overdraft: $o, color: $col, scope: $s) +``` +`account: str` is required; `cap: int`, `overdraft: int`, `color: str`, `scope: str` are optional. Writes the amount actually pulled (`int`) into the destination. No `cap` means uncapped. Canonical dump order: `account, cap, overdraft, color, scope`. + +``` + send_to_account(account: $a, cap: $c, scope: $s) +``` +No destination. All arguments are optional: no `cap` sends everything currently queued; **no `account` refunds the funds to their sources without emitting postings**; no `scope` sends to the unscoped destination. + +``` + save(account: $a, asset: $as, amount: $amt, scope: $s) +``` +No destination. `account: str` and `asset: str` are required, `amount: int` and `scope: str` are optional — omitting `amount` saves the whole balance, omitting `scope` saves the unscoped balance. + +There is no allotment instruction. Splitting an amount across portions is built out of the pure ops above: each share is `portion_to_int(mul_portion($p_i, int_to_portion($amount)))`, and the leftover from flooring is then handed to the earliest shares a unit at a time, using `lt_int` and forward jumps to a shared exit. See `compileAllotmentSplit` in `internal/compiler/compiler.go`. + +### Assertions and checks + +| syntax | operands | +| --- | --- | +| `check_enough_funds($got, $needed)` | `int, int` | +| `assert_leftover($portion)` | `portion` — leftover must be `>= 0` | +| `assert_leftover_exact($portion)` | `portion` — leftover must be exactly `0` | +| `assert_non_negative_portion($portion)` | `portion` — an allotment clause portion must be `>= 0` | +| `assert_same_asset($l, $r)` | `str, str` | +| `assert_valid_account($a)` | `str` | +| `assert_valid_color($c)` | `str` | +| `assert_non_negative_balance($amt, $account)` | `int, str` — the account is only for the error | +| `set_current_asset($asset)` | `str` — required before `pull_account` / `send_to_account` | + +`assert_leftover` / `assert_leftover_exact` are two separate instruction names, not one instruction with an `exact:` flag. + +### Metadata writes + +| syntax | operands | +| --- | --- | +| `set_tx_meta($key, $value)` | `str, str` | +| `set_account_meta($account, $key, $value)` | `str, str, str` | + +### Control flow + +``` + jmp_if_false($cond, #my_label) + jmp_if_true($cond, #my_label) + jmp(#my_label) +#my_label +``` + +`$cond` is `bool`, so a quantity can't be a condition — that is the point of the bool bank, and `ir.Typecheck` rejects `jmp_if_false($some_amount, ..)` where it used to accept it. The two conditional forms are duals, so either edge of a condition is one instruction and there is no negation op. A bool comes from `true`/`false`, from `str_eq`, or from `is_zero` — the last being how a quantity reaches a branch: + +``` + $exhausted = is_zero($remaining_cap) + jmp_if_true($exhausted, #end) +``` + +For all three the target must be a label that is defined in the program, unique, and **after** the jump. The VM only permits forward jumps — that's what guarantees termination — and `ir.Parse` enforces all three rules, so a program that assembles can't loop: + +``` +jmp_if_false($r0, #nope) → label #nope is not defined in the program +#back → label #back is behind the jump (jumps must go forward) + jmp_if_false($r0, #back) +``` + +Together they express an if/else, which is how `@world`'s unboundedness is compiled (see `compiler-architecture.md`): + +``` + $eq = str_eq($account, $world) + jmp_if_false($eq, #not_world) + ; then arm + jmp(#end) +#not_world + ; else arm +#end +``` + +`labelMarker` is a pseudo-instruction: it emits no bytecode, it only feeds the assembler's symbol table. + +### Backtracking (`oneof`) + +``` + mark_push() // opens a region + mark_rewind() // closes the innermost region, undoing what was pulled in it + mark_commit() // closes the innermost region, keeping what was pulled in it +``` + +The three take no operands and write no register; the run state keeps the stack of open regions. `mark_rewind` and `mark_commit` are one instruction (`Op_MarkEnd`) differing in a flag. + +A `oneof` source opens a region per branch. A branch that covers the whole amount jumps to the end with its region still open, and `mark_commit` closes it there; otherwise `mark_rewind` undoes it and the next branch runs. The last branch has no check. `source = oneof { @a @b }` compiles to (simplified: the cap is a constant and the `@world` checks are trimmed): + +``` + $r4 = 10 + $r7 = 0 + mark_push() + $r6 = "a" + $r9 = pull_account(account: $r6, cap: $r4, overdraft: $r7) + $r10 = int_copy($r9) + $r11 = $r4 - $r9 + $r12 = is_zero($r11) + jmp_if_true($r12, #oneof_end_1) + mark_rewind() + mark_push() + $r13 = "b" + $r16 = pull_account(account: $r13, cap: $r4, overdraft: $r7) + $r10 = int_copy($r16) +#oneof_end_1 + mark_commit() +``` + +## A full example + +`send [USD/2 10] (source = @src destination = @dest)` compiles to: + +``` + $r0 = "USD/2" + $r1 = 10 + set_current_asset($r0) + $r2 = "src" + $r3 = 0 + $r4 = pull_account(account: $r2, cap: $r1, overdraft: $r3) + check_enough_funds($r4, $r1) + $r5 = "dest" + send_to_account(account: $r5) +``` + +`$r0` and `$r1` *are* the monetary: `set_current_asset` reads the asset half and the cap is the amount half, with nothing in between. + +## Round-trip caveats + +Known asymmetries between what `ir.Dump` writes and what the parser accepts: + +* **Register names don't survive.** `$src` comes back as `$r`, numbered by first appearance (see [Registers](#registers)). +* **`_` doesn't survive.** It's desugared to a fresh register on the way in, so it dumps as that register (see [Destinations](#destinations)). + +## Error handling + +Text → `[]ir.Instr` never panics: it reports `ir.Error`s. Anything the grammar rejects comes back as a syntax error (and since ANTLR's error recovery leaves partial nodes behind, no AST is built at all in that case). On top of that, `ir.Parse` reports what the grammar can't express: + +* unknown instruction names, and a type parameter on an instruction that doesn't take one +* wrong argument kinds or counts, unknown or duplicate labeled arguments +* duplicate labels, and jumps that don't resolve or don't go forward +* **a register that is read but never written**, reported under the name the text used: + +``` + $a = 42 + $y = lt_int($a, $b) → 3:3: register $b is read but never written +``` + +Since jumps only go forward, text order is execution order, so a read with no earlier write can't be reached by any path — it would hand the VM whatever that register happens to hold. Note this is a linear check: a register written only inside a branch that may be skipped and read afterwards is *not* caught here, which is the job of the path-sensitive bytecode verifier. + +Type errors are **not** checked by `ir.Parse`: writing a `str` register where an `int` is expected parses happily and is caught by `ir.Typecheck` afterwards. diff --git a/numscript.go b/numscript.go index fda5d501..152925dc 100644 --- a/numscript.go +++ b/numscript.go @@ -3,8 +3,10 @@ package numscript import ( "context" + "github.com/formancehq/numscript/internal/compiler" "github.com/formancehq/numscript/internal/interpreter" "github.com/formancehq/numscript/internal/parser" + "github.com/formancehq/numscript/internal/vm" ) // This struct represents a parsed numscript source code @@ -126,3 +128,143 @@ func (p ParseResult) ResolveDependencies(ctx context.Context, vars VariablesMap, func (p ParseResult) GetSource() string { return p.parseResult.Source } + +type ( + VarsEncoder = compiler.VarsEncoder + CompiledProgram = vm.Program + VMStore = vm.Store + Vm = vm.Vm + Vars = vm.Vars +) + +var NewVm = vm.NewVm + +var DecodeVars = vm.DecodeVars + +func (p ParseResult) Compile() (VarsEncoder, CompiledProgram, error) { + return p.CompileWithFeatureFlags(nil) +} + +// CompileWithFeatureFlags compiles the program, rejecting any construct gated +// behind an experimental feature flag that isn't in featureFlags. +func (p ParseResult) CompileWithFeatureFlags(featureFlags map[string]struct{}) (VarsEncoder, CompiledProgram, error) { + if len(p.parseResult.Errors) != 0 { + return VarsEncoder{}, CompiledProgram{}, p.parseResult.Errors[0] + } + + if featureFlags == nil { + featureFlags = make(map[string]struct{}) + } + + return compiler.Compile(p.parseResult.Value, featureFlags) +} + +func Compile(source string) (VarsEncoder, CompiledProgram, error) { + return Parse(source).Compile() +} + +func CompileWithFeatureFlags(source string, featureFlags map[string]struct{}) (VarsEncoder, CompiledProgram, error) { + return Parse(source).CompileWithFeatureFlags(featureFlags) +} + +var DecodeCompiledProgram = vm.DecodeProgram + +// BytecodeVersion is the version of the bytecode wire format a compiled +// program or an encoded Vars was written with — major.minor, versioned +// independently of the library itself: a new library release does not imply a +// new bytecode version. A reader accepts a blob of its own major with a minor +// no newer than its own (BytecodeVersion.CanRead); anything else the decoders +// reject with UnsupportedBytecodeVersionError. +// +// CurrentBytecodeVersion is what this build's Compile and Encode write and the +// newest it can execute. A host that stores bytecode compiled by one build and +// executes it with another can compare it against the stored blob's version +// before trusting the bytecode to run; PeekCompiledProgramVersion and +// PeekVarsVersion read that version from the raw bytes without decoding the +// rest, and without checking that this build can read it. +type ( + BytecodeVersion = vm.BytecodeVersion + UnsupportedBytecodeVersionError = vm.UnsupportedBytecodeVersionError +) + +var ( + CurrentBytecodeVersion = vm.CurrentBytecodeVersion + PeekCompiledProgramVersion = vm.PeekProgramVersion + PeekVarsVersion = vm.PeekVarsVersion +) + +// VerifyCompiledProgram statically checks that a program is safe to execute: +// ExecVm assumes well-formed bytecode and will panic rather than error on a +// program that is not. Compile's output always is, so this is for programs that +// came from somewhere else — DecodeCompiledProgram, most obviously. +// +// VerifyCompiledProgramWithVars additionally checks the program against the vars +// it will be given; prefer it whenever vars are in play, since a program that +// loads a variable is only safe against a pool that actually has it. On +// success it also returns a VerifiedVarsInfo: a caller that reuses the same +// compiled Program across many calls (e.g. an LRU cache keyed on the compiled +// bytes) can keep this and use its CheckVars method to skip re-running +// verification — a whole-program static pass — when it sees a Vars shape it +// already knows is compatible. +type VerifiedVarsInfo = vm.VerifiedVarsInfo + +var ( + VerifyCompiledProgram = vm.Verify + VerifyCompiledProgramWithVars = vm.VerifyWithVars +) + +// VM execution error types, aliased so ExecVm callers can classify failures +// with errors.As without reaching into internal packages — the same pattern as +// the interpreter's error types above. The Vm prefix keeps them apart from the +// interpreter's MissingFundsErr/NegativeAmountErr, which are different types +// with different fields. +type ( + VmMissingFundsError = vm.MissingFundsError + VmNegativeAmountError = vm.NegativeAmountError + VmNegativeBalanceError = vm.NegativeBalanceError + VmAssetMismatchError = vm.AssetMismatchError + VmInvalidAllotmentSum = vm.InvalidAllotmentSum + VmNegativePortionError = vm.NegativePortionError + VmDivideByZeroError = vm.DivideByZeroError + VmInvalidAccountName = vm.InvalidAccountName + VmInvalidColor = vm.InvalidColor + VmInvalidScope = vm.InvalidScope + VmCannotCastScopedAccountToString = vm.CannotCastScopedAccountToString + VmInvalidUncappedSource = vm.InvalidUncappedSource + VmMetadataNotFoundError = vm.MetadataNotFoundError + VmBadMetaValueError = vm.BadMetaValueError + VmInternalError = vm.InternalError + VmInvalidPostingError = vm.InvalidPostingError + VmStoreError = vm.StoreError +) + +func ExecVm[S VMStore](ctx context.Context, machine *Vm, vars *Vars, store S) (ExecutionResult, error) { + res, execErr := vm.Exec(ctx, machine, vars, store) + if execErr != nil { + return ExecutionResult{}, execErr + } + + // Postings share one type (funds.Posting), scope fields included, so they + // pass through unchanged. Metadata is normalized to the interpreter's + // contract: non-nil maps/slices, account rows in the SetAccountsMetadata + // shape. + txMeta := res.Metadata + if txMeta == nil { + txMeta = Metadata{} + } + accountsMeta := make(SetAccountsMetadata, 0, len(res.AccountsMetadata)) + for _, e := range res.AccountsMetadata { + accountsMeta = append(accountsMeta, SetAccountMetadataRow{ + Account: e.Account, + Key: e.Key, + Value: e.Value, + Scope: e.Scope, + }) + } + + return ExecutionResult{ + Postings: res.Postings, + Metadata: txMeta, + AccountsMetadata: accountsMeta, + }, nil +} diff --git a/numscript_test.go b/numscript_test.go index eaab4a45..22665226 100644 --- a/numscript_test.go +++ b/numscript_test.go @@ -637,3 +637,60 @@ func TestResolveDependenciesPublicAPI(t *testing.T) { {Account: "bob", Asset: "USD"}: {}, }, deps.AccountsWrites) } + +// vmTestStore is a minimal VMStore for the ExecVm contract tests below. +type vmTestStore map[string]int64 + +func (s vmTestStore) GetBalance(_ context.Context, account, _, _, _ string) (*big.Int, error) { + return big.NewInt(s[account]), nil +} + +func (s vmTestStore) GetMetadata(_ context.Context, _, _, _ string) (string, bool, error) { + return "", false, nil +} + +// ExecVm must report the same contract Run does: postings, tx metadata and +// account metadata rows — not postings alone. +func TestExecVmReportsMetadata(t *testing.T) { + enc, program, err := numscript.Compile(`send [COIN 30] ( + source = @src + destination = @dest +) + +set_tx_meta("k", "v") +set_account_meta(@dest, "owner", "alice")`) + require.NoError(t, err) + + vars, encErr := enc.Encode(nil) + require.NoError(t, encErr) + + res, execErr := numscript.ExecVm(context.Background(), numscript.NewVm(program), &vars, vmTestStore{"src": 100}) + require.Nil(t, execErr) + + require.Equal(t, []numscript.Posting{ + {Source: "src", Destination: "dest", Asset: "COIN", Amount: big.NewInt(30)}, + }, res.Postings) + require.Equal(t, numscript.Metadata{"k": "v"}, res.Metadata) + require.Equal(t, numscript.SetAccountsMetadata{ + {Account: "dest", Key: "owner", Value: "alice"}, + }, res.AccountsMetadata) +} + +// ExecVm's failures must be classifiable through the public aliases, without +// reaching into internal packages. +func TestExecVmErrorsAreClassifiable(t *testing.T) { + enc, program, err := numscript.Compile(`send [COIN 30] ( + source = @src + destination = @dest +)`) + require.NoError(t, err) + + vars, encErr := enc.Encode(nil) + require.NoError(t, encErr) + + _, execErr := numscript.ExecVm(context.Background(), numscript.NewVm(program), &vars, vmTestStore{"src": 10}) + require.NotNil(t, execErr) + + var missingFunds numscript.VmMissingFundsError + require.True(t, errors.As(execErr, &missingFunds)) +}