main cpu_dsl.py
   1#!/usr/bin/env python3
   2
   3
   4class Block:
   5	def addOp(self, op):
   6		pass
   7	
   8	def processLine(self, parts):
   9		if parts[0] == 'switch':
  10			o = Switch(self, parts[1])
  11			self.addOp(o)
  12			return o
  13		elif parts[0] == 'if':
  14			o = If(self, parts[1])
  15			self.addOp(o)
  16			return o
  17		elif parts[0] == 'end':
  18			raise Exception('end is only allowed inside a switch or if block')
  19		else:
  20			self.addOp(NormalOp(parts))
  21		return self
  22		
  23	def processOps(self, prog, fieldVals, output, otype, oplist):
  24		for i in range(0, len(oplist)):
  25			if i + 1 < len(oplist) and oplist[i+1].op == 'update_flags':
  26				flagUpdates, _ = prog.flags.parseFlagUpdate(oplist[i+1].params[0])
  27			else:
  28				flagUpdates = None
  29			oplist[i].generate(prog, self, fieldVals, output, otype, flagUpdates)
  30		
  31	def resolveLocal(self, name):
  32		return None
  33			
  34class ChildBlock(Block):
  35	def processLine(self, parts):
  36		if parts[0] == 'end':
  37			return self.parent
  38		return super().processLine(parts)
  39
  40#Represents an instruction of the emulated CPU
  41class Instruction(Block):
  42	def __init__(self, value, fields, name):
  43		self.value = value
  44		self.fields = fields
  45		self.name = name
  46		self.implementation = []
  47		self.locals = {}
  48		self.regValues = {}
  49		self.varyingBits = 0
  50		self.invalidFieldValues = {}
  51		self.newLocals = []
  52		for field in fields:
  53			self.varyingBits += fields[field][1]
  54	
  55	def addOp(self, op):
  56		if op.op == 'local':
  57			name = op.params[0]
  58			size = int(op.params[1])
  59			self.locals[name] = size
  60		elif op.op == 'invalid':
  61			name = op.params[0]
  62			value = int(op.params[1])
  63			self.invalidFieldValues.setdefault(name, set()).add(value)
  64		else:
  65			self.implementation.append(op)
  66			
  67	def resolveLocal(self, name):
  68		if name in self.locals:
  69			return name
  70		return None
  71	
  72	def addLocal(self, name, size):
  73		self.locals[name] = size
  74		self.newLocals.append(name)
  75		
  76	def localSize(self, name):
  77		return self.locals.get(name)
  78			
  79	def __lt__(self, other):
  80		if isinstance(other, Instruction):
  81			if self.varyingBits != other.varyingBits:
  82				return self.varyingBits < other.varyingBits
  83			return self.value < other.value
  84		else:
  85			return NotImplemented
  86			
  87	def allValues(self):
  88		values = []
  89		for i in range(0, 1 << self.varyingBits):
  90			iword = self.value
  91			doIt = True
  92			for field in self.fields:
  93				shift,bits = self.fields[field]
  94				val = i & ((1 << bits) - 1)
  95				if field in self.invalidFieldValues and val in self.invalidFieldValues[field]:
  96					doIt = False
  97					break
  98				i >>= bits
  99				iword |= val << shift
 100			if doIt:
 101				values.append(iword)
 102		return values
 103		
 104	def getFieldVals(self, value):
 105		fieldVals = {}
 106		fieldBits = {}
 107		for field in self.fields:
 108			shift,bits = self.fields[field]
 109			val = (value >> shift) & ((1 << bits) - 1)
 110			fieldVals[field] = val
 111			fieldBits[field] = bits
 112		return (fieldVals, fieldBits)
 113	
 114	def generateName(self, value):
 115		fieldVals,fieldBits = self.getFieldVals(value)
 116		names = list(fieldVals.keys())
 117		names.sort()
 118		funName = self.name
 119		for name in names:
 120			funName += '_{0}_{1:0>{2}}'.format(name, bin(fieldVals[name])[2:], fieldBits[name])
 121		return funName
 122		
 123	def generateBody(self, value, prog, otype):
 124		output = []
 125		prog.meta = {}
 126		prog.pushScope(self)
 127		self.regValues = {}
 128		for var in self.locals:
 129			output.append('\n\tuint{sz}_t {name};'.format(sz=self.locals[var], name=var))
 130		self.newLocals = []
 131		fieldVals,_ = self.getFieldVals(value)
 132		self.processOps(prog, fieldVals, output, otype, self.implementation)
 133		
 134		if prog.dispatch == 'call':
 135			begin = '\nvoid ' + self.generateName(value) + '(' + prog.context_type + ' *context)\n{'
 136		elif prog.dispatch == 'goto':
 137			begin = '\n' + self.generateName(value) + ': {'
 138		else:
 139			raise Exception('Unsupported dispatch type ' + prog.dispatch)
 140		if prog.needFlagCoalesce:
 141			begin += prog.flags.coalesceFlags(prog, otype)
 142		if prog.needFlagDisperse:
 143			output.append(prog.flags.disperseFlags(prog, otype))
 144		for var in self.newLocals:
 145			begin += '\n\tuint{sz}_t {name};'.format(sz=self.locals[var], name=var)
 146		prog.popScope()
 147		if prog.dispatch == 'goto':
 148			output += prog.nextInstruction(otype)
 149		return begin + ''.join(output) + '\n}'
 150		
 151	def __str__(self):
 152		pieces = [self.name + ' ' + hex(self.value) + ' ' + str(self.fields)]
 153		for name in self.locals:
 154			pieces.append('\n\tlocal {0} {1}'.format(name, self.locals[name]))
 155		for op in self.implementation:
 156			pieces.append(str(op))
 157		return ''.join(pieces)
 158	
 159#Represents the definition of a helper function
 160class SubRoutine(Block):
 161	def __init__(self, name):
 162		self.name = name
 163		self.implementation = []
 164		self.args = []
 165		self.arg_map = {}
 166		self.locals = {}
 167		self.regValues = {}
 168		self.argValues = {}
 169	
 170	def addOp(self, op):
 171		if op.op == 'arg':
 172			name = op.params[0]
 173			size = int(op.params[1])
 174			self.arg_map[name] = len(self.args)
 175			self.args.append((name, size))
 176		elif op.op == 'local':
 177			name = op.params[0]
 178			size = int(op.params[1])
 179			self.locals[name] = size
 180		else:
 181			self.implementation.append(op)
 182			
 183	def resolveLocal(self, name):
 184		if name in self.locals:
 185			return self.name + '_' + name
 186		return None
 187	
 188	def addLocal(self, name, size):
 189		self.locals[name] = size
 190	
 191	def localSize(self, name):
 192		if name in self.locals:
 193			return self.locals[name]
 194		if name in self.arg_map:
 195			argIndex = self.arg_map[name]
 196			return self.args[argIndex][1]
 197		return None
 198			
 199	def inline(self, prog, params, output, otype, parent):
 200		if len(params) != len(self.args):
 201			raise Exception('{0} expects {1} arguments, but was called with {2}'.format(self.name, len(self.args), len(params)))
 202		argValues = {}
 203		if parent:
 204			self.regValues = parent.regValues
 205		prog.pushScope(self)
 206		i = 0
 207		for name,size in self.args:
 208			argValues[name] = params[i]
 209			i += 1
 210		for name in self.locals:
 211			size = self.locals[name]
 212			output.append('\n\tuint{size}_t {sub}_{local};'.format(size=size, sub=self.name, local=name))
 213		self.argValues = argValues
 214		self.processOps(prog, argValues, output, otype, self.implementation)
 215		prog.popScope()
 216		
 217	def __str__(self):
 218		pieces = [self.name]
 219		for name,size in self.args:
 220			pieces.append('\n\targ {0} {1}'.format(name, size))
 221		for name in self.locals:
 222			pieces.append('\n\tlocal {0} {1}'.format(name, self.locals[name]))
 223		for op in self.implementation:
 224			pieces.append(str(op))
 225		return ''.join(pieces)
 226	
 227class Op:
 228	def __init__(self, evalFun = None):
 229		self.evalFun = evalFun
 230		self.impls = {}
 231		self.outOp = ()
 232	def cBinaryOperator(self, op):
 233		def _impl(prog, params, rawParams, flagUpdates):
 234			if op == '-':
 235				a = params[1]
 236				b = params[0]
 237			else:
 238				a = params[0]
 239				b = params[1]
 240			needsCarry = needsOflow = needsHalf = False
 241			if flagUpdates:
 242				for flag in flagUpdates:
 243					calc = prog.flags.flagCalc[flag]
 244					if calc == 'carry':
 245						needsCarry = True
 246					elif calc == 'half-carry':
 247						needsHalf = True
 248					elif calc == 'overflow':
 249						needsOflow = True
 250			decl = ''
 251			if needsCarry or needsOflow or needsHalf:
 252				size = prog.paramSize(rawParams[2])
 253				if needsCarry and op != 'lsr':
 254					size *= 2
 255				decl,name = prog.getTemp(size)
 256				dst = prog.carryFlowDst = name
 257				prog.lastA = a
 258				prog.lastB = b
 259				prog.lastBFlow = b if op == '-' else '(~{b})'.format(b=b)
 260			else:
 261				dst = params[2]
 262			return decl + '\n\t{dst} = {a} {op} {b};'.format(
 263				dst = dst, a = a, b = b, op = op
 264			)
 265		self.impls['c'] = _impl
 266		self.outOp = (2,)
 267		return self
 268	def cUnaryOperator(self, op):
 269		def _impl(prog, params, rawParams, flagUpdates):
 270			dst = params[1]
 271			decl = ''
 272			if op == '-':
 273				if flagUpdates:
 274					for flag in flagUpdates:
 275						calc = prog.flags.flagCalc[flag]
 276						if calc == 'carry':
 277							needsCarry = True
 278						elif calc == 'half-carry':
 279							needsHalf = True
 280						elif calc == 'overflow':
 281							needsOflow = True
 282				if needsCarry or needsOflow or needsHalf:
 283					size = prog.paramSize(rawParams[1])
 284					if needsCarry:
 285						size *= 2
 286					decl,name = prog.getTemp(size)
 287					dst = prog.carryFlowDst = name
 288					prog.lastA = 0
 289					prog.lastB = params[0]
 290					prog.lastBFlow = params[0]
 291			return decl + '\n\t{dst} = {op}{a};'.format(
 292				dst = dst, a = params[0], op = op
 293			)
 294		self.impls['c'] = _impl
 295		self.outOp = (1,)
 296		return self
 297	def addImplementation(self, lang, outOp, impl):
 298		self.impls[lang] = impl
 299		if not outOp is None:
 300			if type(outOp) is tuple:
 301				self.outOp = outOp
 302			else:
 303				self.outOp = (outOp,)
 304		return self
 305	def evaluate(self, params):
 306		return self.evalFun(*params)
 307	def canEval(self):
 308		return not self.evalFun is None
 309	def numArgs(self):
 310		return self.evalFun.__code__.co_argcount
 311	def numParams(self):
 312		if self.outOp:
 313			params = max(self.outOp) + 1
 314		else:
 315			params = 0
 316		if self.evalFun:
 317			params = max(params, self.numArgs())
 318		return params
 319	def generate(self, otype, prog, params, rawParams, flagUpdates):
 320		if self.impls[otype].__code__.co_argcount == 2:
 321			return self.impls[otype](prog, params)
 322		elif self.impls[otype].__code__.co_argcount == 3:
 323			return self.impls[otype](prog, params, rawParams)
 324		else:
 325			return self.impls[otype](prog, params, rawParams, flagUpdates)
 326		
 327		
 328def _xchgCImpl(prog, params, rawParams):
 329	size = prog.paramSize(rawParams[0])
 330	decl,name = prog.getTemp(size)
 331	return decl + '\n\t{tmp} = {a};\n\t{a} = {b};\n\t{b} = {tmp};'.format(a = params[0], b = params[1], tmp = name)
 332
 333def _dispatchCImpl(prog, params):
 334	if len(params) == 1:
 335		table = 'main'
 336	else:
 337		table = params[1]
 338	if prog.dispatch == 'call':
 339		return '\n\timpl_{tbl}[{op}](context);'.format(tbl = table, op = params[0])
 340	elif prog.dispatch == 'goto':
 341		return '\n\tgoto *impl_{tbl}[{op}];'.format(tbl = table, op = params[0])
 342	else:
 343		raise Exception('Unsupported dispatch type ' + prog.dispatch)
 344
 345def _updateFlagsCImpl(prog, params, rawParams):
 346	autoUpdate, explicit = prog.flags.parseFlagUpdate(params[0])
 347	output = []
 348	parity = None
 349	directFlags = {}
 350	for flag in autoUpdate:
 351		calc = prog.flags.flagCalc[flag]
 352		calc,_,resultBit = calc.partition('-')
 353		if prog.carryFlowDst:
 354			lastDst = prog.carryFlowDst
 355		else:
 356			lastDst = prog.resolveParam(prog.lastDst, prog.currentScope, {})
 357		storage = prog.flags.getStorage(flag)
 358		if calc == 'bit' or calc == 'sign' or calc == 'carry' or calc == 'half' or calc == 'overflow':
 359			myRes = lastDst
 360			if calc == 'sign':
 361				resultBit = prog.paramSize(prog.lastDst) - 1
 362			elif calc == 'carry':
 363				if prog.lastOp.op in ('asr', 'lsr'):
 364					resultBit = 0
 365					myRes = prog.lastA
 366				else:
 367					resultBit = prog.paramSize(prog.lastDst)
 368					if prog.lastOp.op == 'ror':
 369						resultBit -= 1
 370			elif calc == 'half':
 371				resultBit = prog.paramSize(prog.lastDst) - 4
 372				myRes = '({a} ^ {b} ^ {res})'.format(a = prog.lastA, b = prog.lastB, res = lastDst)
 373			elif calc == 'overflow':
 374				resultBit = prog.paramSize(prog.lastDst) - 1
 375				myRes = '((({a} ^ {b})) & ({a} ^ {res}))'.format(a = prog.lastA, b = prog.lastBFlow, res = lastDst)
 376			else:
 377				#Note: offsetting this by the operation size - 8 makes sense for the Z80
 378				#but might not for other CPUs with this kind of fixed bit flag behavior
 379				resultBit = int(resultBit) + prog.paramSize(prog.lastDst) - 8
 380			if type(storage) is tuple:
 381				reg,storageBit = storage
 382				if storageBit == resultBit:
 383					directFlags.setdefault((reg, myRes), []).append(resultBit)
 384				else:
 385					reg = prog.resolveParam(reg, None, {})
 386					if resultBit > storageBit:
 387						op = '>>'
 388						shift = resultBit - storageBit
 389					else:
 390						op = '<<'
 391						shift = storageBit - resultBit
 392					output.append('\n\t{reg} = ({reg} & ~{mask}U) | ({res} {op} {shift}U & {mask}U);'.format(
 393						reg = reg, mask = 1 << storageBit, res = myRes, op = op, shift = shift
 394					))
 395			else:
 396				reg = prog.resolveParam(storage, None, {})
 397				maxBit = prog.paramSize(storage) - 1
 398				if resultBit > maxBit:
 399					output.append('\n\t{reg} = {res} >> {shift} & {mask}U;'.format(reg=reg, res=myRes, shift = resultBit - maxBit, mask = 1 << maxBit))
 400				else:
 401					output.append('\n\t{reg} = {res} & {mask}U;'.format(reg=reg, res=myRes, mask = 1 << resultBit))
 402		elif calc == 'zero':
 403			if prog.carryFlowDst:
 404				realSize = prog.paramSize(prog.lastDst)
 405				if realSize != prog.paramSize(prog.carryFlowDst):
 406					lastDst = '({res} & {mask})'.format(res=lastDst, mask = (1 << realSize) - 1)
 407			if type(storage) is tuple:
 408				reg,storageBit = storage
 409				reg = prog.resolveParam(reg, None, {})
 410				output.append('\n\t{reg} = {res} ? ({reg} & {mask}U) : ({reg} | {bit}U);'.format(
 411					reg = reg, mask = ~(1 << storageBit), res = lastDst, bit = 1 << storageBit
 412				))
 413			else:
 414				reg = prog.resolveParam(storage, None, {})
 415				output.append('\n\t{reg} = {res} == 0;'.format(
 416					reg = reg, res = lastDst
 417				))
 418		elif calc == 'parity':
 419			parity = storage
 420			paritySize = prog.paramSize(prog.lastDst)
 421			if prog.carryFlowDst:
 422				parityDst = paritySrc = prog.carryFlowDst
 423			else:
 424				paritySrc = lastDst
 425				decl,name = prog.getTemp(paritySize)
 426				output.append(decl)
 427				parityDst = name
 428		else:
 429			raise Exception('Unknown flag calc type: ' + calc)
 430	for reg, myRes in directFlags:
 431		bits = directFlags[(reg, myRes)]
 432		resolved = prog.resolveParam(reg, None, {})
 433		if len(bits) == len(prog.flags.storageToFlags[reg]):
 434			output.append('\n\t{reg} = {res};'.format(reg = resolved, res = myRes))
 435		else:
 436			mask = 0
 437			for bit in bits:
 438				mask |= 1 << bit
 439			output.append('\n\t{reg} = ({reg} & ~{mask}U) | ({res} & {mask}U);'.format(
 440				reg = resolved, mask = mask, res = myRes
 441			))
 442	if prog.carryFlowDst:
 443		if prog.lastOp.op != 'cmp':
 444			output.append('\n\t{dst} = {tmpdst};'.format(dst = prog.resolveParam(prog.lastDst, prog.currentScope, {}), tmpdst = prog.carryFlowDst))
 445		prog.carryFlowDst = None
 446	if parity:
 447		if paritySize > 8:
 448			if paritySize > 16:
 449				output.append('\n\t{dst} = {src} ^ ({src} >> 16);'.format(dst=parityDst, src=paritySrc))
 450				paritySrc = parityDst
 451			output.append('\n\t{dst} = {src} ^ ({src} >> 8);'.format(dst=parityDst, src=paritySrc))
 452			paritySrc = parityDst
 453		output.append('\n\t{dst} = ({src} ^ ({src} >> 4)) & 0xF;'.format(dst=parityDst, src=paritySrc))
 454		if type(parity) is tuple:
 455			reg,bit = parity
 456			reg = prog.resolveParam(reg, None, {})
 457			output.append('\n\t{flag} = ({flag} & ~{mask}U) | ((0x6996 >> {parity}) << {bit} & {mask}U);'.format(
 458				flag=reg, mask = 1 << bit, bit = bit, parity = parityDst
 459			))
 460		else:
 461			reg = prog.resolveParam(parity, None, {})
 462			output.append('\n\t{flag} = 0x9669 >> {parity} & 1;'.format(flag=reg, parity=parityDst))
 463			
 464	#TODO: combine explicit flags targeting the same storage location
 465	for flag in explicit:
 466		location = prog.flags.getStorage(flag)
 467		if type(location) is tuple:
 468			reg,bit = location
 469			reg = prog.resolveReg(reg, None, {})
 470			value = str(1 << bit)
 471			if explicit[flag]:
 472				operator = '|='
 473			else:
 474				operator = '&='
 475				value = '~' + value
 476			output.append('\n\t{reg} {op} {val};'.format(reg=reg, op=operator, val=value))
 477		else:
 478			reg = prog.resolveReg(location, None, {})
 479			output.append('\n\t{reg} = {val};'.format(reg=reg, val=explicit[flag]))
 480	return ''.join(output)
 481	
 482def _cmpCImpl(prog, params, rawParams, flagUpdates):
 483	size = prog.paramSize(rawParams[1])
 484	needsCarry = False
 485	if flagUpdates:
 486		for flag in flagUpdates:
 487			calc = prog.flags.flagCalc[flag]
 488			if calc == 'carry':
 489				needsCarry = True
 490				break
 491	if needsCarry:
 492		size *= 2
 493	tmpvar = 'cmp_tmp{sz}__'.format(sz=size)
 494	if flagUpdates:
 495		prog.carryFlowDst = tmpvar
 496		prog.lastA = params[1]
 497		prog.lastB = params[0]
 498		prog.lastBFlow = params[0]
 499	scope = prog.getRootScope()
 500	if not scope.resolveLocal(tmpvar):
 501		scope.addLocal(tmpvar, size)
 502	prog.lastDst = rawParams[1]
 503	return '\n\t{var} = {b} - {a};'.format(var = tmpvar, a = params[0], b = params[1])
 504
 505def _asrCImpl(prog, params, rawParams, flagUpdates):
 506	needsCarry = False
 507	if flagUpdates:
 508		for flag in flagUpdates:
 509			calc = prog.flags.flagCalc[flag]
 510			if calc == 'carry':
 511				needsCarry = True
 512	decl = ''
 513	size = prog.paramSize(rawParams[2])
 514	if needsCarry:
 515		decl,name = prog.getTemp(size * 2)
 516		dst = prog.carryFlowDst = name
 517		prog.lastA = params[0]
 518	else:
 519		dst = params[2]
 520	mask = 1 << (size - 1)
 521	return decl + '\n\t{dst} = ({a} >> {b}) | ({a} & {mask} ? 0xFFFFFFFFU << ({size} - {b}) : 0);'.format(
 522		a = params[0], b = params[1], dst = dst, mask = mask, size=size)
 523	
 524def _sext(size, src):
 525	if size == 16:
 526		return src | 0xFF00 if src & 0x80 else src
 527	else:
 528		return src | 0xFFFF0000 if src & 0x8000 else src
 529
 530def _sextCImpl(prog, params, rawParms):
 531	if params[0] == 16:
 532		fmt = '\n\t{dst} = {src} & 0x80 ? {src} | 0xFF00 : {src};'
 533	else:
 534		fmt = '\n\t{dst} = {src} & 0x8000 ? {src} | 0xFFFF0000 : {src};'
 535	return fmt.format(src=params[1], dst=params[2])
 536	
 537def _getCarryCheck(prog):
 538	carryFlag = None
 539	for flag in prog.flags.flagCalc:
 540		if prog.flags.flagCalc[flag] == 'carry':
 541			carryFlag = flag
 542	if carryFlag is None:
 543		raise Exception('adc requires a defined carry flag')
 544	carryStorage = prog.flags.getStorage(carryFlag)
 545	if type(carryStorage) is tuple:
 546		reg,bit = carryStorage
 547		reg = prog.resolveReg(reg, None, (), False)
 548		return '({reg} & 1 << {bit})'.format(reg=reg, bit=bit)
 549	else:
 550		return prog.resolveReg(carryStorage, None, (), False)
 551
 552def _adcCImpl(prog, params, rawParams, flagUpdates):
 553	needsCarry = needsOflow = needsHalf = False
 554	if flagUpdates:
 555		for flag in flagUpdates:
 556			calc = prog.flags.flagCalc[flag]
 557			if calc == 'carry':
 558				needsCarry = True
 559			elif calc == 'half-carry':
 560				needsHalf = True
 561			elif calc == 'overflow':
 562				needsOflow = True
 563	decl = ''
 564	carryCheck = _getCarryCheck(prog)
 565	if needsCarry or needsOflow or needsHalf:
 566		size = prog.paramSize(rawParams[2])
 567		if needsCarry:
 568			size *= 2
 569		decl,name = prog.getTemp(size)
 570		dst = prog.carryFlowDst = name
 571		prog.lastA = params[0]
 572		prog.lastB = params[1]
 573		prog.lastBFlow = '(~{b})'.format(b=params[1])
 574	else:
 575		dst = params[2]
 576	return decl + '\n\t{dst} = {a} + {b} + ({check} ? 1 : 0);'.format(dst = dst,
 577		a = params[0], b = params[1], check = carryCheck
 578	)
 579
 580def _sbcCImpl(prog, params, rawParams, flagUpdates):
 581	needsCarry = needsOflow = needsHalf = False
 582	if flagUpdates:
 583		for flag in flagUpdates:
 584			calc = prog.flags.flagCalc[flag]
 585			if calc == 'carry':
 586				needsCarry = True
 587			elif calc == 'half-carry':
 588				needsHalf = True
 589			elif calc == 'overflow':
 590				needsOflow = True
 591	decl = ''
 592	carryCheck = _getCarryCheck(prog)
 593	if needsCarry or needsOflow or needsHalf:
 594		size = prog.paramSize(rawParams[2])
 595		if needsCarry:
 596			size *= 2
 597		decl,name = prog.getTemp(size)
 598		dst = prog.carryFlowDst = name
 599		prog.lastA = params[1]
 600		prog.lastB = params[0]
 601		prog.lastBFlow = params[0]
 602	else:
 603		dst = params[2]
 604	return decl + '\n\t{dst} = {b} - {a} - ({check} ? 1 : 0);'.format(dst = dst,
 605		a = params[0], b = params[1], check=_getCarryCheck(prog)
 606	)
 607	
 608def _rolCImpl(prog, params, rawParams, flagUpdates):
 609	needsCarry = False
 610	if flagUpdates:
 611		for flag in flagUpdates:
 612			calc = prog.flags.flagCalc[flag]
 613			if calc == 'carry':
 614				needsCarry = True
 615	decl = ''
 616	size = prog.paramSize(rawParams[2])
 617	if needsCarry:
 618		decl,name = prog.getTemp(size * 2)
 619		dst = prog.carryFlowDst = name
 620	else:
 621		dst = params[2]
 622	return decl + '\n\t{dst} = {a} << {b} | {a} >> ({size} - {b});'.format(dst = dst,
 623		a = params[0], b = params[1], size=size
 624	)
 625	
 626def _rlcCImpl(prog, params, rawParams, flagUpdates):
 627	needsCarry = False
 628	if flagUpdates:
 629		for flag in flagUpdates:
 630			calc = prog.flags.flagCalc[flag]
 631			if calc == 'carry':
 632				needsCarry = True
 633	decl = ''
 634	carryCheck = _getCarryCheck(prog)
 635	size = prog.paramSize(rawParams[2])
 636	if needsCarry:
 637		decl,name = prog.getTemp(size * 2)
 638		dst = prog.carryFlowDst = name
 639	else:
 640		dst = params[2]
 641	return decl + '\n\t{dst} = {a} << {b} | {a} >> ({size} + 1 - {b}) | ({check} ? 1 : 0) << ({b} - 1);'.format(dst = dst,
 642		a = params[0], b = params[1], size=size, check=carryCheck
 643	)
 644	
 645def _rorCImpl(prog, params, rawParams, flagUpdates):
 646	size = prog.paramSize(rawParams[2])
 647	return '\n\t{dst} = {a} >> {b} | {a} << ({size} - {b});'.format(dst = params[2],
 648		a = params[0], b = params[1], size=size
 649	)
 650
 651def _rrcCImpl(prog, params, rawParams, flagUpdates):
 652	needsCarry = False
 653	if flagUpdates:
 654		for flag in flagUpdates:
 655			calc = prog.flags.flagCalc[flag]
 656			if calc == 'carry':
 657				needsCarry = True
 658	decl = ''
 659	carryCheck = _getCarryCheck(prog)
 660	size = prog.paramSize(rawParams[2])
 661	if needsCarry:
 662		decl,name = prog.getTemp(size * 2)
 663		dst = prog.carryFlowDst = name
 664	else:
 665		dst = params[2]
 666	return decl + '\n\t{dst} = {a} >> {b} | {a} << ({size} + 1 - {b}) | ({check} ? 1 : 0) << ({size}-{b});'.format(dst = dst,
 667		a = params[0], b = params[1], size=size, check=carryCheck
 668	)
 669	
 670def _updateSyncCImpl(prog, params):
 671	return '\n\t{sync}(context, target_cycle);'.format(sync=prog.sync_cycle)
 672
 673_opMap = {
 674	'mov': Op(lambda val: val).cUnaryOperator(''),
 675	'not': Op(lambda val: ~val).cUnaryOperator('~'),
 676	'lnot': Op(lambda val: 0 if val else 1).cUnaryOperator('!'),
 677	'neg': Op(lambda val: -val).cUnaryOperator('-'),
 678	'add': Op(lambda a, b: a + b).cBinaryOperator('+'),
 679	'adc': Op().addImplementation('c', 2, _adcCImpl),
 680	'sub': Op(lambda a, b: b - a).cBinaryOperator('-'),
 681	'sbc': Op().addImplementation('c', 2, _sbcCImpl),
 682	'lsl': Op(lambda a, b: a << b).cBinaryOperator('<<'),
 683	'lsr': Op(lambda a, b: a >> b).cBinaryOperator('>>'),
 684	'asr': Op(lambda a, b: a >> b).addImplementation('c', 2, _asrCImpl),
 685	'rol': Op().addImplementation('c', 2, _rolCImpl),
 686	'rlc': Op().addImplementation('c', 2, _rlcCImpl),
 687	'ror': Op().addImplementation('c', 2, _rorCImpl),
 688	'rrc': Op().addImplementation('c', 2, _rrcCImpl),
 689	'and': Op(lambda a, b: a & b).cBinaryOperator('&'),
 690	'or':  Op(lambda a, b: a | b).cBinaryOperator('|'),
 691	'xor': Op(lambda a, b: a ^ b).cBinaryOperator('^'),
 692	'abs': Op(lambda val: abs(val)).addImplementation(
 693		'c', 1, lambda prog, params: '\n\t{dst} = abs({src});'.format(dst=params[1], src=params[0])
 694	),
 695	'cmp': Op().addImplementation('c', None, _cmpCImpl),
 696	'sext': Op(_sext).addImplementation('c', 2, _sextCImpl),
 697	'ocall': Op().addImplementation('c', None, lambda prog, params: '\n\t{pre}{fun}({args});'.format(
 698		pre = prog.prefix, fun = params[0], args = ', '.join(['context'] + [str(p) for p in params[1:]])
 699	)),
 700	'cycles': Op().addImplementation('c', None,
 701		lambda prog, params: '\n\tcontext->cycles += context->opts->gen.clock_divider * {0};'.format(
 702			params[0]
 703		)
 704	),
 705	'addsize': Op(
 706		lambda a, b: b + (2 * a if a else 1)
 707	).addImplementation('c', 2, lambda prog, params: '\n\t{dst} = {val} + {sz} ? {sz} * 2 : 1;'.format(
 708		dst = params[2], sz = params[0], val = params[1]
 709	)),
 710	'decsize': Op(
 711		lambda a, b: b - (2 * a if a else 1)
 712	).addImplementation('c', 2, lambda prog, params: '\n\t{dst} = {val} - {sz} ? {sz} * 2 : 1;'.format(
 713		dst = params[2], sz = params[0], val = params[1]
 714	)),
 715	'xchg': Op().addImplementation('c', (0,1), _xchgCImpl),
 716	'dispatch': Op().addImplementation('c', None, _dispatchCImpl),
 717	'update_flags': Op().addImplementation('c', None, _updateFlagsCImpl),
 718	'update_sync': Op().addImplementation('c', None, _updateSyncCImpl)
 719}
 720
 721#represents a simple DSL instruction
 722class NormalOp:
 723	def __init__(self, parts):
 724		self.op = parts[0]
 725		self.params = parts[1:]
 726		
 727	def generate(self, prog, parent, fieldVals, output, otype, flagUpdates):
 728		procParams = []
 729		allParamsConst = flagUpdates is None and not prog.conditional
 730		opDef = _opMap.get(self.op)
 731		for param in self.params:
 732			allowConst = (self.op in prog.subroutines or len(procParams) != len(self.params) - 1) and param in parent.regValues
 733			isDst = (not opDef is None) and len(procParams) in opDef.outOp
 734			if isDst and self.op == 'xchg':
 735				#xchg uses its regs as both source and destination
 736				#we need to resolve as both so that disperse/coalesce flag stuff gets done
 737				prog.resolveParam(param, parent, fieldVals, allowConst, False)
 738			param = prog.resolveParam(param, parent, fieldVals, allowConst, isDst)
 739			
 740			if (not type(param) is int) and len(procParams) != len(self.params) - 1:
 741				allParamsConst = False
 742			procParams.append(param)
 743			
 744		if self.op == 'meta':
 745			param,_,index = self.params[1].partition('.')
 746			if index:
 747				index = (parent.resolveLocal(index) or index)
 748				if index in fieldVals:
 749					index = str(fieldVals[index])
 750				param = param + '.' + index
 751			else:
 752				param = parent.resolveLocal(param) or param
 753				if param in fieldVals:
 754					param = fieldVals[index]
 755			prog.meta[self.params[0]] = param
 756		elif self.op == 'dis':
 757			#TODO: Disassembler
 758			pass
 759		elif not opDef is None:
 760			if opDef.numParams() > len(procParams):
 761				raise Exception('Insufficient params for ' + self.op + ' (' + ', '.join(self.params) + ')')
 762			if opDef.canEval() and allParamsConst:
 763				#do constant folding
 764				if opDef.numArgs() >= len(procParams):
 765					raise Exception('Insufficient args for ' + self.op + ' (' + ', '.join(self.params) + ')')
 766				dst = self.params[opDef.numArgs()]
 767				result = opDef.evaluate(procParams[:opDef.numArgs()])
 768				while dst in prog.meta:
 769					dst = prog.meta[dst]
 770				maybeLocal = parent.resolveLocal(dst)
 771				if maybeLocal:
 772					dst = maybeLocal
 773				parent.regValues[dst] = result
 774				if prog.isReg(dst):
 775					shortProc = (procParams[0], procParams[-1])
 776					shortParams = (self.params[0], self.params[-1])
 777					output.append(_opMap['mov'].generate(otype, prog, shortProc, shortParams, None))
 778			else:
 779				output.append(opDef.generate(otype, prog, procParams, self.params, flagUpdates))
 780				for dstIdx in opDef.outOp:
 781					dst = self.params[dstIdx]
 782					while dst in prog.meta:
 783						dst = prog.meta[dst]
 784					if dst in parent.regValues:
 785						del parent.regValues[dst]
 786					
 787		elif self.op in prog.subroutines:
 788			procParams = []
 789			for param in self.params:
 790				begin,sep,end = param.partition('.')
 791				if sep:
 792					if end in fieldVals:
 793						param = begin + '.' + str(fieldVals[end])
 794				else:
 795					if param in fieldVals:
 796						param = fieldVals[param]
 797				procParams.append(param)
 798			prog.subroutines[self.op].inline(prog, procParams, output, otype, parent)
 799		else:
 800			output.append('\n\t' + self.op + '(' + ', '.join([str(p) for p in procParams]) + ');')
 801		prog.lastOp = self
 802	
 803	def __str__(self):
 804		return '\n\t' + self.op + ' ' + ' '.join(self.params)
 805		
 806#represents a DSL switch construct
 807class Switch(ChildBlock):
 808	def __init__(self, parent, param):
 809		self.op = 'switch'
 810		self.parent = parent
 811		self.param = param
 812		self.cases = {}
 813		self.regValues = None
 814		self.current_locals = {}
 815		self.case_locals = {}
 816		self.current_case = None
 817		self.default = None
 818		self.default_locals = None
 819	
 820	def addOp(self, op):
 821		if op.op == 'case':
 822			val = int(op.params[0], 16) if op.params[0].startswith('0x') else int(op.params[0])
 823			self.cases[val] = self.current_case = []
 824			self.case_locals[val] = self.current_locals = {}
 825		elif op.op == 'default':
 826			self.default = self.current_case = []
 827			self.default_locals = self.current_locals = {}
 828		elif self.current_case == None:
 829			raise ion('Orphan instruction in switch')
 830		elif op.op == 'local':
 831			name = op.params[0]
 832			size = op.params[1]
 833			self.current_locals[name] = size
 834		else:
 835			self.current_case.append(op)
 836			
 837	def resolveLocal(self, name):
 838		if name in self.current_locals:
 839			return name
 840		return self.parent.resolveLocal(name)
 841	
 842	def addLocal(self, name, size):
 843		self.current_locals[name] = size
 844		
 845	def localSize(self, name):
 846		if name in self.current_locals:
 847			return self.current_locals[name]
 848		return self.parent.localSize(name)
 849			
 850	def generate(self, prog, parent, fieldVals, output, otype, flagUpdates):
 851		prog.pushScope(self)
 852		param = prog.resolveParam(self.param, parent, fieldVals)
 853		if type(param) is int:
 854			self.regValues = self.parent.regValues
 855			if param in self.cases:
 856				self.current_locals = self.case_locals[param]
 857				output.append('\n\t{')
 858				for local in self.case_locals[param]:
 859					output.append('\n\tuint{0}_t {1};'.format(self.case_locals[param][local], local))
 860				self.processOps(prog, fieldVals, output, otype, self.cases[param])
 861				output.append('\n\t}')
 862			elif self.default:
 863				self.current_locals = self.default_locals
 864				output.append('\n\t{')
 865				for local in self.default_locals:
 866					output.append('\n\tuint{0}_t {1};'.format(self.default[local], local))
 867				self.processOps(prog, fieldVals, output, otype, self.default)
 868				output.append('\n\t}')
 869		else:
 870			oldCond = prog.conditional
 871			prog.conditional = True
 872			output.append('\n\tswitch(' + param + ')')
 873			output.append('\n\t{')
 874			for case in self.cases:
 875				temp = prog.temp.copy()
 876				self.current_locals = self.case_locals[case]
 877				self.regValues = dict(self.parent.regValues)
 878				output.append('\n\tcase {0}U: '.format(case) + '{')
 879				for local in self.case_locals[case]:
 880					output.append('\n\tuint{0}_t {1};'.format(self.case_locals[case][local], local))
 881				self.processOps(prog, fieldVals, output, otype, self.cases[case])
 882				output.append('\n\tbreak;')
 883				output.append('\n\t}')
 884				prog.temp = temp
 885			if self.default:
 886				temp = prog.temp.copy()
 887				self.current_locals = self.default_locals
 888				self.regValues = dict(self.parent.regValues)
 889				output.append('\n\tdefault: {')
 890				for local in self.default_locals:
 891					output.append('\n\tuint{0}_t {1};'.format(self.default_locals[local], local))
 892				self.processOps(prog, fieldVals, output, otype, self.default)
 893				prog.temp = temp
 894			output.append('\n\t}')
 895			prog.conditional = oldCond
 896		prog.popScope()
 897	
 898	def __str__(self):
 899		keys = self.cases.keys()
 900		keys.sort()
 901		lines = ['\n\tswitch']
 902		for case in keys:
 903			lines.append('\n\tcase {0}'.format(case))
 904			lines.append(''.join([str(op) for op in self.cases[case]]))
 905		lines.append('\n\tend')
 906		return ''.join(lines)
 907
 908		
 909def _geuCImpl(prog, parent, fieldVals, output):
 910	if prog.lastOp.op == 'cmp':
 911		output.pop()
 912		params = [prog.resolveParam(p, parent, fieldVals) for p in prog.lastOp.params]
 913		return '\n\tif ({a} >= {b}) '.format(a=params[1], b = params[0]) + '{'
 914	else:
 915		raise Exception(">=U not implemented in the general case yet")
 916
 917def _eqCImpl(prog, parent, fieldVals, output):
 918	return '\n\tif (!{a}) {'.format(a=prog.resolveParam(prog.lastDst, None, {}))
 919
 920def _neqCImpl(prog, parent, fieldVals, output):
 921	return '\n\tif ({a}) {'.format(a=prog.resolveParam(prog.lastDst, None, {}))
 922	
 923_ifCmpImpl = {
 924	'c': {
 925		'>=U': _geuCImpl,
 926		'=': _eqCImpl,
 927		'!=': _neqCImpl
 928	}
 929}
 930#represents a DSL conditional construct
 931class If(ChildBlock):
 932	def __init__(self, parent, cond):
 933		self.op = 'if'
 934		self.parent = parent
 935		self.cond = cond
 936		self.body = []
 937		self.elseBody = []
 938		self.curBody = self.body
 939		self.locals = {}
 940		self.elseLocals = {}
 941		self.curLocals = self.locals
 942		self.regValues = None
 943		
 944	def addOp(self, op):
 945		if op.op in ('case', 'arg'):
 946			raise Exception(self.op + ' is not allows inside an if block')
 947		if op.op == 'local':
 948			name = op.params[0]
 949			size = op.params[1]
 950			self.curLocals[name] = size
 951		elif op.op == 'else':
 952			self.curLocals = self.elseLocals
 953			self.curBody = self.elseBody
 954		else:
 955			self.curBody.append(op)
 956			
 957	def localSize(self, name):
 958		return self.curLocals.get(name)
 959		
 960	def resolveLocal(self, name):
 961		if name in self.curLocals:
 962			return name
 963		return self.parent.resolveLocal(name)
 964		
 965	def _genTrueBody(self, prog, fieldVals, output, otype):
 966		self.curLocals = self.locals
 967		subOut = []
 968		self.processOps(prog, fieldVals, subOut, otype, self.body)
 969		for local in self.locals:
 970			output.append('\n\tuint{sz}_t {nm};'.format(sz=self.locals[local], nm=local))
 971		output += subOut
 972			
 973	def _genFalseBody(self, prog, fieldVals, output, otype):
 974		self.curLocals = self.elseLocals
 975		subOut = []
 976		self.processOps(prog, fieldVals, subOut, otype, self.elseBody)
 977		for local in self.elseLocals:
 978			output.append('\n\tuint{sz}_t {nm};'.format(sz=self.elseLocals[local], nm=local))
 979		output += subOut
 980	
 981	def _genConstParam(self, param, prog, fieldVals, output, otype):
 982		if param:
 983			self._genTrueBody(prog, fieldVals, output, otype)
 984		else:
 985			self._genFalseBody(prog, fieldVals, output, otype)
 986			
 987	def generate(self, prog, parent, fieldVals, output, otype, flagUpdates):
 988		self.regValues = parent.regValues
 989		try:
 990			self._genConstParam(prog.checkBool(self.cond), prog, fieldVals, output, otype)
 991		except Exception:
 992			if self.cond in _ifCmpImpl[otype]:
 993				oldCond = prog.conditional
 994				prog.conditional = True
 995				temp = prog.temp.copy()
 996				output.append(_ifCmpImpl[otype][self.cond](prog, parent, fieldVals, output))
 997				self._genTrueBody(prog, fieldVals, output, otype)
 998				prog.temp = temp
 999				if self.elseBody:
1000					temp = prog.temp.copy()
1001					output.append('\n\t} else {')
1002					self._genFalseBody(prog, fieldVals, output, otype)
1003					prog.temp = temp
1004				output.append('\n\t}')
1005				prog.conditional = oldCond
1006			else:
1007				cond = prog.resolveParam(self.cond, parent, fieldVals)
1008				if type(cond) is int:
1009					self._genConstParam(cond, prog, fieldVals, output, otype)
1010				else:
1011					temp = prog.temp.copy()
1012					output.append('\n\tif ({cond}) '.format(cond=cond) + '{')
1013					oldCond = prog.conditional
1014					prog.conditional = True
1015					self._genTrueBody(prog, fieldVals, output, otype)
1016					prog.temp = temp
1017					if self.elseBody:
1018						temp = prog.temp.copy()
1019						output.append('\n\t} else {')
1020						self._genFalseBody(prog, fieldVals, output, otype)
1021						prog.temp = temp
1022					output.append('\n\t}')
1023					prog.conditional = oldCond
1024						
1025	
1026	def __str__(self):
1027		lines = ['\n\tif']
1028		for op in self.body:
1029			lines.append(str(op))
1030		lines.append('\n\tend')
1031		return ''.join(lines)
1032
1033class Registers:
1034	def __init__(self):
1035		self.regs = {}
1036		self.pointers = {}
1037		self.regArrays = {}
1038		self.regToArray = {}
1039		self.addReg('cycles', 32)
1040		self.addReg('sync_cycle', 32)
1041	
1042	def addReg(self, name, size):
1043		self.regs[name] = size
1044		
1045	def addPointer(self, name, size, count):
1046		self.pointers[name] = (size, count)
1047	
1048	def addRegArray(self, name, size, regs):
1049		self.regArrays[name] = (size, regs)
1050		idx = 0
1051		if not type(regs) is int:
1052			for reg in regs:
1053				self.regs[reg] = size
1054				self.regToArray[reg] = (name, idx)
1055				idx += 1
1056	
1057	def isReg(self, name):
1058		return name in self.regs
1059	
1060	def isRegArray(self, name):
1061		return name in self.regArrays
1062		
1063	def isRegArrayMember(self, name):
1064		return name in self.regToArray
1065		
1066	def arrayMemberParent(self, name):
1067		return self.regToArray[name][0]
1068	
1069	def arrayMemberIndex(self, name):
1070		return self.regToArray[name][1]
1071	
1072	def arrayMemberName(self, array, index):
1073		if type(index) is int and not type(self.regArrays[array][1]) is int:
1074			return self.regArrays[array][1][index]
1075		else:
1076			return None
1077			
1078	def isNamedArray(self, array):
1079		return array in self.regArrays and type(self.regArrays[array][1]) is int
1080	
1081	def processLine(self, parts):
1082		if len(parts) == 3:
1083			if parts[1].startswith('ptr'):
1084				self.addPointer(parts[0], parts[1][3:], int(parts[2]))
1085			else:
1086				self.addRegArray(parts[0], int(parts[1]), int(parts[2]))
1087		elif len(parts) > 2:
1088			self.addRegArray(parts[0], int(parts[1]), parts[2:])
1089		else:
1090			if parts[1].startswith('ptr'):
1091				self.addPointer(parts[0], parts[1][3:], 1)
1092			else:
1093				self.addReg(parts[0], int(parts[1]))
1094		return self
1095
1096	def writeHeader(self, otype, hFile):
1097		fieldList = []
1098		for pointer in self.pointers:
1099			stars = '*'
1100			ptype, count = self.pointers[pointer]
1101			while ptype.startswith('ptr'):
1102				stars += '*'
1103				ptype = ptype[3:]
1104			if ptype.isdigit():
1105				ptype = 'uint{sz}_t'.format(sz=ptype)
1106			if count > 1:
1107				arr = '[{n}]'.format(n=count)
1108			else:
1109				arr = ''
1110			hFile.write('\n\t{ptype} {stars}{nm}{arr};'.format(nm=pointer, ptype=ptype, stars=stars, arr=arr))
1111		for reg in self.regs:
1112			if not self.isRegArrayMember(reg):
1113				fieldList.append((self.regs[reg], 1, reg))
1114		for arr in self.regArrays:
1115			size,regs = self.regArrays[arr]
1116			if not type(regs) is int:
1117				regs = len(regs)
1118			fieldList.append((size, regs, arr))
1119		fieldList.sort()
1120		fieldList.reverse()
1121		for size, count, name in fieldList:
1122			if count > 1:
1123				hFile.write('\n\tuint{sz}_t {nm}[{ct}];'.format(sz=size, nm=name, ct=count))
1124			else:
1125				hFile.write('\n\tuint{sz}_t {nm};'.format(sz=size, nm=name))
1126	
1127class Flags:
1128	def __init__(self):
1129		self.flagBits = {}
1130		self.flagCalc = {}
1131		self.flagStorage = {}
1132		self.flagReg = None
1133		self.storageToFlags = {}
1134		self.maxBit = -1
1135	
1136	def processLine(self, parts):
1137		if parts[0] == 'register':
1138			self.flagReg = parts[1]
1139		else:
1140			flag,bit,calc,storage = parts
1141			bit,_,top = bit.partition('-')
1142			bit = int(bit)
1143			if top:
1144				top = int(bit)
1145				if top > self.maxBit:
1146					self.maxBit = top
1147				self.flagBits[flag] = (bit,top)
1148			else:
1149				if bit > self.maxBit:
1150					self.maxBit = bit
1151				self.flagBits[flag] = bit
1152			self.flagCalc[flag] = calc
1153			self.flagStorage[flag] = storage
1154			storage,_,storebit = storage.partition('.')
1155			self.storageToFlags.setdefault(storage, []).append((storebit, flag))
1156		return self
1157	
1158	def getStorage(self, flag):
1159		if not flag in self.flagStorage:
1160			raise Exception('Undefined flag ' + flag)
1161		loc,_,bit = self.flagStorage[flag].partition('.')
1162		if bit:
1163			return (loc, int(bit))
1164		else:
1165			return loc 
1166	
1167	def parseFlagUpdate(self, flagString):
1168		last = ''
1169		autoUpdate = set()
1170		explicit = {}
1171		for c in flagString:
1172			if c.isdigit():
1173				if last.isalpha():
1174					num = int(c)
1175					if num > 1:
1176						raise Exception(c + ' is not a valid digit for update_flags')
1177					explicit[last] = num
1178					last = c
1179				else:
1180					raise Exception('Digit must follow flag letter in update_flags')
1181			else:
1182				if last.isalpha():
1183					autoUpdate.add(last)
1184				last = c
1185		if last.isalpha():
1186			autoUpdate.add(last)
1187		return (autoUpdate, explicit)
1188	
1189	def disperseFlags(self, prog, otype):
1190		bitToFlag = [None] * (self.maxBit+1)
1191		src = prog.resolveReg(self.flagReg, None, {})
1192		output = []
1193		for flag in self.flagBits:
1194			bit = self.flagBits[flag]
1195			if type(bit) is tuple:
1196				bot,top = bit
1197				mask = ((1 << (top + 1 - bot)) - 1) << bot
1198				output.append('\n\t{dst} = {src} & mask;'.format(
1199					dst=prog.resolveReg(self.flagStorage[flag], None, {}), src=src, mask=mask
1200				))
1201			else:
1202				bitToFlag[self.flagBits[flag]] = flag		
1203		multi = {}
1204		for bit in range(len(bitToFlag)-1,-1,-1):
1205			flag = bitToFlag[bit]
1206			if not flag is None:
1207				field,_,dstbit = self.flagStorage[flag].partition('.')
1208				dst = prog.resolveReg(field, None, {})
1209				if dstbit:
1210					dstbit = int(dstbit)
1211					multi.setdefault(dst, []).append((dstbit, bit))
1212				else:
1213					output.append('\n\t{dst} = {src} & {mask};'.format(dst=dst, src=src, mask=(1 << bit)))
1214		for dst in multi:
1215			didClear = False
1216			direct = []
1217			for dstbit, bit in multi[dst]:
1218				if dstbit == bit:
1219					direct.append(bit)
1220				else:
1221					if not didClear:
1222						output.append('\n\t{dst} = 0;'.format(dst=dst))
1223						didClear = True
1224					if dstbit > bit:
1225						shift = '<<'
1226						diff = dstbit - bit
1227					else:
1228						shift = '>>'
1229						diff = bit - dstbit
1230					output.append('\n\t{dst} |= {src} {shift} {diff} & {mask};'.format(
1231						src=src, dst=dst, shift=shift, diff=diff, mask=(1 << dstbit)
1232					))
1233			if direct:
1234				if len(direct) == len(multi[dst]):
1235					output.append('\n\t{dst} = {src};'.format(dst=dst, src=src))
1236				else:
1237					mask = 0
1238					for bit in direct:
1239						mask = mask | (1 << bit)
1240					output.append('\n\t{dst} = {src} & {mask};'.format(dst=dst, src=src, mask=mask))
1241		return ''.join(output)
1242	
1243	def coalesceFlags(self, prog, otype):
1244		dst = prog.resolveReg(self.flagReg, None, {})
1245		output = ['\n\t{dst} = 0;'.format(dst=dst)]
1246		bitToFlag = [None] * (self.maxBit+1)
1247		for flag in self.flagBits:
1248			bit = self.flagBits[flag]
1249			if type(bit) is tuple:
1250				bot,_ = bit
1251				src = prog.resolveReg(self.flagStorage[flag], None, {})
1252				if bot:
1253					output.append('\n\t{dst} |= {src} << {shift};'.format(
1254						dst=dst, src = src, shift = bot
1255					))
1256				else:
1257					output.append('\n\t{dst} |= {src};'.format(
1258						dst=dst, src = src
1259					))
1260			else:
1261				bitToFlag[bit] = flag
1262		multi = {}
1263		for bit in range(len(bitToFlag)-1,-1,-1):
1264			flag = bitToFlag[bit]
1265			if not flag is None:
1266				field,_,srcbit = self.flagStorage[flag].partition('.')
1267				src = prog.resolveReg(field, None, {})
1268				if srcbit:
1269					srcbit = int(srcbit)
1270					multi.setdefault(src, []).append((srcbit,bit))
1271				else:
1272					output.append('\n\tif ({src}) {{\n\t\t{dst} |= 1 << {bit};\n\t}}'.format(
1273						dst=dst, src=src, bit=bit
1274					))
1275		for src in multi:
1276			direct = 0
1277			for srcbit, dstbit in multi[src]:
1278				if srcbit == dstbit:
1279					direct = direct | (1 << srcbit)
1280				else:
1281					output.append('\n\tif ({src} & (1 << {srcbit})) {{\n\t\t{dst} |= 1 << {dstbit};\n\t}}'.format(
1282						src=src, dst=dst, srcbit=srcbit, dstbit=dstbit
1283					))
1284			if direct:
1285				output.append('\n\t{dst} |= {src} & {mask};'.format(
1286					dst=dst, src=src, mask=direct
1287				))
1288		return ''.join(output)
1289		
1290		
1291class Program:
1292	def __init__(self, regs, instructions, subs, info, flags):
1293		self.regs = regs
1294		self.instructions = instructions
1295		self.subroutines = subs
1296		self.meta = {}
1297		self.booleans = {}
1298		self.prefix = info.get('prefix', [''])[0]
1299		self.opsize = int(info.get('opcode_size', ['8'])[0])
1300		self.extra_tables = info.get('extra_tables', [])
1301		self.context_type = self.prefix + 'context'
1302		self.body = info.get('body', [None])[0]
1303		self.interrupt = info.get('interrupt', [None])[0]
1304		self.sync_cycle = info.get('sync_cycle', [None])[0]
1305		self.includes = info.get('include', [])
1306		self.flags = flags
1307		self.lastDst = None
1308		self.scopes = []
1309		self.currentScope = None
1310		self.lastOp = None
1311		self.carryFlowDst = None
1312		self.lastA = None
1313		self.lastB = None
1314		self.lastBFlow = None
1315		self.conditional = False
1316		self.declares = []
1317		
1318	def __str__(self):
1319		pieces = []
1320		for reg in self.regs:
1321			pieces.append(str(self.regs[reg]))
1322		for name in self.subroutines:
1323			pieces.append('\n'+str(self.subroutines[name]))
1324		for instruction in self.instructions:
1325			pieces.append('\n'+str(instruction))
1326		return ''.join(pieces)
1327		
1328	def writeHeader(self, otype, header):
1329		hFile = open(header, 'w')
1330		macro = header.upper().replace('.', '_')
1331		hFile.write('#ifndef {0}_'.format(macro))
1332		hFile.write('\n#define {0}_'.format(macro))
1333		hFile.write('\n#include "backend.h"')
1334		hFile.write('\n\ntypedef struct {')
1335		hFile.write('\n\tcpu_options gen;')
1336		hFile.write('\n}} {0}options;'.format(self.prefix))
1337		hFile.write('\n\ntypedef struct {')
1338		hFile.write('\n\t{0}options *opts;'.format(self.prefix))
1339		self.regs.writeHeader(otype, hFile)
1340		hFile.write('\n}} {0}context;'.format(self.prefix))
1341		hFile.write('\n')
1342		hFile.write('\nvoid {pre}execute({type} *context, uint32_t target_cycle);'.format(pre = self.prefix, type = self.context_type))
1343		for decl in self.declares:
1344			hFile.write('\n' + decl)
1345		hFile.write('\n#endif //{0}_'.format(macro))
1346		hFile.write('\n')
1347		hFile.close()
1348		
1349	def _buildTable(self, otype, table, body, lateBody):
1350		pieces = []
1351		opmap = [None] * (1 << self.opsize)
1352		bodymap = {}
1353		if table in self.instructions:
1354			instructions = self.instructions[table]
1355			instructions.sort()
1356			for inst in instructions:
1357				for val in inst.allValues():
1358					if opmap[val] is None:
1359						self.meta = {}
1360						self.temp = {}
1361						self.needFlagCoalesce = False
1362						self.needFlagDisperse = False
1363						self.lastOp = None
1364						opmap[val] = inst.generateName(val)
1365						bodymap[val] = inst.generateBody(val, self, otype)
1366		
1367		if self.dispatch == 'call':
1368			pieces.append('\nstatic impl_fun impl_{name}[{sz}] = {{'.format(name = table, sz=len(opmap)))
1369			for inst in range(0, len(opmap)):
1370				op = opmap[inst]
1371				if op is None:
1372					pieces.append('\n\tunimplemented,')
1373				else:
1374					pieces.append('\n\t' + op + ',')
1375					body.append(bodymap[inst])
1376			pieces.append('\n};')
1377		elif self.dispatch == 'goto':
1378			body.append('\n\tstatic void *impl_{name}[{sz}] = {{'.format(name = table, sz=len(opmap)))
1379			for inst in range(0, len(opmap)):
1380				op = opmap[inst]
1381				if op is None:
1382					body.append('\n\t\t&&unimplemented,')
1383				else:
1384					body.append('\n\t\t&&' + op + ',')
1385					lateBody.append(bodymap[inst])
1386			body.append('\n\t};')
1387		else:
1388			raise Exception("unimplmeneted dispatch type " + self.dispatch)
1389		body.extend(pieces)
1390		
1391	def nextInstruction(self, otype):
1392		output = []
1393		if self.dispatch == 'goto':
1394			if self.interrupt in self.subroutines:
1395				output.append('\n\tif (context->cycles >= context->sync_cycle) {')
1396			output.append('\n\tif (context->cycles >= target_cycle) { return; }')
1397			if self.interrupt in self.subroutines:
1398				self.meta = {}
1399				self.temp = {}
1400				self.subroutines[self.interrupt].inline(self, [], output, otype, None)
1401				output.append('\n\t}')
1402			
1403			self.meta = {}
1404			self.temp = {}
1405			self.subroutines[self.body].inline(self, [], output, otype, None)
1406		return output
1407	
1408	def build(self, otype):
1409		body = []
1410		pieces = []
1411		for include in self.includes:
1412			body.append('#include "{0}"\n'.format(include))
1413		if self.dispatch == 'call':
1414			body.append('\nstatic void unimplemented({pre}context *context)'.format(pre = self.prefix))
1415			body.append('\n{')
1416			body.append('\n\tfatal_error("Unimplemented instruction\\n");')
1417			body.append('\n}\n')
1418			body.append('\ntypedef void (*impl_fun)({pre}context *context);'.format(pre=self.prefix))
1419			for table in self.extra_tables:
1420				body.append('\nstatic impl_fun impl_{name}[{sz}];'.format(name = table, sz=(1 << self.opsize)))
1421			body.append('\nstatic impl_fun impl_main[{sz}];'.format(sz=(1 << self.opsize)))
1422		elif self.dispatch == 'goto':
1423			body.append('\nvoid {pre}execute({type} *context, uint32_t target_cycle)'.format(pre = self.prefix, type = self.context_type))
1424			body.append('\n{')
1425			
1426		for table in self.extra_tables:
1427			self._buildTable(otype, table, body, pieces)
1428		self._buildTable(otype, 'main', body, pieces)
1429		if self.dispatch == 'call' and self.body in self.subroutines:
1430			pieces.append('\nvoid {pre}execute({type} *context, uint32_t target_cycle)'.format(pre = self.prefix, type = self.context_type))
1431			pieces.append('\n{')
1432			pieces.append('\n\t{sync}(context, target_cycle);'.format(sync=self.sync_cycle))
1433			pieces.append('\n\twhile (context->cycles < target_cycle)')
1434			pieces.append('\n\t{')
1435			#TODO: Handle interrupts in call dispatch mode
1436			self.meta = {}
1437			self.temp = {}
1438			self.subroutines[self.body].inline(self, [], pieces, otype, None)
1439			pieces.append('\n\t}')
1440			pieces.append('\n}')
1441		elif self.dispatch == 'goto':
1442			body.append('\n\t{sync}(context, target_cycle);'.format(sync=self.sync_cycle))
1443			body += self.nextInstruction(otype)
1444			pieces.append('\nunimplemented:')
1445			pieces.append('\n\tfatal_error("Unimplemented instruction\\n");')
1446			pieces.append('\n}')
1447		return ''.join(body) +  ''.join(pieces)
1448		
1449	def checkBool(self, name):
1450		if not name in self.booleans:
1451			raise Exception(name + ' is not a defined boolean flag')
1452		return self.booleans[name]
1453	
1454	def getTemp(self, size):
1455		if size in self.temp:
1456			return ('', self.temp[size])
1457		self.temp[size] = 'gen_tmp{sz}__'.format(sz=size);
1458		return ('\n\tuint{sz}_t gen_tmp{sz}__;'.format(sz=size), self.temp[size])
1459		
1460	def resolveParam(self, param, parent, fieldVals, allowConstant=True, isdst=False):
1461		keepGoing = True
1462		while keepGoing:
1463			keepGoing = False
1464			try:
1465				if type(param) is int:
1466					pass
1467				elif param.startswith('0x'):
1468					param = int(param, 16)
1469				else:
1470					param = int(param)
1471			except ValueError:
1472				
1473				if parent:
1474					if param in parent.regValues and allowConstant:
1475						return parent.regValues[param]
1476					maybeLocal = parent.resolveLocal(param)
1477					if maybeLocal:
1478						if isdst:
1479							self.lastDst = param
1480						return maybeLocal
1481				if param in fieldVals:
1482					param = fieldVals[param]
1483					fieldVals = {}
1484					keepGoing = True
1485				elif param in self.meta:
1486					param = self.meta[param]
1487					keepGoing = True
1488				elif self.isReg(param):
1489					return self.resolveReg(param, parent, fieldVals, isdst)
1490		if isdst:
1491			self.lastDst = param
1492		return param
1493	
1494	def isReg(self, name):
1495		if not type(name) is str:
1496			return False
1497		begin,sep,_ = name.partition('.')
1498		if sep:
1499			if begin in self.meta:
1500				begin = self.meta[begin]
1501			return self.regs.isRegArray(begin)
1502		else:
1503			return self.regs.isReg(name)
1504	
1505	def resolveReg(self, name, parent, fieldVals, isDst=False):
1506		begin,sep,end = name.partition('.')
1507		if sep:
1508			if begin in self.meta:
1509				begin = self.meta[begin]
1510			if not self.regs.isRegArrayMember(end):
1511				end = self.resolveParam(end, parent, fieldVals)
1512			if not type(end) is int and self.regs.isRegArrayMember(end):
1513				arrayName = self.regs.arrayMemberParent(end)
1514				end = self.regs.arrayMemberIndex(end)
1515				if arrayName != begin:
1516					end = 'context->{0}[{1}]'.format(arrayName, end)
1517			if self.regs.isNamedArray(begin):
1518				regName = self.regs.arrayMemberName(begin, end)
1519			else:
1520				regName = '{0}.{1}'.format(begin, end)
1521			ret = 'context->{0}[{1}]'.format(begin, end)
1522		else:
1523			regName = name
1524			if self.regs.isRegArrayMember(name):
1525				arr,idx = self.regs.regToArray[name]
1526				ret = 'context->{0}[{1}]'.format(arr, idx)
1527			else:
1528				ret = 'context->' + name
1529		if regName == self.flags.flagReg:
1530			if isDst:
1531				self.needFlagDisperse = True
1532			else:
1533				self.needFlagCoalesce = True
1534		if isDst:
1535			self.lastDst = regName
1536		return ret
1537		
1538	
1539	
1540	def paramSize(self, name):
1541		if name in self.meta:
1542			return self.paramSize(self.meta[name])
1543		for i in range(len(self.scopes) -1, -1, -1):
1544			size = self.scopes[i].localSize(name)
1545			if size:
1546				return size
1547		begin,sep,_ = name.partition('.')
1548		if sep and self.regs.isRegArray(begin):
1549			return self.regs.regArrays[begin][0]
1550		if self.regs.isReg(name):
1551			return self.regs.regs[name]
1552		return 32
1553	
1554	def pushScope(self, scope):
1555		self.scopes.append(scope)
1556		self.currentScope = scope
1557		
1558	def popScope(self):
1559		ret = self.scopes.pop()
1560		self.currentScope = self.scopes[-1] if self.scopes else None
1561		return ret
1562		
1563	def getRootScope(self):
1564		return self.scopes[0]
1565
1566def parse(args):
1567	f = args.source
1568	instructions = {}
1569	subroutines = {}
1570	registers = None
1571	flags = None
1572	declares = []
1573	errors = []
1574	info = {}
1575	line_num = 0
1576	cur_object = None
1577	for line in f:
1578		line_num += 1
1579		line,_,comment = line.partition('#')
1580		if not line.strip():
1581			continue
1582		if line[0].isspace():
1583			if not cur_object is None:
1584				sep = True
1585				parts = []
1586				while sep:
1587					before,sep,after = line.partition('"')
1588					before = before.strip()
1589					if before:
1590						parts += [el.strip() for el in before.split(' ')]
1591					if sep:
1592						#TODO: deal with escaped quotes
1593						inside,sep,after = after.partition('"')
1594						parts.append('"' + inside + '"')
1595					line = after
1596				if type(cur_object) is dict:
1597					cur_object[parts[0]] = parts[1:]
1598				elif type(cur_object) is list:
1599					cur_object.append(' '.join(parts))
1600				else:
1601					cur_object = cur_object.processLine(parts)
1602				
1603#				if type(cur_object) is Registers:
1604#					if len(parts) > 2:
1605#						cur_object.addRegArray(parts[0], int(parts[1]), parts[2:])
1606#					else:
1607#						cur_object.addReg(parts[0], int(parts[1]))
1608#				elif type(cur_object) is dict:
1609#					cur_object[parts[0]] = parts[1:]
1610#				elif parts[0] == 'switch':
1611#					o = Switch(cur_object, parts[1])
1612#					cur_object.addOp(o)
1613#					cur_object = o
1614#				elif parts[0] == 'if':
1615#					o = If(cur_object, parts[1])
1616#					cur_object.addOp(o)
1617#					cur_object = o
1618#				elif parts[0] == 'end':
1619#					cur_object = cur_object.parent
1620#				else:
1621#					cur_object.addOp(NormalOp(parts))
1622			else:
1623				errors.append("Orphan instruction on line {0}".format(line_num))
1624		else:
1625			parts = line.split(' ')
1626			if len(parts) > 1:
1627				if len(parts) > 2:
1628					table,bitpattern,name = parts
1629				else:
1630					bitpattern,name = parts
1631					table = 'main'
1632				value = 0
1633				fields = {}
1634				curbit = len(bitpattern) - 1
1635				for char in bitpattern:
1636					value <<= 1
1637					if char in ('0', '1'):
1638						value |= int(char)
1639					else:
1640						if char in fields:
1641							fields[char] = (curbit, fields[char][1] + 1)
1642						else:
1643							fields[char] = (curbit, 1)
1644					curbit -= 1
1645				cur_object = Instruction(value, fields, name.strip())
1646				instructions.setdefault(table, []).append(cur_object)
1647			elif line.strip() == 'regs':
1648				if registers is None:
1649					registers = Registers()
1650				cur_object = registers
1651			elif line.strip() == 'info':
1652				cur_object = info
1653			elif line.strip() == 'flags':
1654				if flags is None:
1655					flags = Flags()
1656				cur_object = flags
1657			elif line.strip() == 'declare':
1658				cur_object = declares
1659			else:
1660				cur_object = SubRoutine(line.strip())
1661				subroutines[cur_object.name] = cur_object
1662	if errors:
1663		print(errors)
1664	else:
1665		p = Program(registers, instructions, subroutines, info, flags)
1666		p.dispatch = args.dispatch
1667		p.declares = declares
1668		p.booleans['dynarec'] = False
1669		p.booleans['interp'] = True
1670		if args.define:
1671			for define in args.define:
1672				name,sep,val = define.partition('=')
1673				name = name.strip()
1674				val = val.strip()
1675				if sep:
1676					p.booleans[name] = bool(val)
1677				else:
1678					p.booleans[name] = True
1679		
1680		if 'header' in info:
1681			print('#include "{0}"'.format(info['header'][0]))
1682			p.writeHeader('c', info['header'][0])
1683		print('#include "util.h"')
1684		print('#include <stdlib.h>')
1685		print(p.build('c'))
1686
1687def main(argv):
1688	from argparse import ArgumentParser, FileType
1689	argParser = ArgumentParser(description='CPU emulator DSL compiler')
1690	argParser.add_argument('source', type=FileType('r'))
1691	argParser.add_argument('-D', '--define', action='append')
1692	argParser.add_argument('-d', '--dispatch', choices=('call', 'switch', 'goto'), default='call')
1693	parse(argParser.parse_args(argv[1:]))
1694
1695if __name__ == '__main__':
1696	from sys import argv
1697	main(argv)