crypto.go 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376
  1. package ike
  2. import (
  3. "crypto/aes"
  4. "crypto/cipher"
  5. "crypto/hmac"
  6. "crypto/rand"
  7. "crypto/sha1"
  8. "crypto/sha256"
  9. "crypto/subtle"
  10. "encoding/binary"
  11. "encoding/hex"
  12. "errors"
  13. "fmt"
  14. "hash"
  15. "io"
  16. "math/big"
  17. "net"
  18. )
  19. type ikeKeys struct {
  20. SKd []byte
  21. SKai []byte
  22. SKar []byte
  23. SKei []byte
  24. SKer []byte
  25. SKpi []byte
  26. SKpr []byte
  27. }
  28. func (suite negotiatedSuite) prf() (func() hash.Hash, int, error) {
  29. switch suite.PRFID {
  30. case prfHMACSHA1:
  31. return sha1.New, sha1.Size, nil
  32. case prfHMACSHA256:
  33. return sha256.New, sha256.Size, nil
  34. default:
  35. return nil, 0, fmt.Errorf("%w: PRF id %d", errUnsupportedSuite, suite.PRFID)
  36. }
  37. }
  38. func (suite negotiatedSuite) encryptionKeyLength() (int, error) {
  39. switch suite.EncryptionBits {
  40. case 128, 256:
  41. return suite.EncryptionBits / 8, nil
  42. default:
  43. return 0, fmt.Errorf("%w: AES key length %d", errUnsupportedSuite, suite.EncryptionBits)
  44. }
  45. }
  46. func (suite negotiatedSuite) integrityLengths() (keyLength int, checksumLength int, err error) {
  47. switch suite.IntegrityID {
  48. case integrityHMACSHA1_96:
  49. return sha1.Size, 12, nil
  50. case integrityHMACSHA256_128:
  51. return sha256.Size, 16, nil
  52. default:
  53. return 0, 0, fmt.Errorf("%w: integrity id %d", errUnsupportedSuite, suite.IntegrityID)
  54. }
  55. }
  56. func prf(suite negotiatedSuite, key, data []byte) ([]byte, error) {
  57. hashFactory, _, err := suite.prf()
  58. if err != nil {
  59. return nil, err
  60. }
  61. mac := hmac.New(hashFactory, key)
  62. _, _ = mac.Write(data)
  63. return mac.Sum(nil), nil
  64. }
  65. func prfPlus(suite negotiatedSuite, key, seed []byte, length int) ([]byte, error) {
  66. if length < 0 {
  67. return nil, errors.New("ike: negative key stream length")
  68. }
  69. result := make([]byte, 0, length)
  70. var previous []byte
  71. for counter := byte(1); len(result) < length; counter++ {
  72. if counter == 0 {
  73. return nil, errors.New("ike: PRF+ output is too long")
  74. }
  75. input := make([]byte, 0, len(previous)+len(seed)+1)
  76. input = append(input, previous...)
  77. input = append(input, seed...)
  78. input = append(input, counter)
  79. block, err := prf(suite, key, input)
  80. if err != nil {
  81. return nil, err
  82. }
  83. result = append(result, block...)
  84. previous = block
  85. }
  86. return result[:length], nil
  87. }
  88. func deriveIKEKeys(
  89. suite negotiatedSuite,
  90. sharedSecret []byte,
  91. initiatorNonce []byte,
  92. responderNonce []byte,
  93. initiatorSPI [8]byte,
  94. responderSPI [8]byte,
  95. ) (ikeKeys, error) {
  96. _, preferredLength, err := suite.prf()
  97. if err != nil {
  98. return ikeKeys{}, err
  99. }
  100. encryptionLength, err := suite.encryptionKeyLength()
  101. if err != nil {
  102. return ikeKeys{}, err
  103. }
  104. integrityLength, _, err := suite.integrityLengths()
  105. if err != nil {
  106. return ikeKeys{}, err
  107. }
  108. nonceKey := append(append([]byte(nil), initiatorNonce...), responderNonce...)
  109. skeyseed, err := prf(suite, nonceKey, sharedSecret)
  110. if err != nil {
  111. return ikeKeys{}, err
  112. }
  113. seed := make([]byte, 0, len(initiatorNonce)+len(responderNonce)+16)
  114. seed = append(seed, initiatorNonce...)
  115. seed = append(seed, responderNonce...)
  116. seed = append(seed, initiatorSPI[:]...)
  117. seed = append(seed, responderSPI[:]...)
  118. total := preferredLength + integrityLength*2 + encryptionLength*2 + preferredLength*2
  119. stream, err := prfPlus(suite, skeyseed, seed, total)
  120. if err != nil {
  121. return ikeKeys{}, err
  122. }
  123. take := func(length int) []byte {
  124. value := append([]byte(nil), stream[:length]...)
  125. stream = stream[length:]
  126. return value
  127. }
  128. return ikeKeys{
  129. SKd: take(preferredLength),
  130. SKai: take(integrityLength),
  131. SKar: take(integrityLength),
  132. SKei: take(encryptionLength),
  133. SKer: take(encryptionLength),
  134. SKpi: take(preferredLength),
  135. SKpr: take(preferredLength),
  136. }, nil
  137. }
  138. func integrityMAC(suite negotiatedSuite, key, packetWithoutChecksum []byte) ([]byte, error) {
  139. var hashFactory func() hash.Hash
  140. switch suite.IntegrityID {
  141. case integrityHMACSHA1_96:
  142. hashFactory = sha1.New
  143. case integrityHMACSHA256_128:
  144. hashFactory = sha256.New
  145. default:
  146. return nil, fmt.Errorf("%w: integrity id %d", errUnsupportedSuite, suite.IntegrityID)
  147. }
  148. _, checksumLength, err := suite.integrityLengths()
  149. if err != nil {
  150. return nil, err
  151. }
  152. mac := hmac.New(hashFactory, key)
  153. _, _ = mac.Write(packetWithoutChecksum)
  154. return mac.Sum(nil)[:checksumLength], nil
  155. }
  156. func encryptPayloads(
  157. header ikeHeader,
  158. inner []payload,
  159. suite negotiatedSuite,
  160. encryptionKey []byte,
  161. integrityKey []byte,
  162. random io.Reader,
  163. ) ([]byte, error) {
  164. if random == nil {
  165. random = rand.Reader
  166. }
  167. first, plaintext, err := marshalPayloadChain(inner)
  168. if err != nil {
  169. return nil, err
  170. }
  171. block, err := aes.NewCipher(encryptionKey)
  172. if err != nil {
  173. return nil, fmt.Errorf("ike: initialize AES: %w", err)
  174. }
  175. paddingLength := block.BlockSize() - (len(plaintext)+1)%block.BlockSize()
  176. if paddingLength == block.BlockSize() {
  177. paddingLength = 0
  178. }
  179. padding := make([]byte, paddingLength)
  180. if _, err := io.ReadFull(random, padding); err != nil {
  181. return nil, fmt.Errorf("ike: generate encrypted payload padding: %w", err)
  182. }
  183. plaintext = append(plaintext, padding...)
  184. plaintext = append(plaintext, byte(paddingLength))
  185. iv := make([]byte, block.BlockSize())
  186. if _, err := io.ReadFull(random, iv); err != nil {
  187. return nil, fmt.Errorf("ike: generate encrypted payload IV: %w", err)
  188. }
  189. ciphertext := make([]byte, len(plaintext))
  190. cipher.NewCBCEncrypter(block, iv).CryptBlocks(ciphertext, plaintext)
  191. _, checksumLength, err := suite.integrityLengths()
  192. if err != nil {
  193. return nil, err
  194. }
  195. skLength := 4 + len(iv) + len(ciphertext) + checksumLength
  196. if skLength > 65535 {
  197. return nil, errors.New("ike: encrypted payload exceeds 65535 bytes")
  198. }
  199. body := make([]byte, skLength)
  200. body[0] = first
  201. body[1] = 0
  202. binary.BigEndian.PutUint16(body[2:4], uint16(skLength))
  203. copy(body[4:], iv)
  204. copy(body[4+len(iv):], ciphertext)
  205. header.NextPayload = payloadEncrypted
  206. packet := header.marshal(body)
  207. checksum, err := integrityMAC(suite, integrityKey, packet[:len(packet)-checksumLength])
  208. if err != nil {
  209. return nil, err
  210. }
  211. copy(packet[len(packet)-checksumLength:], checksum)
  212. return packet, nil
  213. }
  214. func decryptPayloads(
  215. packet []byte,
  216. suite negotiatedSuite,
  217. encryptionKey []byte,
  218. integrityKey []byte,
  219. ) (ikeHeader, []payload, error) {
  220. header, body, err := parseIKEPacket(packet)
  221. if err != nil {
  222. return ikeHeader{}, nil, err
  223. }
  224. if header.NextPayload != payloadEncrypted || len(body) < 4 {
  225. return ikeHeader{}, nil, fmt.Errorf("%w: message is not an encrypted IKE payload", errUnexpectedPacket)
  226. }
  227. skLength := int(binary.BigEndian.Uint16(body[2:4]))
  228. if skLength != len(body) {
  229. return ikeHeader{}, nil, fmt.Errorf("%w: encrypted payload length mismatch", errMalformedPacket)
  230. }
  231. block, err := aes.NewCipher(encryptionKey)
  232. if err != nil {
  233. return ikeHeader{}, nil, fmt.Errorf("ike: initialize AES: %w", err)
  234. }
  235. _, checksumLength, err := suite.integrityLengths()
  236. if err != nil {
  237. return ikeHeader{}, nil, err
  238. }
  239. if len(body) < 4+block.BlockSize()+block.BlockSize()+checksumLength {
  240. return ikeHeader{}, nil, fmt.Errorf("%w: encrypted payload is too short", errMalformedPacket)
  241. }
  242. expected, err := integrityMAC(suite, integrityKey, packet[:len(packet)-checksumLength])
  243. if err != nil {
  244. return ikeHeader{}, nil, err
  245. }
  246. actual := packet[len(packet)-checksumLength:]
  247. if subtle.ConstantTimeCompare(actual, expected) != 1 {
  248. return ikeHeader{}, nil, errIntegrityMismatch
  249. }
  250. ivStart := 4
  251. ciphertextStart := ivStart + block.BlockSize()
  252. ciphertextEnd := len(body) - checksumLength
  253. ciphertext := body[ciphertextStart:ciphertextEnd]
  254. if len(ciphertext) == 0 || len(ciphertext)%block.BlockSize() != 0 {
  255. return ikeHeader{}, nil, fmt.Errorf("%w: ciphertext is not block aligned", errMalformedPacket)
  256. }
  257. plaintext := make([]byte, len(ciphertext))
  258. cipher.NewCBCDecrypter(block, body[ivStart:ciphertextStart]).CryptBlocks(plaintext, ciphertext)
  259. paddingLength := int(plaintext[len(plaintext)-1])
  260. if paddingLength+1 > len(plaintext) {
  261. return ikeHeader{}, nil, fmt.Errorf("%w: invalid encrypted payload padding", errMalformedPacket)
  262. }
  263. plaintext = plaintext[:len(plaintext)-paddingLength-1]
  264. payloads, err := parsePayloadChain(body[0], plaintext)
  265. if err != nil {
  266. return ikeHeader{}, nil, err
  267. }
  268. return header, payloads, nil
  269. }
  270. var modpPrimes = map[uint16]string{
  271. dhMODP1024: "FFFFFFFFFFFFFFFFC90FDAA22168C234C4C6628B80DC1CD1" +
  272. "29024E088A67CC74020BBEA63B139B22514A08798E3404DD" +
  273. "EF9519B3CD3A431B302B0A6DF25F14374FE1356D6D51C245" +
  274. "E485B576625E7EC6F44C42E9A637ED6B0BFF5CB6F406B7ED" +
  275. "EE386BFB5A899FA5AE9F24117C4B1FE649286651ECE65381" +
  276. "FFFFFFFFFFFFFFFF",
  277. dhMODP2048: "FFFFFFFFFFFFFFFFC90FDAA22168C234C4C6628B80DC1CD1" +
  278. "29024E088A67CC74020BBEA63B139B22514A08798E3404DD" +
  279. "EF9519B3CD3A431B302B0A6DF25F14374FE1356D6D51C245" +
  280. "E485B576625E7EC6F44C42E9A637ED6B0BFF5CB6F406B7ED" +
  281. "EE386BFB5A899FA5AE9F24117C4B1FE649286651ECE45B3D" +
  282. "C2007CB8A163BF0598DA48361C55D39A69163FA8FD24CF5F" +
  283. "83655D23DCA3AD961C62F356208552BB9ED529077096966D" +
  284. "670C354E4ABC9804F1746C08CA18217C32905E462E36CE3B" +
  285. "E39E772C180E86039B2783A2EC07A28FB5C55DF06F4C52C9" +
  286. "DE2BCBF6955817183995497CEA956AE515D2261898FA0510" +
  287. "15728E5A8AACAA68FFFFFFFFFFFFFFFF",
  288. }
  289. type dhExchange struct {
  290. Group uint16
  291. prime *big.Int
  292. private *big.Int
  293. Public []byte
  294. }
  295. func newDHExchange(group uint16, random io.Reader) (*dhExchange, error) {
  296. primeHex, ok := modpPrimes[group]
  297. if !ok {
  298. return nil, fmt.Errorf("%w: DH group %d", errUnsupportedSuite, group)
  299. }
  300. primeBytes, err := hex.DecodeString(primeHex)
  301. if err != nil {
  302. return nil, fmt.Errorf("ike: internal MODP constant: %w", err)
  303. }
  304. prime := new(big.Int).SetBytes(primeBytes)
  305. if random == nil {
  306. random = rand.Reader
  307. }
  308. sample := make([]byte, len(primeBytes))
  309. if _, err := io.ReadFull(random, sample); err != nil {
  310. return nil, fmt.Errorf("ike: generate DH private value: %w", err)
  311. }
  312. private := new(big.Int).SetBytes(sample)
  313. private.Mod(private, new(big.Int).Sub(prime, big.NewInt(3)))
  314. private.Add(private, big.NewInt(2))
  315. publicInteger := new(big.Int).Exp(big.NewInt(2), private, prime)
  316. public := publicInteger.FillBytes(make([]byte, len(primeBytes)))
  317. return &dhExchange{Group: group, prime: prime, private: private, Public: public}, nil
  318. }
  319. func (exchange *dhExchange) shared(peerPublic []byte) ([]byte, error) {
  320. if exchange == nil || exchange.prime == nil || exchange.private == nil {
  321. return nil, errors.New("ike: DH exchange is not initialized")
  322. }
  323. if len(peerPublic) != len(exchange.Public) {
  324. return nil, fmt.Errorf("ike: peer DH value length %d does not match group length %d", len(peerPublic), len(exchange.Public))
  325. }
  326. peer := new(big.Int).SetBytes(peerPublic)
  327. upper := new(big.Int).Sub(exchange.prime, big.NewInt(2))
  328. if peer.Cmp(big.NewInt(2)) < 0 || peer.Cmp(upper) > 0 {
  329. return nil, errors.New("ike: peer DH public value is outside the safe range")
  330. }
  331. shared := new(big.Int).Exp(peer, exchange.private, exchange.prime)
  332. if shared.Sign() == 0 || shared.Cmp(big.NewInt(1)) == 0 {
  333. return nil, errors.New("ike: invalid trivial DH shared secret")
  334. }
  335. return shared.FillBytes(make([]byte, len(exchange.Public))), nil
  336. }
  337. func natDetectionHash(
  338. initiatorSPI [8]byte,
  339. responderSPI [8]byte,
  340. ip net.IP,
  341. port uint16,
  342. ) ([]byte, error) {
  343. if ip4 := ip.To4(); ip4 != nil {
  344. ip = ip4
  345. } else if ip16 := ip.To16(); ip16 != nil {
  346. ip = ip16
  347. } else {
  348. return nil, errors.New("ike: NAT detection address is not an IP address")
  349. }
  350. input := make([]byte, 0, 16+len(ip)+2)
  351. input = append(input, initiatorSPI[:]...)
  352. input = append(input, responderSPI[:]...)
  353. input = append(input, ip...)
  354. var encodedPort [2]byte
  355. binary.BigEndian.PutUint16(encodedPort[:], port)
  356. input = append(input, encodedPort[:]...)
  357. sum := sha1.Sum(input)
  358. return sum[:], nil
  359. }