stk500v2.py 8.0 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220
  1. """
  2. STK500v2 protocol implementation for programming AVR chips.
  3. The STK500v2 protocol is used by the ArduinoMega2560 and a few other Arduino platforms to load firmware.
  4. This is a python 3 conversion of the code created by David Braam for the Cura project.
  5. """
  6. import os
  7. import struct
  8. import sys
  9. import time
  10. from serial import Serial
  11. from serial import SerialException
  12. from serial import SerialTimeoutException
  13. from . import ispBase, intelHex
  14. class Stk500v2(ispBase.IspBase):
  15. def __init__(self):
  16. self.serial = None
  17. self.seq = 1
  18. self.last_addr = -1
  19. self.progress_callback = None
  20. def connect(self, port = "COM22", speed = 115200):
  21. if self.serial is not None:
  22. self.close()
  23. try:
  24. self.serial = Serial(str(port), speed, timeout=1, writeTimeout=10000)
  25. except SerialException as e:
  26. raise ispBase.IspError("Failed to open serial port")
  27. except:
  28. raise ispBase.IspError("Unexpected error while connecting to serial port:" + port + ":" + str(sys.exc_info()[0]))
  29. self.seq = 1
  30. #Reset the controller
  31. for n in range(0, 2):
  32. self.serial.setDTR(True)
  33. time.sleep(0.1)
  34. self.serial.setDTR(False)
  35. time.sleep(0.1)
  36. time.sleep(0.2)
  37. self.serial.flushInput()
  38. self.serial.flushOutput()
  39. if self.sendMessage([0x10, 0xc8, 0x64, 0x19, 0x20, 0x00, 0x53, 0x03, 0xac, 0x53, 0x00, 0x00]) != [0x10, 0x00]:
  40. self.close()
  41. raise ispBase.IspError("Failed to enter programming mode")
  42. self.sendMessage([0x06, 0x80, 0x00, 0x00, 0x00])
  43. if self.sendMessage([0xEE])[1] == 0x00:
  44. self._has_checksum = True
  45. else:
  46. self._has_checksum = False
  47. self.serial.timeout = 5
  48. def close(self):
  49. if self.serial is not None:
  50. self.serial.close()
  51. self.serial = None
  52. #Leave ISP does not reset the serial port, only resets the device, and returns the serial port after disconnecting it from the programming interface.
  53. # This allows you to use the serial port without opening it again.
  54. def leaveISP(self):
  55. if self.serial is not None:
  56. if self.sendMessage([0x11]) != [0x11, 0x00]:
  57. raise ispBase.IspError("Failed to leave programming mode")
  58. ret = self.serial
  59. self.serial = None
  60. return ret
  61. return None
  62. def isConnected(self):
  63. return self.serial is not None
  64. def hasChecksumFunction(self):
  65. return self._has_checksum
  66. def sendISP(self, data):
  67. recv = self.sendMessage([0x1D, 4, 4, 0, data[0], data[1], data[2], data[3]])
  68. return recv[2:6]
  69. def writeFlash(self, flash_data):
  70. #Set load addr to 0, in case we have more then 64k flash we need to enable the address extension
  71. page_size = self.chip["pageSize"] * 2
  72. flash_size = page_size * self.chip["pageCount"]
  73. print("Writing flash")
  74. if flash_size > 0xFFFF:
  75. self.sendMessage([0x06, 0x80, 0x00, 0x00, 0x00])
  76. else:
  77. self.sendMessage([0x06, 0x00, 0x00, 0x00, 0x00])
  78. load_count = (len(flash_data) + page_size - 1) / page_size
  79. for i in range(0, int(load_count)):
  80. recv = self.sendMessage([0x13, page_size >> 8, page_size & 0xFF, 0xc1, 0x0a, 0x40, 0x4c, 0x20, 0x00, 0x00] + flash_data[(i * page_size):(i * page_size + page_size)])
  81. if self.progress_callback is not None:
  82. if self._has_checksum:
  83. self.progress_callback(i + 1, load_count)
  84. else:
  85. self.progress_callback(i + 1, load_count*2)
  86. def verifyFlash(self, flash_data):
  87. if self._has_checksum:
  88. self.sendMessage([0x06, 0x00, (len(flash_data) >> 17) & 0xFF, (len(flash_data) >> 9) & 0xFF, (len(flash_data) >> 1) & 0xFF])
  89. res = self.sendMessage([0xEE])
  90. checksum_recv = res[2] | (res[3] << 8)
  91. checksum = 0
  92. for d in flash_data:
  93. checksum += d
  94. checksum &= 0xFFFF
  95. if hex(checksum) != hex(checksum_recv):
  96. raise ispBase.IspError("Verify checksum mismatch: 0x%x != 0x%x" % (checksum & 0xFFFF, checksum_recv))
  97. else:
  98. #Set load addr to 0, in case we have more then 64k flash we need to enable the address extension
  99. flash_size = self.chip["pageSize"] * 2 * self.chip["pageCount"]
  100. if flash_size > 0xFFFF:
  101. self.sendMessage([0x06, 0x80, 0x00, 0x00, 0x00])
  102. else:
  103. self.sendMessage([0x06, 0x00, 0x00, 0x00, 0x00])
  104. load_count = (len(flash_data) + 0xFF) / 0x100
  105. for i in range(0, int(load_count)):
  106. recv = self.sendMessage([0x14, 0x01, 0x00, 0x20])[2:0x102]
  107. if self.progress_callback is not None:
  108. self.progress_callback(load_count + i + 1, load_count*2)
  109. for j in range(0, 0x100):
  110. if i * 0x100 + j < len(flash_data) and flash_data[i * 0x100 + j] != recv[j]:
  111. raise ispBase.IspError("Verify error at: 0x%x" % (i * 0x100 + j))
  112. def sendMessage(self, data):
  113. message = struct.pack(">BBHB", 0x1B, self.seq, len(data), 0x0E)
  114. for c in data:
  115. message += struct.pack(">B", c)
  116. checksum = 0
  117. for c in message:
  118. checksum ^= c
  119. message += struct.pack(">B", checksum)
  120. try:
  121. self.serial.write(message)
  122. self.serial.flush()
  123. except SerialTimeoutException:
  124. raise ispBase.IspError("Serial send timeout")
  125. self.seq = (self.seq + 1) & 0xFF
  126. return self.recvMessage()
  127. def recvMessage(self):
  128. state = "Start"
  129. checksum = 0
  130. while True:
  131. s = self.serial.read()
  132. if len(s) < 1:
  133. raise ispBase.IspError("Timeout")
  134. b = struct.unpack(">B", s)[0]
  135. checksum ^= b
  136. #print(hex(b))
  137. if state == "Start":
  138. if b == 0x1B:
  139. state = "GetSeq"
  140. checksum = 0x1B
  141. elif state == "GetSeq":
  142. state = "MsgSize1"
  143. elif state == "MsgSize1":
  144. msg_size = b << 8
  145. state = "MsgSize2"
  146. elif state == "MsgSize2":
  147. msg_size |= b
  148. state = "Token"
  149. elif state == "Token":
  150. if b != 0x0E:
  151. state = "Start"
  152. else:
  153. state = "Data"
  154. data = []
  155. elif state == "Data":
  156. data.append(b)
  157. if len(data) == msg_size:
  158. state = "Checksum"
  159. elif state == "Checksum":
  160. if checksum != 0:
  161. state = "Start"
  162. else:
  163. return data
  164. def portList():
  165. ret = []
  166. import _winreg
  167. key=_winreg.OpenKey(_winreg.HKEY_LOCAL_MACHINE,"HARDWARE\\DEVICEMAP\\SERIALCOMM")
  168. i=0
  169. while True:
  170. try:
  171. values = _winreg.EnumValue(key, i)
  172. except:
  173. return ret
  174. if "USBSER" in values[0]:
  175. ret.append(values[1])
  176. i+=1
  177. return ret
  178. def runProgrammer(port, filename):
  179. """ Run an STK500v2 program on serial port 'port' and write 'filename' into flash. """
  180. programmer = Stk500v2()
  181. programmer.connect(port = port)
  182. programmer.programChip(intelHex.readHex(filename))
  183. programmer.close()
  184. def main():
  185. """ Entry point to call the stk500v2 programmer from the commandline. """
  186. import threading
  187. if sys.argv[1] == "AUTO":
  188. print(portList())
  189. for port in portList():
  190. threading.Thread(target=runProgrammer, args=(port,sys.argv[2])).start()
  191. time.sleep(5)
  192. else:
  193. programmer = Stk500v2()
  194. programmer.connect(port = sys.argv[1])
  195. programmer.programChip(intelHex.readHex(sys.argv[2]))
  196. sys.exit(1)
  197. if __name__ == "__main__":
  198. main()