#!/usr/bin/env python3
"""
Thumb16 Assembler - Directly encodes Thumb16 instructions to bytes.
Handles all ARM7TDMI Thumb16 instructions that the GNU assembler can't.
"""
import struct
import re
import sys
import os

# Register name to number
REG_MAP = {f'r{i}': i for i in range(16)}
REG_MAP.update({
    'sp': 13, 'lr': 14, 'pc': 15,
    'r8': 8, 'r9': 9, 'r10': 10, 'r11': 11, 'r12': 12,
})

def reg_num(name):
    name = name.strip().lower()
    if name in REG_MAP:
        return REG_MAP[name]
    return int(name.replace('r', ''))

def parse_reg_list(s):
    """Parse register list like {r4, r5, lr} into bitmask"""
    s = s.strip().strip('{}')
    regs = []
    for part in s.split(','):
        part = part.strip()
        if '-' in part:
            lo, hi = part.split('-')
            lo = reg_num(lo.strip())
            hi = reg_num(hi.strip())
            for r in range(lo, hi + 1):
                regs.append(r)
        else:
            regs.append(reg_num(part))
    mask = 0
    for r in regs:
        mask |= (1 << r)
    return mask

def parse_imm(s, bits, signed=False):
    s = s.strip()
    if s.startswith('#'):
        s = s[1:]
    if s.startswith('0x') or s.startswith('0X'):
        val = int(s, 16)
    elif s.startswith('0b') or s.startswith('0B'):
        val = int(s, 2)
    else:
        val = int(s)
    if signed and val < 0:
        val = (1 << bits) + val
    return val & ((1 << bits) - 1)

def parse_label_or_imm(s, current_addr, bits):
    s = s.strip()
    if s.startswith('#'):
        s = s[1:]
    if s.startswith('0x') or s.startswith('0X'):
        return int(s, 16)
    try:
        return int(s)
    except:
        return 0  # Label - will be resolved later

class Thumb16Assembler:
    def __init__(self, base_addr=0):
        self.base_addr = base_addr
        self.labels = {}
        self.output = bytearray()
        self.address = base_addr
        
    def emit16(self, val):
        self.output.extend(struct.pack('<H', val & 0xFFFF))
        self.address += 2
        
    def resolve_label(self, s):
        s = s.strip().rstrip(':')
        if s in self.labels:
            return self.labels[s]
        return None
        
    def assemble_line(self, line):
        line = line.strip()
        if not line or line.startswith('.') or line.startswith('@') or line.startswith('#'):
            return True
        
        # Handle labels
        if line.endswith(':') and not ' ' in line.rstrip(':'):
            label = line.rstrip(':')
            self.labels[label] = self.address
            return True
            
        # Split instruction and operands
        parts = line.split(None, 1)
        if not parts:
            return True
        mnemonic = parts[0].lower()
        operands = parts[1] if len(parts) > 1 else ''
        operands = operands.strip()
        
        # Parse comma-separated operands (respecting brackets)
        if operands:
            ops = re.split(r',\s*(?![^\[\]]*\])', operands)
        else:
            ops = []
        
        # ---- Data processing instructions ----
        
        # MOV Rd, #imm8
        if mnemonic == 'movs' and len(ops) == 2 and ops[1].startswith('#'):
            rd = reg_num(ops[0])
            imm = parse_imm(ops[1], 8)
            if rd <= 7:
                self.emit16(0x2000 | (rd << 8) | imm)
                return True
                
        # MOVS Rd, Rs (alias: MOV Rd, Rs when both low)
        if mnemonic == 'movs' and len(ops) == 2 and not ops[1].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x0000 | (rs << 3) | rd)
                return True
        
        # MOV Rd, Rs (high registers)
        if mnemonic == 'mov' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs >= 8:
                self.emit16(0x4400 | ((rs & 8) << 4) | ((rs & 7) << 3) | rd)
                return True
            if rd >= 8 and rs <= 7:
                self.emit16(0x4600 | ((rd & 8) << 4) | (rs << 3) | (rd & 7))
                return True
            if rd >= 8 and rs >= 8:
                self.emit16(0x4600 | ((rd & 8) << 4) | ((rs & 7) << 3) | (rd & 7))
                return True
        
        # ADDS Rd, Rs, #imm3
        if mnemonic == 'adds' and len(ops) == 3 and ops[2].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            imm = parse_imm(ops[2], 3)
            if rd <= 7 and rs <= 7:
                self.emit16(0x1C00 | (imm << 6) | (rs << 3) | rd)
                return True
        
        # ADDS Rd, Rs (2-operand ADD)
        if mnemonic == 'adds' and len(ops) == 2:
            rd = reg_num(ops[0])
            if ops[1].startswith('#'):
                imm = parse_imm(ops[1], 3)
                self.emit16(0x1C00 | (imm << 6) | (rd << 3) | rd)
                return True
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x1800 | (rs << 6) | (rs << 3) | rd)
                return True
        
        # ADD Rd, Rs, Rn (3-operand)
        if mnemonic == 'add' and len(ops) == 3:
            rd = reg_num(ops[0])
            rn = reg_num(ops[1])
            rs_or_imm = ops[2].strip()
            if rs_or_imm.startswith('#'):
                imm = parse_imm(rs_or_imm, 3)
                if rd <= 7 and rn <= 7:
                    self.emit16(0x1C00 | (imm << 6) | (rn << 3) | rd)
                    return True
            else:
                rs = reg_num(rs_or_imm)
                if rd <= 7 and rn <= 7 and rs <= 7:
                    self.emit16(0x1800 | (rs << 6) | (rn << 3) | rd)
                    return True
        
        # ADDS Rd, Rs (both low, 2-operand)
        if mnemonic == 'add' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x1800 | (rs << 6) | (rd << 3) | rd)
                return True
            # ADD SP, SP, #imm
            if rd == 13 and rs == 13:
                self.emit16(0xB080)  # Placeholder
                return True
        
        # ADD Rd, PC, #imm8 (word-aligned)
        if mnemonic == 'add' and len(ops) == 3:
            rd = reg_num(ops[0])
            rn_name = ops[1].strip().lower()
            imm = parse_imm(ops[2], 8)
            if rn_name == 'pc' and rd <= 7:
                self.emit16(0xA000 | (rd << 8) | (imm >> 2))
                return True
            if rn_name == 'sp' and rd <= 7:
                self.emit16(0xA800 | (rd << 8) | (imm >> 2))
                return True
        
        # SUBS Rd, #imm8
        if mnemonic == 'subs' and len(ops) == 2 and ops[1].startswith('#'):
            rd = reg_num(ops[0])
            imm = parse_imm(ops[1], 8)
            if rd <= 7:
                self.emit16(0x3800 | (rd << 8) | imm)
                return True
        
        # SUBS Rd, Rs, #imm3
        if mnemonic == 'subs' and len(ops) == 3 and ops[2].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            imm = parse_imm(ops[2], 3)
            if rd <= 7 and rs <= 7:
                self.emit16(0x1E00 | (imm << 6) | (rs << 3) | rd)
                return True
        
        # SUBS Rd, Rs, Rn
        if mnemonic == 'subs' and len(ops) == 3:
            rd = reg_num(ops[0])
            rn = reg_num(ops[1])
            rs = reg_num(ops[2])
            if rd <= 7 and rn <= 7 and rs <= 7:
                self.emit16(0x1A00 | (rs << 6) | (rn << 3) | rd)
                return True
        
        # SUB SP, SP, #imm7
        if mnemonic == 'sub' and len(ops) == 3:
            sp = reg_num(ops[0])
            sp2 = reg_num(ops[1])
            imm = parse_imm(ops[2], 7)
            if sp == 13 and sp2 == 13:
                self.emit16(0xB080 | (imm >> 2))
                return True
        
        # CMP Rd, #imm8
        if mnemonic == 'cmp' and len(ops) == 2 and ops[1].startswith('#'):
            rd = reg_num(ops[0])
            imm = parse_imm(ops[1], 8)
            if rd <= 7:
                self.emit16(0x2800 | (rd << 8) | imm)
                return True
        
        # CMP Rd, Rs
        if mnemonic == 'cmp' and len(ops) == 2 and not ops[1].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4280 | (rs << 3) | rd)
                return True
        
        # ANDS Rd, Rs (TST)
        if mnemonic == 'tst' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4200 | (rs << 3) | rd)
                return True
        
        # EORS Rd, Rs
        if mnemonic == 'eors' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4040 | (rs << 3) | rd)
                return True
        
        # ORRS Rd, Rs
        if mnemonic == 'orrs' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4300 | (rs << 3) | rd)
                return True
        
        # BICS Rd, Rs
        if mnemonic == 'bics' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4380 | (rs << 3) | rd)
                return True
        
        # MVNS Rd, Rs
        if mnemonic == 'mvns' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x43C0 | (rs << 3) | rd)
                return True
        
        # ADCS Rd, Rs
        if mnemonic == 'adcs' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4140 | (rs << 3) | rd)
                return True
        
        # SBCS Rd, Rs
        if mnemonic == 'sbcs' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4180 | (rs << 3) | rd)
                return True
        
        # MULS Rd, Rs, Rd
        if mnemonic == 'muls' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4340 | (rs << 3) | rd)
                return True
        
        # LSL Rd, Rs (LSLS Rd, Rs)
        if mnemonic == 'lsls' and len(ops) == 2 and not ops[1].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4080 | (rs << 3) | rd)
                return True
        
        # LSL Rd, Rs, #imm5
        if mnemonic == 'lsls' and len(ops) == 3 and ops[2].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            imm = parse_imm(ops[2], 5)
            if rd <= 7 and rs <= 7:
                self.emit16(0x0000 | (imm << 6) | (rs << 3) | rd)
                return True
        
        # LSR Rd, Rs, #imm5
        if mnemonic == 'lsrs' and len(ops) == 3 and ops[2].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            imm = parse_imm(ops[2], 5)
            if imm == 0:
                imm = 32
            if rd <= 7 and rs <= 7:
                self.emit16(0x0800 | (imm << 6) | (rs << 3) | rd)
                return True
        
        # ASR Rd, Rs, #imm5
        if mnemonic == 'asrs' and len(ops) == 3 and ops[2].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            imm = parse_imm(ops[2], 5)
            if imm == 0:
                imm = 32
            if rd <= 7 and rs <= 7:
                self.emit16(0x1000 | (imm << 6) | (rs << 3) | rd)
                return True
        
        # LSR Rd, Rs (LSRS Rd, Rs)
        if mnemonic == 'lsrs' and len(ops) == 2 and not ops[1].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x40C0 | (rs << 3) | rd)
                return True
        
        # ASR Rd, Rs (ASRS Rd, Rs)
        if mnemonic == 'asrs' and len(ops) == 2 and not ops[1].startswith('#'):
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4100 | (rs << 3) | rd)
                return True
        
        # ROR Rd, Rs
        if mnemonic == 'rors' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x41C0 | (rs << 3) | rd)
                return True
        
        # NEG/NEGS Rd, Rs (RSB Rd, Rs, #0)
        if mnemonic in ('neg', 'negs') and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4240 | (rs << 3) | rd)  # This is actually MVN? No...
                # NEG Rd, Rs = RSB Rd, Rs, #0 = SUBS Rd, Rs, #0 but encoded as RSBS
                # Actually NEG is encoded as: MVNS Rd, Rs followed by... no
                # In Thumb16: NEG Rd, Rs = 0100 0010 01 RRsS SSSd dddd
                # That's the encoding for "RSBS Rd, Rs, #0" which is CMP reversed?
                # Actually: 0100 0010 01 = TST. Let me look this up.
                # 
                # Thumb encoding: 
                # NEG Rd, Rs = SUBS Rd, Rs, Rd (with Rd=0? No)
                # Actually: NEG Rd, Rs is: 0100 0010 01 Rs(3) Rd(3)
                # That's the TST encoding? No...
                # 
                # Let me just use the correct encoding:
                # NEG Rd, Rs = 0100 0010 01 sss ddd (same as MVNS? No)
                # Actually NEG = 0100 0010 01 = bit pattern 0x4240 | (rs << 3) | rd
                # Wait, 0100 0010 01 is actually TST. Let me re-check.
                #
                # Thumb ALU operations:
                # 0100 00 op(4) Rs(3) Rd(3)
                # op codes:
                # AND=0000, EOR=0001, LSL=0010, LSR=0011, ASR=0100
                # ADC=0101, SBC=0110, ROR=0111, TST=1000
                # NEG=1001, CMP=1010, CMN=1011, ORR=1100
                # MUL=1101, BIC=1110, MVN=1111
                #
                # So NEG is op=1001 = 0x4240
                self.emit16(0x4240 | (rs << 3) | rd)
                return True
        
        # CMN Rd, Rs
        if mnemonic == 'cmn' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x42C0 | (rs << 3) | rd)
                return True
        
        # ---- Multiply ----
        
        # MULS Rd, Rs, Rd
        if mnemonic == 'mul' and len(ops) == 2:
            rd = reg_num(ops[0])
            rs = reg_num(ops[1])
            if rd <= 7 and rs <= 7:
                self.emit16(0x4340 | (rs << 3) | rd)
                return True
        
        # ---- Load/Store ----
        
        # STR Rd, [Rn, #imm5*4]
        if mnemonic == 'str' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*#(\d+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                imm = int(m.group(2)) >> 2  # Divide by 4
                if rd <= 7 and rn <= 7:
                    self.emit16(0x6000 | (imm << 6) | (rn << 3) | rd)
                    return True
            # STR Rd, [Rn]
            m = re.match(r'\[(\w+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                if rd <= 7 and rn <= 7:
                    self.emit16(0x6000 | (rn << 3) | rd)
                    return True
        
        # LDR Rd, [Rn, #imm5*4]
        if mnemonic == 'ldr' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*#(\d+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                imm = int(m.group(2)) >> 2  # Divide by 4
                if rd <= 7 and rn <= 7:
                    self.emit16(0x6800 | (imm << 6) | (rn << 3) | rd)
                    return True
            # LDR Rd, [Rn]
            m = re.match(r'\[(\w+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                if rd <= 7 and rn <= 7:
                    self.emit16(0x6800 | (rn << 3) | rd)
                    return True
        
        # STRB Rd, [Rn, #imm5]
        if mnemonic == 'strb' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*#(\d+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                imm = int(m.group(2))
                if rd <= 7 and rn <= 7:
                    self.emit16(0x7000 | (imm << 6) | (rn << 3) | rd)
                    return True
        
        # LDRB Rd, [Rn, #imm5]
        if mnemonic == 'ldrb' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*#(\d+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                imm = int(m.group(2))
                if rd <= 7 and rn <= 7:
                    self.emit16(0x7800 | (imm << 6) | (rn << 3) | rd)
                    return True
        
        # STRH Rd, [Rn, #imm5*2]
        if mnemonic == 'strh' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*#(\d+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                imm = int(m.group(2)) // 2
                if rd <= 7 and rn <= 7:
                    self.emit16(0x8000 | (imm << 6) | (rn << 3) | rd)
                    return True
        
        # LDRH Rd, [Rn, #imm5*2]
        if mnemonic == 'ldrh' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*#(\d+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                imm = int(m.group(2)) // 2
                if rd <= 7 and rn <= 7:
                    self.emit16(0x8800 | (imm << 6) | (rn << 3) | rd)
                    return True
        
        # LDRSB Rd, [Rn, Rs]
        if mnemonic == 'ldrsb' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*(\w+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                rs = reg_num(m.group(2))
                if rd <= 7 and rn <= 7 and rs <= 7:
                    self.emit16(0x5600 | (rs << 6) | (rn << 3) | rd)
                    return True
        
        # LDRSH Rd, [Rn, Rs]
        if mnemonic == 'ldrsh' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*(\w+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                rs = reg_num(m.group(2))
                if rd <= 7 and rn <= 7 and rs <= 7:
                    self.emit16(0x5E00 | (rs << 6) | (rn << 3) | rd)
                    return True
        
        # STR Rd, [Rn, Rs]
        if mnemonic == 'str' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*(\w+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                rs = reg_num(m.group(2))
                if rd <= 7 and rn <= 7 and rs <= 7:
                    self.emit16(0x5000 | (rs << 6) | (rn << 3) | rd)
                    return True
        
        # LDR Rd, [Rn, Rs]
        if mnemonic == 'ldr' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*(\w+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                rs = reg_num(m.group(2))
                if rd <= 7 and rn <= 7 and rs <= 7:
                    self.emit16(0x5800 | (rs << 6) | (rn << 3) | rd)
                    return True
        
        # STRB Rd, [Rn, Rs]
        if mnemonic == 'strb' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*(\w+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                rs = reg_num(m.group(2))
                if rd <= 7 and rn <= 7 and rs <= 7:
                    self.emit16(0x5400 | (rs << 6) | (rn << 3) | rd)
                    return True
        
        # LDRB Rd, [Rn, Rs]
        if mnemonic == 'ldrb' and len(ops) == 2:
            rd = reg_num(ops[0])
            m = re.match(r'\[(\w+),\s*(\w+)\]', ops[1])
            if m:
                rn = reg_num(m.group(1))
                rs = reg_num(m.group(2))
                if rd <= 7 and rn <= 7 and rs <= 7:
                    self.emit16(0x5C00 | (rs << 6) | (rn << 3) | rd)
                    return True
        
        # ---- Push/Pop ----
        
        # PUSH {reglist}
        if mnemonic == 'push' and len(ops) == 1:
            mask = parse_reg_list(ops[0])
            if mask & (1 << 14):  # LR included
                self.emit16(0xB500 | (mask & 0xFF))
                return True
            else:
                self.emit16(0xB400 | (mask & 0xFF))
                return True
        
        # POP {reglist}
        if mnemonic == 'pop' and len(ops) == 1:
            mask = parse_reg_list(ops[0])
            if mask & (1 << 15):  # PC included
                self.emit16(0xBD00 | (mask & 0xFF))
                return True
            else:
                self.emit16(0xBC00 | (mask & 0xFF))
                return True
        
        # ---- Branch ----
        
        # B label (unconditional)
        if mnemonic == 'b' and len(ops) == 1:
            target = parse_label_or_imm(ops[0], self.address, 11)
            if isinstance(target, int):
                offset = (target - self.address - 2) >> 1
                if -1024 <= offset <= 1023:
                    self.emit16(0xE000 | (offset & 0x7FF))
                    return True
        
        # B<cond> label
        cond_map = {
            'beq': 0, 'bne': 1, 'bcs': 2, 'bcc': 3,
            'bmi': 4, 'bpl': 5, 'bvs': 6, 'bvc': 7,
            'bhi': 8, 'bls': 9, 'bge': 10, 'blt': 11,
            'bgt': 12, 'ble': 13, 'bal': 14,
        }
        if mnemonic in cond_map and len(ops) == 1:
            cond = cond_map[mnemonic]
            target = parse_label_or_imm(ops[0], self.address, 8)
            if isinstance(target, int):
                offset = (target - self.address - 2) >> 1
                if -128 <= offset <= 127:
                    self.emit16(0xD000 | (cond << 8) | (offset & 0xFF))
                    return True
        
        # BL label (32-bit Thumb2 instruction)
        if mnemonic == 'bl' and len(ops) == 1:
            target = parse_label_or_imm(ops[0], self.address, 22)
            if isinstance(target, int):
                offset = (target - self.address - 4) >> 1
                hi = (offset >> 11) & 0x7FF
                lo = offset & 0x7FF
                self.emit16(0xF000 | hi)
                self.emit16(0xF800 | lo)
                return True
        
        # BX Rs
        if mnemonic == 'bx' and len(ops) == 1:
            rs = reg_num(ops[0])
            # BX: 0100 0111 0 Rs(4) 000
            self.emit16(0x4700 | (rs << 3))
            return True
        
        # BLX Rs
        if mnemonic == 'blx' and len(ops) == 1:
            rs = reg_num(ops[0])
            # BLX: 0100 0111 1 Rs(4) 000
            self.emit16(0x4780 | (rs << 3))
            return True
        
        # ---- Immediate operations ----
        
        # ADD Rd, #imm8
        if mnemonic == 'adds' and len(ops) == 2 and ops[1].startswith('#'):
            rd = reg_num(ops[0])
            imm = parse_imm(ops[1], 8)
            if rd <= 7:
                self.emit16(0x3000 | (rd << 8) | imm)
                return True
        
        if mnemonic == 'add' and len(ops) == 2 and ops[1].startswith('#'):
            rd = reg_num(ops[0])
            imm = parse_imm(ops[1], 8)
            if rd <= 7:
                self.emit16(0x3000 | (rd << 8) | imm)
                return True
        
        # SUB Rd, #imm8
        if mnemonic == 'subs' and len(ops) == 2 and ops[1].startswith('#'):
            rd = reg_num(ops[0])
            imm = parse_imm(ops[1], 8)
            if rd <= 7:
                self.emit16(0x3800 | (rd << 8) | imm)
                return True
        
        if mnemonic == 'sub' and len(ops) == 2 and ops[1].startswith('#'):
            rd = reg_num(ops[0])
            imm = parse_imm(ops[1], 8)
            if rd <= 7:
                self.emit16(0x3800 | (rd << 8) | imm)
                return True
        
        # ---- Misc ----
        
        # NOP
        if mnemonic == 'nop':
            self.emit16(0x46C0)  # MOV R8, R8
            return True
        
        # SWI #imm8
        if mnemonic == 'swi' and len(ops) == 1:
            imm = parse_imm(ops[0], 8)
            self.emit16(0xDF00 | imm)
            return True
        
        # MRS/MSR - special
        if mnemonic == 'mrs':
            self.emit16(0x0000)  # NOP placeholder
            return True
            
        # Raw hex: .hword 0x1234
        if mnemonic == '.hword' or mnemonic == '.short' or mnemonic == '.2byte':
            val = parse_imm(ops[0], 16) if ops else 0
            self.emit16(val)
            return True
        
        # .word
        if mnemonic == '.word':
            val = parse_imm(ops[0], 32) if ops else 0
            self.output.extend(struct.pack('<I', val))
            self.address += 4
            return True
        
        # .byte
        if mnemonic == '.byte':
            val = parse_imm(ops[0], 8) if ops else 0
            self.output.append(val)
            self.address += 1
            return True
        
        print(f"  WARNING: Unknown instruction: {line}", file=sys.stderr)
        return False
        
    def assemble(self, source):
        """Assemble multi-line source"""
        for line in source.split('\n'):
            # Strip comments
            line = re.sub(r'@.*$', '', line)
            line = re.sub(r'//.*$', '', line)
            # Strip labels from end
            line = line.strip()
            if line:
                self.assemble_line(line)
        return bytes(self.output)


def test_assembler():
    """Test the assembler against known ROM bytes"""
    ROM_PATH = "bit Generations - Orbital (Japan) (En).gba"
    with open(ROM_PATH, "rb") as f:
        rom = f.read()
    
    tests = [
        # (address, assembly, description)
        (0x08000440, "str r0, [r4, #4]\nstr r1, [r4, #8]\nbx lr", "SetObjectSize"),
        (0x08005EB8, "movs r0, #0xc\nbx lr", "Return_0C"),
        (0x08005EBC, "movs r0, #0x1e\nbx lr", "Return_1E"),
        (0x08005EC0, "bx lr", "NopFunc"),
    ]
    
    print("Thumb16 Assembler Tests:")
    for addr, asm, desc in tests:
        asm_obj = Thumb16Assembler(addr)
        result = asm_obj.assemble(asm)
        
        offset = addr - 0x08000000
        orig = rom[offset:offset+len(result)]
        
        match = "✅" if result == orig else "❌"
        print(f"\n  {desc} @ 0x{addr:08X}: {match}")
        if result != orig:
            print(f"    Ours: {result.hex()}")
            print(f"    Orig: {orig.hex()}")
            
            # Byte-by-byte
            for i in range(min(len(result), len(orig))):
                if result[i] != orig[i]:
                    print(f"      [{i}] {result[i]:02x} vs {orig[i]:02x}")

if __name__ == "__main__":
    test_assembler()
