#!/usr/bin/env python3
import sys
from pathlib import Path


def annotate(text, labels={}, cellsize=56):
    address, prog = 0, False
    for line in text.split('\n'):
        length, words = 0, line.split(';')[0].split()
        match words:
            case ['.code']:
                prog = True
            case ['.code', addr]:
                prog = True
                address = eval(addr)
            case [directive, *_] if directive.startswith('.'):
                prog = False
            case [label, *cells] if prog:
                if label.startswith('_'):
                    labels[label] = address
                    length = len("".join(cells))
                    line = " " * len(label) + line[len(label):]
                else:
                    length = len("".join(words))
        yield address, prog, line
        address += length * 4 // cellsize


def replace(annotext, labels, asize=4):
    for addr, prog, line in ann:
        words = line.split(";")
        if words:
            for label, address in labels.items():
                words[0] = words[0].replace(label, f"{address:0{asize}x}")
        yield ";".join(words)


if __name__ == "__main__":
    labels = {}
    if len(sys.argv) > 1:
        source = Path(sys.argv[1]).read_text()
    else:
        source = sys.stdin.read()
    if '.mm-3' in source:
        cellsize = 56
    elif '.mm-2' in source:
        cellsize = 40
    elif '.mm-1' in source:
        cellsize = 24
    elif '.mm-m' in source or 'mm-r' in source:
        cellsize = 16
    else:
        cellsize = 8

    ann = [*annotate(source, labels, cellsize)]
    res = "\n".join(replace(ann, labels, 4))
    if len(sys.argv) > 2:
        Path(sys.argv[2]).write_text(res)
    else:
        sys.stdout.write(res)
    if len(sys.argv) > 3:
        print(labels, file=sys.stderr)
