From 5790bef0c31cdc5572efc2b889dca06d8c4cedf4 Mon Sep 17 00:00:00 2001 From: ascandone Date: Fri, 25 Sep 2026 17:22:56 +0200 Subject: [PATCH] feat: numscript compiler, public compile/exec API and CLI The numscript-to-IR compiler (internal/compiler) with its own bidirectional checker (internal/typecheck) and the shared builtin names (internal/builtins); the public API in numscript.go (Compile, VarsEncoder, NewVm, ExecVm, program/vars encode-decode, opt-in verification) and the CLI wiring. Cut from feat/vm's final tree state. Two cleanups called out in the series plan: internal/analysis' builtin name constants now alias internal/builtins instead of re-declaring the literals, and typecheck's package doc no longer claims analysis consumes it (analysis keeps its own unification-based checker). --- compiler-architecture.md | 531 ++++++ internal/analysis/check.go | 19 +- internal/builtins/builtins.go | 12 + internal/cmd/assemble.go | 116 ++ internal/cmd/bytecode_run.go | 276 +++ internal/cmd/bytecode_run_test.go | 292 ++++ internal/cmd/root.go | 9 + internal/cmd/run.go | 48 +- internal/compiler/bench_test.go | 335 ++++ internal/compiler/compile_error_test.go | 243 +++ internal/compiler/compiler.go | 1492 +++++++++++++++++ internal/compiler/compiler_error.go | 146 ++ internal/compiler/compiler_example_test.go | 85 + internal/compiler/compiler_test.go | 821 +++++++++ internal/compiler/e2e_test.go | 1344 +++++++++++++++ internal/compiler/fuzz_mutate_test.go | 111 ++ internal/compiler/scripts_test.go | 176 ++ .../fuzz/FuzzMutatedBytecode/db56cb61b1ca3263 | 2 + internal/compiler/value.go | 38 + internal/compiler/vars_e2e_test.go | 154 ++ internal/compiler/vars_encoder.go | 89 + internal/compiler/verify_corpus_test.go | 51 + internal/typecheck/typecheck.go | 447 +++++ internal/typecheck/typecheck_test.go | 81 + numscript.go | 67 + 25 files changed, 6956 insertions(+), 29 deletions(-) create mode 100644 compiler-architecture.md create mode 100644 internal/builtins/builtins.go create mode 100644 internal/cmd/assemble.go create mode 100644 internal/cmd/bytecode_run.go create mode 100644 internal/cmd/bytecode_run_test.go create mode 100644 internal/compiler/bench_test.go create mode 100644 internal/compiler/compile_error_test.go create mode 100644 internal/compiler/compiler.go create mode 100644 internal/compiler/compiler_error.go create mode 100644 internal/compiler/compiler_example_test.go create mode 100644 internal/compiler/compiler_test.go create mode 100644 internal/compiler/e2e_test.go create mode 100644 internal/compiler/fuzz_mutate_test.go create mode 100644 internal/compiler/scripts_test.go create mode 100644 internal/compiler/testdata/fuzz/FuzzMutatedBytecode/db56cb61b1ca3263 create mode 100644 internal/compiler/value.go create mode 100644 internal/compiler/vars_e2e_test.go create mode 100644 internal/compiler/vars_encoder.go create mode 100644 internal/compiler/verify_corpus_test.go create mode 100644 internal/typecheck/typecheck.go create mode 100644 internal/typecheck/typecheck_test.go diff --git a/compiler-architecture.md b/compiler-architecture.md new file mode 100644 index 00000000..938e356e --- /dev/null +++ b/compiler-architecture.md @@ -0,0 +1,531 @@ +# 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: + +``` +INT_ADD 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 (int min, string add, int add, portion sub, etc) +- a few domain instructions which call the `runtime.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. +- conditional jumps (`JMP_IF_ZERO`), which can only jump forward (so that the vm always halts by design) +- a `MK_ALLOTMENT` instruction which computes the allotment-related calculations +- constant pool loading instructions: `LOAD_STR(dest:u8, idx:u16)`, which performs `str_regs[dest] = program.str_pool[idx]`, and `LOAD_INT`. +- `LOAD_VAR_STRING(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 + +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 a format version and the number of sections. The version lets a decoder reject a payload encoded by a newer, incompatible version instead of silently misreading it. + +``` +| "NUMB" 4 B | magic ++-------------------------------+ +| version : u16 2 B | header +| 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 + +### 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 ++-------------------------------+ +| version : u16 2 B | header +| 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 + +> [!NOTE] +> This isn't yet implemented in the `feat/exp/vm` branch. There is a branch with a POC of those checks. + +Even if there are bugs in the compiler, we can analyse the bytecode to prove statically that the bytecode can't make the vm crash, that the computation always halts (the instruction set is designed so that this is a decidable problem). We can also prova statically most of the interesting properties that ensure that the bytecode isn't resulting in undefined behaviour. +Some of the examples are: + +- No undefined opcodes. Ensures no panic +- Extended instructions aren't truncated. Ensures no panic +- Const idx doesn't overflow the const pool array. Ensures no panic +- Var idx doesn't overflow the vars pool array. Ensures no panic +- We don't overflow the max register declared by the compiler output. Ensures no panic +- Only jump forward. Ensures termination +- No read before write (undefined behaviour). This ensures we can re-use vm instances. Note this has to be checked on every path (including possible jumps) + +Vm is simple enough that we can easily audit every line of code that could panic (e.g. array access), and perform static checks on bytecode. + +The static check is optional: the compiler should emit valid bytecode anyway. +Still, we can use this as a sanity-check right after program is compiled, or after the raft node receives the bytes payload, to make sure nothing went wrong in the meanwhile. + +## 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. + +``` +// 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 yet implemented in the `feat/exp/vm` branch. 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/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..d75d1d4a --- /dev/null +++ b/internal/cmd/bytecode_run.go @@ -0,0 +1,276 @@ +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, + }, 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..ba5a14af --- /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{}, "a 'remaining' clause should be the last 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..a29136d3 --- /dev/null +++ b/internal/compiler/compiler.go @@ -0,0 +1,1492 @@ +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/utils" + "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 + } + portions[i] = p + case *parser.RemainingAllotment: + if remainingIdx != -1 { + return nil, DuplicateRemaining{Range: al.Range} + } + remainingIdx = i + default: + utils.NonExhaustiveMatchPanic[any](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 +} + +// 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: mon.Asset, Right: *st.currentAssetReg}) + 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} + } + 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: + panic("TODO compileExpr binary op " + string(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: + panic("TODO compileExpr prefix op " + string(expr.Operator)) + } + + case *parser.FnCall: + return st.compileFnCall(expr, false) + + default: + return utils.NonExhaustiveMatchPanic[ir.Reg](expr), nil + } +} + +// 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: + panic("TODO compileExpr fn call " + 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: + panic("TODO compileMonetaryExpr binary op " + string(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 { + panic("TODO compileMonetaryExpr prefix op " + string(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 utils.NonExhaustiveMatchPanic[monetaryValue](expr), nil + } +} + +// 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: + panic("TODO compileMonetaryExpr fn call " + 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 + } + + var innerCapReg ir.Reg + if capReg == nil { + innerCapReg = clauseCapIntReg + } else { + innerCapReg = st.minInt(clauseCapIntReg, *capReg) + } + + 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, + } + }) + + endLabel := st.FreshLabel("inorder_end") + inorderCap := st.PushWithDest(func(dest ir.Reg) ir.Instr { + return ir.UnaryOp{ + Op: ir.OpIntCopy{}, + Arg: *capReg, + Dest: dest, + } + }) + + 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, + }) + st.jmpIfAmountZero(inorderCap, endLabel) + } + } + st.Push(ir.LabelMarker{ + Label: endLabel, + }) + 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 utils.NonExhaustiveMatchPanic[ir.Reg](src), nil + } +} + +func (st *state) compileSourceWithRequiredAmount( + capReg ir.Reg, + src parser.Source, +) (ir.Reg, CompilerError) { + got, err := st.compileSource(&capReg, 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} + }) + for _, clause := range dest.Clauses { + capAmtReg, err := st.compileCapAmount(clause.Cap) + if err != nil { + return err + } + // mirrors internal/interpreter's sendTo, *parser.DestinationInorder + // case: max(min(cap, remaining), 0), so 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) + 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}) + } + + return st.compileKeptOrDestination(dest.Remaining, pulledAmtReg, remaining) + + default: + utils.NonExhaustiveMatchPanic[any](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: + utils.NonExhaustiveMatchPanic[any](keptOrDest) + } + + return nil +} + +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 utils.NonExhaustiveMatchPanic[ir.Reg](sentValue), nil + } + +} + +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: + utils.NonExhaustiveMatchPanic[any](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 utils.NonExhaustiveMatchPanic[CompilerError](stmt.Caller.Name) + } + + default: + return utils.NonExhaustiveMatchPanic[CompilerError](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: + panic("TODO meta value of 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 { + st.compileExternalVar(decl) + return nil + } + 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: + panic("unexpected meta var type: " + 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) { + 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: + panic("unexpected var type: " + decl.Type.Name) + } +} + +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..a8cb77de --- /dev/null +++ b/internal/compiler/compiler_error.go @@ -0,0 +1,146 @@ +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 + } +) + +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 (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 "a 'remaining' clause should be the last 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) +} + +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) +) 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..610b5d8c --- /dev/null +++ b/internal/compiler/compiler_test.go @@ -0,0 +1,821 @@ +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 = "src" + $r4 = 0 + $r5 = str_eq($r3, $r0) + jmp_if_false($r5, #not_world_0) + $r6 = pull_account(account: $r3, cap: $r2) + jmp(#pull_end_1) +#not_world_0 + $r6 = pull_account(account: $r3, cap: $r2, overdraft: $r4) +#pull_end_1 + check_enough_funds($r6, $r2) + $r7 = "dest" + send_to_account(account: $r7) +`)) +} + +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 = "src" + $r6 = 0 + $r7 = str_eq($r5, $r0) + jmp_if_false($r7, #not_world_0) + $r8 = pull_account(account: $r5, cap: $r4) + jmp(#pull_end_1) +#not_world_0 + $r8 = pull_account(account: $r5, cap: $r4, overdraft: $r6) +#pull_end_1 + check_enough_funds($r8, $r4) + $r9 = "dest" + send_to_account(account: $r9) +`)) +} + +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 = "src" + $r6 = 0 + $r7 = str_eq($r5, $r0) + jmp_if_false($r7, #not_world_0) + $r8 = pull_account(account: $r5, cap: $r4) + jmp(#pull_end_1) +#not_world_0 + $r8 = pull_account(account: $r5, cap: $r4, overdraft: $r6) +#pull_end_1 + check_enough_funds($r8, $r4) + $r9 = "dest" + send_to_account(account: $r9) +`)) +} + +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 = "src" + $r7 = 0 + $r8 = str_eq($r6, $r0) + jmp_if_false($r8, #not_world_0) + $r9 = pull_account(account: $r6, cap: $r5) + jmp(#pull_end_1) +#not_world_0 + $r9 = pull_account(account: $r6, cap: $r5, overdraft: $r7) +#pull_end_1 + check_enough_funds($r9, $r5) + $r10 = "dest" + send_to_account(account: $r10) +`)) +} + +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 = "src" + $r7 = 0 + $r8 = str_eq($r6, $r0) + jmp_if_false($r8, #not_world_0) + $r9 = pull_account(account: $r6, cap: $r5) + jmp(#pull_end_1) +#not_world_0 + $r9 = pull_account(account: $r6, cap: $r5, overdraft: $r7) +#pull_end_1 + check_enough_funds($r9, $r5) + $r10 = "dest" + send_to_account(account: $r10) +`)) +} + +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 = "src" + $r5 = 0 + $r6 = str_eq($r4, $r0) + jmp_if_false($r6, #not_world_0) + $r7 = pull_account(account: $r4, cap: $r2) + jmp(#pull_end_1) +#not_world_0 + $r7 = pull_account(account: $r4, cap: $r2, overdraft: $r5) +#pull_end_1 + check_enough_funds($r7, $r2) + $r8 = "dest" + send_to_account(account: $r8) +`)) +} + +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 = "src" + $r5 = 0 + $r6 = str_eq($r4, $r0) + jmp_if_false($r6, #not_world_0) + $r7 = pull_account(account: $r4, cap: $r3) + jmp(#pull_end_1) +#not_world_0 + $r7 = pull_account(account: $r4, cap: $r3, overdraft: $r5) +#pull_end_1 + check_enough_funds($r7, $r3) + $r8 = "dest" + send_to_account(account: $r8) +`)) +} + +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 = "src" + $r6 = 0 + $r7 = str_eq($r5, $r0) + jmp_if_false($r7, #not_world_0) + $r8 = pull_account(account: $r5, cap: $r4) + jmp(#pull_end_1) +#not_world_0 + $r8 = pull_account(account: $r5, cap: $r4, overdraft: $r6) +#pull_end_1 + check_enough_funds($r8, $r4) + $r9 = "dest" + send_to_account(account: $r9) +`)) +} + +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 = "src" + $r5 = 0 + $r6 = str_eq($r4, $r0) + jmp_if_false($r6, #not_world_0) + $r7 = pull_account(account: $r4, cap: $r3) + jmp(#pull_end_1) +#not_world_0 + $r7 = pull_account(account: $r4, cap: $r3, overdraft: $r5) +#pull_end_1 + check_enough_funds($r7, $r3) + $r8 = "dest" + send_to_account(account: $r8) +`)) +} + +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 = "world" + $r5 = 0 + $r6 = str_eq($r4, $r0) + jmp_if_false($r6, #not_world_0) + $r7 = pull_account(account: $r4, cap: $r3) + jmp(#pull_end_1) +#not_world_0 + $r7 = pull_account(account: $r4, cap: $r3, overdraft: $r5) +#pull_end_1 + check_enough_funds($r7, $r3) + $r8 = "users" + $r9 = ":" + $r10 = ":" + $r11 = "wallet" + $r12 = add_string($r8, $r9) + $r13 = add_string($r12, $r1) + $r14 = add_string($r13, $r10) + $r15 = add_string($r14, $r11) + assert_valid_account($r15) + send_to_account(account: $r15) +`)) +} + +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 = "world" + $r5 = 0 + $r6 = str_eq($r4, $r0) + jmp_if_false($r6, #not_world_0) + $r7 = pull_account(account: $r4, cap: $r3) + jmp(#pull_end_1) +#not_world_0 + $r7 = pull_account(account: $r4, cap: $r3, overdraft: $r5) +#pull_end_1 + check_enough_funds($r7, $r3) + $r8 = "account" + $r9 = ":" + $r10 = int_to_string($r1) + $r11 = add_string($r8, $r9) + $r12 = add_string($r11, $r10) + assert_valid_account($r12) + send_to_account(account: $r12) +`)) +} + +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 = "a" + $r6 = 0 + $r7 = str_eq($r5, $r0) + jmp_if_false($r7, #not_world_1) + $r8 = pull_account(account: $r5, cap: $r4) + jmp(#pull_end_2) +#not_world_1 + $r8 = pull_account(account: $r5, cap: $r4, overdraft: $r6) +#pull_end_2 + $r3 += $r8 + $r4 -= $r8 + $r9 = is_zero($r4) + jmp_if_true($r9, #inorder_end_0) + $r10 = "b" + $r11 = 0 + $r12 = str_eq($r10, $r0) + jmp_if_false($r12, #not_world_3) + $r13 = pull_account(account: $r10, cap: $r4) + jmp(#pull_end_4) +#not_world_3 + $r13 = pull_account(account: $r10, cap: $r4, overdraft: $r11) +#pull_end_4 + $r3 += $r13 + $r4 -= $r13 + $r14 = is_zero($r4) + jmp_if_true($r14, #inorder_end_0) + $r15 = "c" + $r16 = 0 + $r17 = str_eq($r15, $r0) + jmp_if_false($r17, #not_world_5) + $r18 = pull_account(account: $r15, cap: $r4) + jmp(#pull_end_6) +#not_world_5 + $r18 = pull_account(account: $r15, cap: $r4, overdraft: $r16) +#pull_end_6 + $r3 += $r18 +#inorder_end_0 + check_enough_funds($r3, $r2) + $r19 = "dest" + send_to_account(account: $r19) +`)) +} + +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 = "a" + $r6 = 0 + $r7 = str_eq($r5, $r0) + jmp_if_false($r7, #not_world_1) + $r8 = pull_account(account: $r5, cap: $r4) + jmp(#pull_end_2) +#not_world_1 + $r8 = pull_account(account: $r5, cap: $r4, overdraft: $r6) +#pull_end_2 + $r3 += $r8 + $r4 -= $r8 + $r9 = is_zero($r4) + jmp_if_true($r9, #inorder_end_0) + $r10 = "USD/2" + $r11 = 5 + assert_same_asset($r10, $r1) + $r12 = int_copy($r11) + $r13 = lt_int($r11, $r4) + jmp_if_true($r13, #min_end_3) + $r12 = int_copy($r4) +#min_end_3 + $r14 = "b" + $r15 = 0 + $r16 = str_eq($r14, $r0) + jmp_if_false($r16, #not_world_4) + $r17 = pull_account(account: $r14, cap: $r12) + jmp(#pull_end_5) +#not_world_4 + $r17 = pull_account(account: $r14, cap: $r12, overdraft: $r15) +#pull_end_5 + $r3 += $r17 + $r4 -= $r17 + $r18 = is_zero($r4) + jmp_if_true($r18, #inorder_end_0) + $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 + $r3 += $r22 +#inorder_end_0 + check_enough_funds($r3, $r2) + $r23 = "dest" + send_to_account(account: $r23) +`)) +} + +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 = "world" + $r4 = 0 + $r5 = str_eq($r3, $r0) + jmp_if_false($r5, #not_world_0) + $r6 = pull_account(account: $r3, cap: $r2) + jmp(#pull_end_1) +#not_world_0 + $r6 = pull_account(account: $r3, cap: $r2, overdraft: $r4) +#pull_end_1 + check_enough_funds($r6, $r2) + $r7 = int_copy($r6) + $r8 = "USD/2" + $r9 = 4 + assert_same_asset($r8, $r1) + $r10 = 0 + $r11 = int_copy($r7) + $r12 = lt_int($r7, $r9) + jmp_if_true($r12, #min_end_2) + $r11 = int_copy($r9) +#min_end_2 + $r13 = int_copy($r11) + $r14 = lt_int($r10, $r11) + jmp_if_true($r14, #max_end_3) + $r13 = int_copy($r10) +#max_end_3 + $r15 = "d1" + send_to_account(account: $r15, cap: $r13) + $r7 -= $r13 + $r16 = "d2" + send_to_account(account: $r16, cap: $r7) +`)) +} + +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) + mark_push() + $r3 = "a" + $r4 = 0 + $r5 = str_eq($r3, $r0) + jmp_if_false($r5, #not_world_1) + $r6 = pull_account(account: $r3, cap: $r2) + jmp(#pull_end_2) +#not_world_1 + $r6 = pull_account(account: $r3, cap: $r2, overdraft: $r4) +#pull_end_2 + $r7 = int_copy($r6) + $r8 = $r2 - $r6 + $r9 = is_zero($r8) + jmp_if_true($r9, #oneof_end_0) + mark_rewind() + mark_push() + $r10 = "b" + $r11 = 0 + $r12 = str_eq($r10, $r0) + jmp_if_false($r12, #not_world_3) + $r13 = pull_account(account: $r10, cap: $r2) + jmp(#pull_end_4) +#not_world_3 + $r13 = pull_account(account: $r10, cap: $r2, overdraft: $r11) +#pull_end_4 + $r7 = int_copy($r13) + $r14 = $r2 - $r13 + $r15 = is_zero($r14) + jmp_if_true($r15, #oneof_end_0) + mark_rewind() + mark_push() + $r16 = "c" + $r17 = 0 + $r18 = str_eq($r16, $r0) + jmp_if_false($r18, #not_world_5) + $r19 = pull_account(account: $r16, cap: $r2) + jmp(#pull_end_6) +#not_world_5 + $r19 = pull_account(account: $r16, cap: $r2, overdraft: $r17) +#pull_end_6 + $r7 = int_copy($r19) +#oneof_end_0 + mark_commit() + check_enough_funds($r7, $r2) + $r20 = "dest" + send_to_account(account: $r20) +`)) +} + +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) + mark_push() + $r3 = "a" + $r4 = 0 + $r5 = str_eq($r3, $r0) + jmp_if_false($r5, #not_world_1) + $r6 = pull_account(account: $r3, cap: $r2) + jmp(#pull_end_2) +#not_world_1 + $r6 = pull_account(account: $r3, cap: $r2, overdraft: $r4) +#pull_end_2 + $r7 = int_copy($r6) + $r8 = $r2 - $r6 + $r9 = is_zero($r8) + jmp_if_true($r9, #oneof_end_0) + mark_rewind() + mark_push() + $r10 = "b" + $r11 = 0 + $r12 = str_eq($r10, $r0) + jmp_if_false($r12, #not_world_3) + $r13 = pull_account(account: $r10, cap: $r2) + jmp(#pull_end_4) +#not_world_3 + $r13 = pull_account(account: $r10, cap: $r2, overdraft: $r11) +#pull_end_4 + $r7 = int_copy($r13) +#oneof_end_0 + mark_commit() + check_enough_funds($r7, $r2) + $r14 = "dest" + send_to_account(account: $r14) +`)) +} + +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 = "world" + $r4 = 0 + $r5 = str_eq($r3, $r0) + jmp_if_false($r5, #not_world_0) + $r6 = pull_account(account: $r3, cap: $r2) + jmp(#pull_end_1) +#not_world_0 + $r6 = pull_account(account: $r3, cap: $r2, overdraft: $r4) +#pull_end_1 + check_enough_funds($r6, $r2) + $r7 = "USD/2" + $r8 = 4 + assert_same_asset($r7, $r1) + $r9 = int_copy($r6) + $r10 = lt_int($r6, $r8) + jmp_if_true($r10, #min_end_4) + $r9 = int_copy($r8) +#min_end_4 + $r11 = $r6 - $r9 + $r12 = is_zero($r11) + jmp_if_true($r12, #oneof_dest_clause_3) + $r13 = "b" + send_to_account(account: $r13) + jmp(#oneof_dest_end_2) +#oneof_dest_clause_3 + $r14 = "a" + send_to_account(account: $r14) + jmp(#oneof_dest_end_2) +#oneof_dest_end_2 +`)) +} + +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 = "src" + $r4 = "RED" + assert_valid_color($r4) + $r5 = 0 + $r6 = str_eq($r3, $r0) + jmp_if_false($r6, #not_world_0) + $r7 = pull_account(account: $r3, cap: $r2, color: $r4) + jmp(#pull_end_1) +#not_world_0 + $r7 = pull_account(account: $r3, cap: $r2, overdraft: $r5, color: $r4) +#pull_end_1 + check_enough_funds($r7, $r2) + $r8 = "dest" + send_to_account(account: $r8) +`)) +} + +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 = "src" + $r4 = "RED" + assert_valid_color($r4) + $r5 = pull_account(account: $r3, cap: $r2, color: $r4) + check_enough_funds($r5, $r2) + $r6 = "dest" + send_to_account(account: $r6) +`)) +} diff --git a/internal/compiler/e2e_test.go b/internal/compiler/e2e_test.go new file mode 100644 index 00000000..7e0545ad --- /dev/null +++ b/internal/compiler/e2e_test.go @@ -0,0 +1,1344 @@ +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_CompileAssembleRun exercises the whole pipeline: source -> compiler +// (IR) -> assembler (vm.Program) -> VM execution -> postings. +func TestE2E_CompileAssembleRun(t *testing.T) { + src := ` + send [USD/2 10] ( + 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), + }} + + machine := vm.NewVm(program) + res, execErr := vm.Exec(context.Background(), machine, nil, store) + require.Nil(t, execErr) + + want := []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + } + requirePostingsEqual(t, want, res.Postings) +} + +// TestE2E_Inorder exercises an inorder source { @a @b @c } end-to-end, including +// the early-exit jump: @a has 6, @b has 10, @c has 100; sending 10 pulls 6 from +// @a (cap -> 4), then 4 from @b (cap -> 0 -> jump past @c). @c is never touched. +func TestE2E_Inorder(t *testing.T) { + src := ` + send [USD/2 10] ( + source = { + @a + @b + @c + } + 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: "a", Asset: "USD/2", Color: ""}: big.NewInt(6), + {Account: "b", Asset: "USD/2", Color: ""}: big.NewInt(10), + {Account: "c", Asset: "USD/2", Color: ""}: big.NewInt(100), + }} + + machine := vm.NewVm(program) + res, execErr := vm.Exec(context.Background(), machine, nil, store) + require.Nil(t, execErr) + + want := []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(6)}, + {Source: "b", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(4)}, + } + requirePostingsEqual(t, want, res.Postings) +} + +// TestE2E_InorderWithCap exercises a capped (`max`) source inside an inorder +// end-to-end. @b holds 100 but is capped at 5, so the cap must bind: @a gives 3 +// (remaining 10->7), @b gives only 5 (not 7) -> remaining 2, @c gives 2. +func TestE2E_InorderWithCap(t *testing.T) { + src := ` + send [USD/2 10] ( + source = { + @a + max [USD/2 5] from @b + @c + } + 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: "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), + }} + + machine := vm.NewVm(program) + res, execErr := vm.Exec(context.Background(), machine, nil, store) + require.Nil(t, execErr) + + want := []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(3)}, + {Source: "b", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(5)}, + {Source: "c", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(2)}, + } + requirePostingsEqual(t, want, res.Postings) +} + +// TestE2E_InsufficientFunds checks the failure path: when the source can't cover +// the sent amount, the VM's CheckEnoughFunds must report a MissingFundsError. +func TestE2E_InsufficientFunds(t *testing.T) { + src := ` + send [USD/2 10] ( + source = @src + destination = @dest + ) + ` + + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + + // src only has 4, but 10 is required. + store := e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(4), + }} + + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, store) + require.IsType(t, vm.MissingFundsError{}, execErr) +} + +// TestE2E_DestinationInorder exercises a destination-inorder split end-to-end: +// 100 pulled from @world is distributed as `max [USD/2 30] to @x; remaining to +// @y`, so @x must get 30 and @y the remaining 70. +func TestE2E_DestinationInorder(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + max [USD/2 30] to @x + remaining to @y + } + ) + ` + + 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{}} + + machine := vm.NewVm(program) + res, execErr := vm.Exec(context.Background(), machine, nil, store) + require.Nil(t, execErr) + + want := []funds.Posting{ + {Source: "world", Destination: "x", Asset: "USD/2", Amount: big.NewInt(30)}, + {Source: "world", Destination: "y", Asset: "USD/2", Amount: big.NewInt(70)}, + } + requirePostingsEqual(t, want, res.Postings) +} + +// TestE2E_DestinationKept exercises a `kept` clause: of 100 pulled from @world, +// 30 is kept (refunded, no posting) and the remaining 70 goes to @y. +func TestE2E_DestinationKept(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + max [USD/2 30] kept + remaining to @y + } + ) + ` + + 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{}} + + machine := vm.NewVm(program) + res, execErr := vm.Exec(context.Background(), machine, nil, store) + require.Nil(t, execErr) + + // only the remaining 70 is posted; the kept 30 produces no posting + want := []funds.Posting{ + {Source: "world", Destination: "y", Asset: "USD/2", Amount: big.NewInt(70)}, + } + requirePostingsEqual(t, want, res.Postings) +} + +// TestE2E_DestinationAllotment splits the pulled amount by portions. 100 from +// @world with { 1/2 to @a; remaining to @b } => a=50, b=50. +func TestE2E_DestinationAllotment(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + 1/2 to @a + remaining to @b + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "a", Asset: "USD/2", Amount: big.NewInt(50)}, + {Source: "world", Destination: "b", Asset: "USD/2", Amount: big.NewInt(50)}, + }, postings) +} + +// TestE2E_DestinationAllotmentThirds exercises the floor-then-distribute-leftover +// rounding: 100 by thirds => 34, 33, 33 (the leftover unit goes to the earliest). +func TestE2E_DestinationAllotmentThirds(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + 1/3 to @a + 1/3 to @b + remaining to @c + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "a", Asset: "USD/2", Amount: big.NewInt(34)}, + {Source: "world", Destination: "b", Asset: "USD/2", Amount: big.NewInt(33)}, + {Source: "world", Destination: "c", Asset: "USD/2", Amount: big.NewInt(33)}, + }, postings) +} + +// TestE2E_SourceAllotment splits the requested amount across sub-sources, pulling +// each exactly. 100 with { 1/4 from @s1; remaining from @s2 } => 25 from s1, 75 +// from s2. +func TestE2E_SourceAllotment(t *testing.T) { + src := ` + send [USD/2 100] ( + source = { + 1/4 from @s1 + remaining from @s2 + } + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "s1", Asset: "USD/2", Color: ""}: big.NewInt(1000), + {Account: "s2", Asset: "USD/2", Color: ""}: big.NewInt(1000), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "s1", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(25)}, + {Source: "s2", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(75)}, + }, postings) +} + +// TestE2E_SourceAllotmentThirds checks the rounding split on the source side too: +// 100 by thirds => 34, 33, 33. +func TestE2E_SourceAllotmentThirds(t *testing.T) { + src := ` + send [USD/2 100] ( + source = { + 1/3 from @a + 1/3 from @b + remaining from @c + } + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2", Color: ""}: big.NewInt(1000), + {Account: "b", Asset: "USD/2", Color: ""}: big.NewInt(1000), + {Account: "c", Asset: "USD/2", Color: ""}: big.NewInt(1000), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(34)}, + {Source: "b", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(33)}, + {Source: "c", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(33)}, + }, postings) +} + +// TestE2E_SourceAllotmentInsufficient: a sub-source must provide its exact share, +// else MissingFunds. s1 only has 10 but its 1/2 share of 100 is 50. +func TestE2E_SourceAllotmentInsufficient(t *testing.T) { + src := ` + send [USD/2 100] ( + source = { + 1/2 from @s1 + remaining from @s2 + } + 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: "s1", Asset: "USD/2", Color: ""}: big.NewInt(10), + {Account: "s2", Asset: "USD/2", Color: ""}: big.NewInt(1000), + }} + machine := vm.NewVm(program) + _, execErr := vm.Exec(context.Background(), machine, nil, store) + require.IsType(t, vm.MissingFundsError{}, execErr) +} + +// 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.IsType(t, vm.InvalidAllotmentSum{}, execErr) + allotErr := execErr.(vm.InvalidAllotmentSum) + require.Equal(t, "4/3", allotErr.ActualSum.String()) + require.EqualError(t, allotErr, "invalid allotment: portions must sum to 1, got 4/3") +} + +// 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.IsType(t, vm.InvalidAllotmentSum{}, execErr) +} + +// TestE2E_AllotmentExactNoRemaining: a no-remaining allotment summing to exactly +// 1 is valid. +func TestE2E_AllotmentExactNoRemaining(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + 1/4 to @a + 3/4 to @b + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "a", Asset: "USD/2", Amount: big.NewInt(25)}, + {Source: "world", Destination: "b", Asset: "USD/2", Amount: big.NewInt(75)}, + }, postings) +} + +// TestE2E_AllotmentRemainingOnly: `{ remaining to @dest }` is 100% (leftover = 1), +// which must remain valid (a `< 1` check would wrongly reject it). +func TestE2E_AllotmentRemainingOnly(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + remaining to @dest + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(100)}, + }, postings) +} + +func TestE2E_IntAddition(t *testing.T) { + src := ` + send [USD/2 4 + 6] ( + 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), + }} + + res, execErr := vm.Exec(context.Background(), vm.NewVm(program), nil, store) + require.Nil(t, execErr) + + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, res.Postings) +} + +func TestE2E_IntSubtraction(t *testing.T) { + src := ` + send [USD/2 16 - 6] ( + 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), + }} + + res, execErr := vm.Exec(context.Background(), vm.NewVm(program), nil, store) + require.Nil(t, execErr) + + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, res.Postings) +} + +func TestE2E_MonetaryAddition(t *testing.T) { + src := ` + vars { + monetary $a = [USD/2 3] + monetary $b = [USD/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), + }} + + res, execErr := vm.Exec(context.Background(), vm.NewVm(program), nil, store) + require.Nil(t, execErr) + + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, res.Postings) +} + +func TestE2E_MonetarySubtraction(t *testing.T) { + src := ` + vars { + monetary $a = [USD/2 30] + monetary $b = [USD/2 20] + } + send $a - $b ( + source = @src + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) +} + +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.IsType(t, vm.AssetMismatchError{}, 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.IsType(t, vm.AssetMismatchError{}, execErr) +} + +func TestE2E_GetAmount(t *testing.T) { + src := ` + #![feature("experimental-get-amount-function")] + vars { + monetary $m = [USD/2 42] + number $n = get_amount($m) + } + send [USD/2 $n] ( + 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), + }} + + res, execErr := vm.Exec(context.Background(), vm.NewVm(program), nil, store) + require.Nil(t, execErr) + + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(42)}, + }, res.Postings) +} + +func TestE2E_GetAsset(t *testing.T) { + src := ` + #![feature("experimental-get-asset-function")] + vars { + monetary $m = [USD/2 42] + asset $a = get_asset($m) + } + send [$a 10] ( + 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), + }} + + res, execErr := vm.Exec(context.Background(), vm.NewVm(program), nil, store) + require.Nil(t, execErr) + + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, res.Postings) +} + +func TestE2E_PrefixMinusNumber(t *testing.T) { + // $neg = -10 (prefix on literal), $pos = -$neg = 10 (prefix on var) + src := ` + vars { + number $neg = -10 + number $pos = -$neg + } + send [USD/2 $pos] ( + source = @src + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) +} + +func TestE2E_PrefixMinusMonetary(t *testing.T) { + // $neg_mon = [USD/2 -10], -$neg_mon = [USD/2 10] + src := ` + vars { + monetary $neg_mon = [USD/2 -10] + monetary $pos_mon = -$neg_mon + } + send $pos_mon ( + source = @src + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) +} + +func TestE2E_Balance(t *testing.T) { + // $bal = balance(@src, USD/2) reads @src's balance (100), then sends it all + src := ` + vars { + monetary $bal = balance(@src, USD/2) + } + send $bal ( + source = @src + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(100)}, + }, postings) +} + +func TestE2E_AccountInterpolation(t *testing.T) { + // destination = @users:<$id>:wallet, with $id = "alice" + src := ` + #![feature("experimental-account-interpolation")] + vars { + string $id = "alice" + } + send [USD/2 10] ( + source = @world + destination = @users:$id:wallet + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "users:alice:wallet", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) +} + +func TestE2E_AccountInterpolationInt(t *testing.T) { + // destination = @account:<$n>, with $n = 42 + src := ` + #![feature("experimental-account-interpolation")] + vars { + number $n = 42 + } + send [USD/2 10] ( + source = @world + destination = @account:$n + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "account:42", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) +} + +func TestE2E_BoundedOverdraft(t *testing.T) { + src := ` + send [USD/2 42] ( + source = @a allowing overdraft up to [USD/2 5] + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2", Color: ""}: big.NewInt(40), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(42)}, + }, postings) +} + +func TestE2E_NestedDestination(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + 1/2 to { + max [USD/2 10] to @x + remaining to @a + } + remaining to @b + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "x", Asset: "USD/2", Amount: big.NewInt(10)}, + {Source: "world", Destination: "a", Asset: "USD/2", Amount: big.NewInt(40)}, + {Source: "world", Destination: "b", Asset: "USD/2", Amount: big.NewInt(50)}, + }, postings) +} + +func TestE2E_SendAll(t *testing.T) { + src := `send [USD/2 *] (source = @a destination = @dest)` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2", Color: ""}: big.NewInt(30), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(30)}, + }, postings) +} + +func TestE2E_UncappedBoundedOverdraft(t *testing.T) { + src := ` + send [USD/2 *] ( + source = @a allowing overdraft up to [USD/2 5] + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2", Color: ""}: big.NewInt(40), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(45)}, + }, postings) +} + +func TestE2E_SendAllMultiSource(t *testing.T) { + // unbounded inorder: pull everything from each source in order and sum it + src := ` + send [USD/2 *] ( + source = { + @a + max [USD/2 5] from @b + @c + } + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2", Color: ""}: big.NewInt(10), + {Account: "b", Asset: "USD/2", Color: ""}: big.NewInt(100), + {Account: "c", Asset: "USD/2", Color: ""}: big.NewInt(7), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + {Source: "b", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(5)}, + {Source: "c", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(7)}, + }, postings) +} + +func TestE2E_SendAllNegativeOverdraftBoundClamped(t *testing.T) { + // a negative overdraft bound is clamped to 0 in the unbounded path, so only + // the positive balance is sent (mirrors the interpreter's NonNeg). + src := ` + send [COIN *] ( + source = @s allowing overdraft up to [COIN -10] + destination = @dest + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "s", Asset: "COIN", Color: ""}: big.NewInt(1), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "s", Destination: "dest", Asset: "COIN", Amount: big.NewInt(1)}, + }, postings) +} + +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.IsType(t, vm.AssetMismatchError{}, 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.IsType(t, vm.AssetMismatchError{}, execErr) +} + +func TestE2E_Save(t *testing.T) { + // save 30 of @a's 100, so the send-all only takes the remaining 70 + src := ` + save [USD/2 30] from @a + send [USD/2 *] (source = @a destination = @dest) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(70)}, + }, postings) +} + +func TestE2E_InternalVar(t *testing.T) { + src := ` + vars { account $acc = @src } + send [USD/2 10] (source = $acc destination = @dest) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) +} + +func TestE2E_OverdraftFunction(t *testing.T) { + src := ` + #![feature("experimental-overdraft-function")] + vars { monetary $od = overdraft(@acc, USD/2) } + send $od (source = @world destination = @dest) + ` + // negative balance -> overdraft is the debt + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "acc", Asset: "USD/2", Color: ""}: big.NewInt(-100), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(100)}, + }, postings) + + // positive balance -> overdraft is 0, nothing sent + postings = runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "acc", Asset: "USD/2", Color: ""}: big.NewInt(100), + }}) + requirePostingsEqual(t, []funds.Posting{}, postings) +} + +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.IsType(t, vm.NegativeBalanceError{}, 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.IsType(t, vm.DivideByZeroError{}, execErr) +} + +func TestE2E_ColoredSource(t *testing.T) { + store := e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "src", Asset: "COIN", Color: ""}: big.NewInt(100), + {Account: "src", Asset: "COIN", Color: "RED"}: big.NewInt(30), + }} + + got := runE2E(t, ` + #![feature("experimental-asset-colors")] + send [COIN 10] ( + source = @src \ "RED" + destination = @dest + ) + `, store) + + want := []funds.Posting{ + {Source: "src", Destination: "dest", Asset: "COIN", Color: "RED", Amount: big.NewInt(10)}, + } + requirePostingsEqual(t, want, got) +} + +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.IsType(t, vm.InvalidColor{}, 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 runE2E(t *testing.T, src string, store e2eStore) []funds.Posting { + t.Helper() + parsed := parser.Parse(src) + require.Empty(t, parsed.Errors) + _, program, cErr := compiler.Compile(parsed.Value, nil) + require.Nil(t, cErr) + machine := vm.NewVm(program) + res, execErr := vm.Exec(context.Background(), machine, nil, store) + require.Nil(t, execErr) + return res.Postings +} + +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) + } +} + +// --- Allotment rounding, ported from internal/runtime/allotment_test.go ----- +// These pinned funds.MakeAllotment before the split was lowered into pure +// instructions; they now pin the compiler's lowering of it. + +// A two-unit shortfall: 1/6,1/6,4/6 of 100 floors to 16,16,66 (sum 98), so the +// first two shares each get one unit back. +func TestE2E_AllotmentLeftoverTwoUnits(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + 1/6 to @a + 1/6 to @b + remaining to @c + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "a", Asset: "USD/2", Amount: big.NewInt(17)}, + {Source: "world", Destination: "b", Asset: "USD/2", Amount: big.NewInt(17)}, + {Source: "world", Destination: "c", Asset: "USD/2", Amount: big.NewInt(66)}, + }, postings) +} + +// An odd amount split in half: 7 -> 3,3 (sum 6), leftover unit to the earliest. +func TestE2E_AllotmentHalvesOfOddAmount(t *testing.T) { + src := ` + send [USD/2 7] ( + source = @world + destination = { + 1/2 to @a + remaining to @b + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "a", Asset: "USD/2", Amount: big.NewInt(4)}, + {Source: "world", Destination: "b", Asset: "USD/2", Amount: big.NewInt(3)}, + }, postings) +} + +// A single whole share: the lowering emits no fixup blocks at all for n == 1. +func TestE2E_AllotmentSinglePortionWhole(t *testing.T) { + src := ` + send [USD/2 100] ( + source = @world + destination = { + remaining to @a + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "a", Asset: "USD/2", Amount: big.NewInt(100)}, + }, postings) +} + +// Percentages that divide exactly: no leftover, so no share is adjusted. +func TestE2E_AllotmentPercentagesDivideExactly(t *testing.T) { + src := ` + send [USD/2 10000] ( + source = @world + destination = { + 19/100 to @a + remaining to @b + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "a", Asset: "USD/2", Amount: big.NewInt(1900)}, + {Source: "world", Destination: "b", Asset: "USD/2", Amount: big.NewInt(8100)}, + }, postings) +} + +// Sevenths of 1001 floor awkwardly (143 + 286 + 572 = 1001 exactly here), the +// point being that the shares must always sum back to the amount. +func TestE2E_AllotmentPartsSumToAmount(t *testing.T) { + src := ` + send [USD/2 1001] ( + source = @world + destination = { + 1/7 to @a + 2/7 to @b + remaining to @c + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + + total := new(big.Int) + for _, p := range postings { + total.Add(total, p.Amount) + } + require.Zero(t, total.Cmp(big.NewInt(1001)), "shares sum to %s, want 1001 (%v)", total, postings) +} + +// Beyond int64: ~1e27+1 split in half, the odd unit going to the earliest share. +func TestE2E_AllotmentBeyondInt64(t *testing.T) { + src := ` + send [USD/2 1000000000000000000000000001] ( + source = @world + destination = { + 1/2 to @a + remaining to @b + } + ) + ` + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{}}) + + amount, _ := new(big.Int).SetString("1000000000000000000000000001", 10) + half := new(big.Int).Div(amount, big.NewInt(2)) + requirePostingsEqual(t, []funds.Posting{ + {Source: "world", Destination: "a", Asset: "USD/2", Amount: new(big.Int).Add(half, big.NewInt(1))}, + {Source: "world", Destination: "b", Asset: "USD/2", Amount: half}, + }, postings) +} + +// TestE2E_Oneof runs a compiled `oneof` end to end, which the IR snapshot tests +// cannot: they check the emitted instructions, not that executing them backtracks +// correctly. What matters here is that the single mark_pop at the join is reached +// on every path — the branch that covered the amount jumps straight to it, and the +// last branch falls through to it. +func TestE2E_Oneof(t *testing.T) { + src := ` + #![feature("experimental-oneof")] + send [USD/2 10] ( + source = oneof { + @a + @b + @c + } + destination = @dest + ) + ` + + t.Run("first branch covers it", func(t *testing.T) { + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2"}: big.NewInt(10), + {Account: "b", Asset: "USD/2"}: big.NewInt(10), + {Account: "c", Asset: "USD/2"}: big.NewInt(10), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) + }) + + t.Run("backtracks past a short branch", func(t *testing.T) { + // @a can only cover 3 of the 10, so its partial funding is discarded whole + // rather than combined with @b's + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2"}: big.NewInt(3), + {Account: "b", Asset: "USD/2"}: big.NewInt(10), + {Account: "c", Asset: "USD/2"}: big.NewInt(10), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "b", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) + }) + + t.Run("backtracks twice, to the last branch", func(t *testing.T) { + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2"}: big.NewInt(3), + {Account: "b", Asset: "USD/2"}: big.NewInt(9), + {Account: "c", Asset: "USD/2"}: big.NewInt(10), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "c", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) + }) + + t.Run("no branch covers it", func(t *testing.T) { + // the last branch is not rewound, so check_enough_funds reports what it got + 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: "a", Asset: "USD/2"}: big.NewInt(3), + {Account: "b", Asset: "USD/2"}: big.NewInt(4), + {Account: "c", Asset: "USD/2"}: big.NewInt(5), + }, + }) + require.IsType(t, vm.MissingFundsError{}, execErr) + }) +} + +// A oneof nested inside another must keep the two regions independent: the inner +// backtrack may not discard what the outer branch already pulled. +func TestE2E_OneofNested(t *testing.T) { + src := ` + #![feature("experimental-oneof")] + send [USD/2 10] ( + source = oneof { + { + @a + oneof { @b @c } + } + @d + } + destination = @dest + ) + ` + + t.Run("inner backtrack keeps the outer branch's funds", func(t *testing.T) { + // @a gives 4, so the inner oneof needs 6: @b has only 5 -> rewind -> @c + // covers it. @a's 4 must survive the inner rewind. + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2"}: big.NewInt(4), + {Account: "b", Asset: "USD/2"}: big.NewInt(5), + {Account: "c", Asset: "USD/2"}: big.NewInt(6), + {Account: "d", Asset: "USD/2"}: big.NewInt(10), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "a", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(4)}, + {Source: "c", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(6)}, + }, postings) + }) + + t.Run("outer backtrack discards the whole inner region too", func(t *testing.T) { + // the inorder branch tops out at 4+6=10... but with @c at 5 it reaches 9, + // so the outer oneof rewinds @a and @c together and takes @d instead + postings := runE2E(t, src, e2eStore{balances: map[funds.PairKey]*big.Int{ + {Account: "a", Asset: "USD/2"}: big.NewInt(4), + {Account: "b", Asset: "USD/2"}: big.NewInt(3), + {Account: "c", Asset: "USD/2"}: big.NewInt(5), + {Account: "d", Asset: "USD/2"}: big.NewInt(10), + }}) + requirePostingsEqual(t, []funds.Posting{ + {Source: "d", Destination: "dest", Asset: "USD/2", Amount: big.NewInt(10)}, + }, postings) + }) +} diff --git a/internal/compiler/fuzz_mutate_test.go b/internal/compiler/fuzz_mutate_test.go new file mode 100644 index 00000000..be81278e --- /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 vm.VerifyWithVars(prog, vars) != 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..d8c64fd6 --- /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}, 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/typecheck/typecheck.go b/internal/typecheck/typecheck.go new file mode 100644 index 00000000..88e277fc --- /dev/null +++ b/internal/typecheck/typecheck.go @@ -0,0 +1,447 @@ +// 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}) + 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 { + if i >= len(sig.params) { + break + } + c.checkExpr(arg, sig.params[i]) + } +} diff --git a/internal/typecheck/typecheck_test.go b/internal/typecheck/typecheck_test.go new file mode 100644 index 00000000..5bdd17b8 --- /dev/null +++ b/internal/typecheck/typecheck_test.go @@ -0,0 +1,81 @@ +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") +} diff --git a/numscript.go b/numscript.go index fda5d501..470b0569 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,68 @@ 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 + +// 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. +var ( + VerifyCompiledProgram = vm.Verify + VerifyCompiledProgramWithVars = vm.VerifyWithVars +) + +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 now (funds.Posting); the VM leaves scope fields + // empty. TODO map VM tx/account metadata (stringified) onto the typed + // contract; deferred together with scopes in the VM. + return ExecutionResult{Postings: res.Postings}, nil +}