package vm import ( "encoding/binary" "fmt" "strconv" "strings" ) var mnemonics = map[string]Op{ "nop": NOP, "yield": YIELD, "halt": HALT, "ldi": LDI, "lui": LUI, "mov": MOV, "add": ADD, "sub": SUB, "mul": MUL, "div": DIV, "mod": MOD, "and": AND, "or": OR, "xor": XOR, "shl": SHL, "shr": SHR, "sar": SAR, "addi": ADDI, "jmp": JMP, "beq": BEQ, "bne": BNE, "blt": BLT, "bge": BGE, "call": CALL, "ret": RET, "push": PUSH, "pop": POP, "ldb": LDB, "ldh": LDH, "ldw": LDW, "stb": STB, "sth": STH, "stw": STW, "in": IN, "out": OUT, } // Assemble converts assembly text into a program. Syntax: // // label: define a label // .equ NAME value define a constant // li rd, value pseudo-op: load a 32-bit constant (always 2 instructions) // ldw rd, [rb+off] memory access (also ldb/ldh/stb/sth/stw) // in rd, port | out port, rs // beq ra, rb, label branches and jmp/call take labels // // Registers are r0..r15 (sp = r15). Comments start with ';' or '#'. func Assemble(src string) ([]byte, error) { type line struct { no int text string } var lines []line consts := map[string]int64{} labels := map[string]int{} n := 0 // instruction count for i, raw := range strings.Split(src, "\n") { t := raw if j := strings.IndexAny(t, ";#"); j >= 0 { t = t[:j] } t = strings.TrimSpace(t) for { j := strings.Index(t, ":") if j < 0 || strings.ContainsAny(t[:j], " \t,[") { break } labels[t[:j]] = n t = strings.TrimSpace(t[j+1:]) } if t == "" { continue } f := strings.Fields(t) f[0] = strings.ToLower(f[0]) if f[0] == ".equ" { if len(f) != 3 { return nil, fmt.Errorf("line %d: .equ NAME value", i+1) } v, err := parseNum(f[2], consts) if err != nil { return nil, fmt.Errorf("line %d: %v", i+1, err) } consts[f[1]] = v continue } if f[0] == "li" { n += 2 } else { n++ } lines = append(lines, line{i + 1, t}) } var out []byte emit := func(op Op, ra, rb int, imm int32) { out = binary.LittleEndian.AppendUint32(out, Encode(op, ra, rb, imm)) } for _, l := range lines { name, rest, _ := strings.Cut(l.text, " ") name = strings.ToLower(name) args := splitArgs(rest) err := func() error { pc := len(out) / 4 if name == "li" { if len(args) != 2 { return fmt.Errorf("li rd, value") } rd, err := parseReg(args[0]) if err != nil { return err } v, err := parseNum(args[1], consts) if err != nil { return err } emit(LDI, rd, 0, int32(int16(uint32(v)))) emit(LUI, rd, 0, int32(int16(uint32(v)>>16))) return nil } op, ok := mnemonics[name] if !ok { return fmt.Errorf("unknown mnemonic %q", name) } target := func(s string) (int32, error) { if a, ok := labels[s]; ok { return int32(a - (pc + 1)), nil } v, err := parseNum(s, consts) return int32(v), err } need := func(k int) error { if len(args) != k { return fmt.Errorf("%s takes %d operands", name, k) } return nil } switch op { case NOP, YIELD, HALT, RET: if err := need(0); err != nil { return err } emit(op, 0, 0, 0) case LDI, LUI, ADDI: if err := need(2); err != nil { return err } ra, err := parseReg(args[0]) if err != nil { return err } v, err := parseNum(args[1], consts) if err != nil { return err } if v < -32768 || v > 65535 { return fmt.Errorf("immediate %d out of 16-bit range (use li)", v) } emit(op, ra, 0, int32(int16(v))) case MOV, ADD, SUB, MUL, DIV, MOD, AND, OR, XOR, SHL, SHR, SAR: if err := need(2); err != nil { return err } ra, err := parseReg(args[0]) if err != nil { return err } rb, err := parseReg(args[1]) if err != nil { return err } emit(op, ra, rb, 0) case JMP, CALL: if err := need(1); err != nil { return err } t, err := target(args[0]) if err != nil { return err } emit(op, 0, 0, t) case BEQ, BNE, BLT, BGE: if err := need(3); err != nil { return err } ra, err := parseReg(args[0]) if err != nil { return err } rb, err := parseReg(args[1]) if err != nil { return err } t, err := target(args[2]) if err != nil { return err } emit(op, ra, rb, t) case PUSH, POP: if err := need(1); err != nil { return err } ra, err := parseReg(args[0]) if err != nil { return err } emit(op, ra, 0, 0) case LDB, LDH, LDW, STB, STH, STW: if err := need(2); err != nil { return err } ra, err := parseReg(args[0]) if err != nil { return err } rb, off, err := parseMem(args[1], consts) if err != nil { return err } emit(op, ra, rb, off) case IN: if err := need(2); err != nil { return err } ra, err := parseReg(args[0]) if err != nil { return err } p, err := parseNum(args[1], consts) if err != nil { return err } emit(op, ra, 0, int32(int16(p))) case OUT: if err := need(2); err != nil { return err } p, err := parseNum(args[0], consts) if err != nil { return err } ra, err := parseReg(args[1]) if err != nil { return err } emit(op, ra, 0, int32(int16(p))) } return nil }() if err != nil { return nil, fmt.Errorf("line %d: %v", l.no, err) } } return out, nil } func splitArgs(s string) []string { s = strings.TrimSpace(s) if s == "" { return nil } parts := strings.Split(s, ",") for i := range parts { parts[i] = strings.TrimSpace(parts[i]) } return parts } func parseReg(s string) (int, error) { s = strings.ToLower(s) if s == "sp" { return SP, nil } if strings.HasPrefix(s, "r") { if n, err := strconv.Atoi(s[1:]); err == nil && n >= 0 && n < NumRegs { return n, nil } } return 0, fmt.Errorf("bad register %q", s) } func parseNum(s string, consts map[string]int64) (int64, error) { if v, ok := consts[s]; ok { return v, nil } v, err := strconv.ParseInt(s, 0, 64) if err != nil { return 0, fmt.Errorf("bad number or unknown name %q", s) } return v, nil } // parseMem parses "[rb]" or "[rb+off]" / "[rb-off]". func parseMem(s string, consts map[string]int64) (int, int32, error) { if !strings.HasPrefix(s, "[") || !strings.HasSuffix(s, "]") { return 0, 0, fmt.Errorf("bad memory operand %q", s) } s = strings.TrimSpace(s[1 : len(s)-1]) regPart, offPart := s, "" if i := strings.IndexAny(s, "+-"); i >= 0 { regPart, offPart = strings.TrimSpace(s[:i]), strings.TrimSpace(s[i:]) offPart = strings.ReplaceAll(offPart, " ", "") } rb, err := parseReg(regPart) if err != nil { return 0, 0, err } var off int64 if offPart != "" { sign := int64(1) if offPart[0] == '-' { sign = -1 } v, err := parseNum(offPart[1:], consts) if err != nil { return 0, 0, err } off = sign * v } if off < -32768 || off > 32767 { return 0, 0, fmt.Errorf("offset %d out of range", off) } return rb, int32(off), nil }