mirror of
https://github.com/MengMengCode/VoCat.git
synced 2026-08-20 14:53:42 +08:00
feat: implement IKE session relay and transport layer for ePDG communication
This commit is contained in:
@@ -280,6 +280,239 @@ func decryptPayloads(
|
||||
return header, payloads, nil
|
||||
}
|
||||
|
||||
const defaultIKEFragmentSize = 1100
|
||||
|
||||
func encryptPayloadsFragmented(
|
||||
header ikeHeader,
|
||||
inner []payload,
|
||||
suite negotiatedSuite,
|
||||
encryptionKey []byte,
|
||||
integrityKey []byte,
|
||||
maxFragmentSize int,
|
||||
random io.Reader,
|
||||
) ([][]byte, error) {
|
||||
if random == nil {
|
||||
random = rand.Reader
|
||||
}
|
||||
if maxFragmentSize <= 0 {
|
||||
maxFragmentSize = defaultIKEFragmentSize
|
||||
}
|
||||
first, plaintext, err := marshalPayloadChain(inner)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
block, err := aes.NewCipher(encryptionKey)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ike: initialize AES: %w", err)
|
||||
}
|
||||
_, checksumLength, err := suite.integrityLengths()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
maxChunk := maxFragmentSize - ikeHeaderLength - 8 - block.BlockSize() - block.BlockSize() - checksumLength
|
||||
if maxChunk < 64 {
|
||||
maxChunk = 64
|
||||
}
|
||||
|
||||
var chunks [][]byte
|
||||
for len(plaintext) > 0 {
|
||||
take := len(plaintext)
|
||||
if take > maxChunk {
|
||||
take = maxChunk
|
||||
}
|
||||
chunks = append(chunks, plaintext[:take])
|
||||
plaintext = plaintext[take:]
|
||||
}
|
||||
totalFragments := uint16(len(chunks))
|
||||
if totalFragments == 0 {
|
||||
totalFragments = 1
|
||||
chunks = [][]byte{nil}
|
||||
}
|
||||
|
||||
var packets [][]byte
|
||||
for index, chunk := range chunks {
|
||||
fragNum := uint16(index + 1)
|
||||
fragNext := uint8(payloadNone)
|
||||
if fragNum == 1 {
|
||||
fragNext = first
|
||||
}
|
||||
|
||||
paddingLength := block.BlockSize() - (len(chunk)+1)%block.BlockSize()
|
||||
if paddingLength == block.BlockSize() {
|
||||
paddingLength = 0
|
||||
}
|
||||
padding := make([]byte, paddingLength)
|
||||
if _, err := io.ReadFull(random, padding); err != nil {
|
||||
return nil, fmt.Errorf("ike: generate encrypted payload padding: %w", err)
|
||||
}
|
||||
paddedChunk := append(append([]byte(nil), chunk...), padding...)
|
||||
paddedChunk = append(paddedChunk, byte(paddingLength))
|
||||
|
||||
iv := make([]byte, block.BlockSize())
|
||||
if _, err := io.ReadFull(random, iv); err != nil {
|
||||
return nil, fmt.Errorf("ike: generate encrypted payload IV: %w", err)
|
||||
}
|
||||
ciphertext := make([]byte, len(paddedChunk))
|
||||
cipher.NewCBCEncrypter(block, iv).CryptBlocks(ciphertext, paddedChunk)
|
||||
|
||||
skfLength := 4 + 4 + len(iv) + len(ciphertext) + checksumLength
|
||||
if skfLength > 65535 {
|
||||
return nil, errors.New("ike: encrypted fragment exceeds 65535 bytes")
|
||||
}
|
||||
|
||||
body := make([]byte, skfLength)
|
||||
body[0] = fragNext
|
||||
body[1] = 0
|
||||
binary.BigEndian.PutUint16(body[2:4], uint16(skfLength))
|
||||
binary.BigEndian.PutUint16(body[4:6], fragNum)
|
||||
binary.BigEndian.PutUint16(body[6:8], totalFragments)
|
||||
copy(body[8:], iv)
|
||||
copy(body[8+len(iv):], ciphertext)
|
||||
|
||||
fragHeader := header
|
||||
fragHeader.NextPayload = payloadEncryptedFragment
|
||||
packet := fragHeader.marshal(body)
|
||||
checksum, err := integrityMAC(suite, integrityKey, packet[:len(packet)-checksumLength])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
copy(packet[len(packet)-checksumLength:], checksum)
|
||||
packets = append(packets, packet)
|
||||
}
|
||||
return packets, nil
|
||||
}
|
||||
|
||||
func decryptSingleFragment(
|
||||
packet []byte,
|
||||
suite negotiatedSuite,
|
||||
encryptionKey []byte,
|
||||
integrityKey []byte,
|
||||
) (ikeHeader, uint8, uint16, uint16, []byte, error) {
|
||||
header, body, err := parseIKEPacket(packet)
|
||||
if err != nil {
|
||||
return ikeHeader{}, 0, 0, 0, nil, err
|
||||
}
|
||||
if header.NextPayload != payloadEncryptedFragment || len(body) < 8 {
|
||||
return ikeHeader{}, 0, 0, 0, nil, fmt.Errorf("%w: message is not an encrypted IKE fragment", errUnexpectedPacket)
|
||||
}
|
||||
skfLength := int(binary.BigEndian.Uint16(body[2:4]))
|
||||
if skfLength != len(body) {
|
||||
return ikeHeader{}, 0, 0, 0, nil, fmt.Errorf("%w: encrypted fragment length mismatch", errMalformedPacket)
|
||||
}
|
||||
block, err := aes.NewCipher(encryptionKey)
|
||||
if err != nil {
|
||||
return ikeHeader{}, 0, 0, 0, nil, fmt.Errorf("ike: initialize AES: %w", err)
|
||||
}
|
||||
_, checksumLength, err := suite.integrityLengths()
|
||||
if err != nil {
|
||||
return ikeHeader{}, 0, 0, 0, nil, err
|
||||
}
|
||||
if len(body) < 8+block.BlockSize()+block.BlockSize()+checksumLength {
|
||||
return ikeHeader{}, 0, 0, 0, nil, fmt.Errorf("%w: encrypted fragment is too short", errMalformedPacket)
|
||||
}
|
||||
expected, err := integrityMAC(suite, integrityKey, packet[:len(packet)-checksumLength])
|
||||
if err != nil {
|
||||
return ikeHeader{}, 0, 0, 0, nil, err
|
||||
}
|
||||
actual := packet[len(packet)-checksumLength:]
|
||||
if subtle.ConstantTimeCompare(actual, expected) != 1 {
|
||||
return ikeHeader{}, 0, 0, 0, nil, errIntegrityMismatch
|
||||
}
|
||||
|
||||
fragNext := body[0]
|
||||
fragNum := binary.BigEndian.Uint16(body[4:6])
|
||||
totalFrags := binary.BigEndian.Uint16(body[6:8])
|
||||
if fragNum == 0 || totalFrags == 0 || fragNum > totalFrags {
|
||||
return ikeHeader{}, 0, 0, 0, nil, fmt.Errorf("%w: invalid fragment numbers %d/%d", errMalformedPacket, fragNum, totalFrags)
|
||||
}
|
||||
|
||||
ivStart := 8
|
||||
ciphertextStart := ivStart + block.BlockSize()
|
||||
ciphertextEnd := len(body) - checksumLength
|
||||
ciphertext := body[ciphertextStart:ciphertextEnd]
|
||||
if len(ciphertext) == 0 || len(ciphertext)%block.BlockSize() != 0 {
|
||||
return ikeHeader{}, 0, 0, 0, nil, fmt.Errorf("%w: fragment ciphertext is not block aligned", errMalformedPacket)
|
||||
}
|
||||
plaintext := make([]byte, len(ciphertext))
|
||||
cipher.NewCBCDecrypter(block, body[ivStart:ciphertextStart]).CryptBlocks(plaintext, ciphertext)
|
||||
paddingLength := int(plaintext[len(plaintext)-1])
|
||||
if paddingLength+1 > len(plaintext) {
|
||||
return ikeHeader{}, 0, 0, 0, nil, fmt.Errorf("%w: invalid encrypted fragment padding", errMalformedPacket)
|
||||
}
|
||||
plaintext = plaintext[:len(plaintext)-paddingLength-1]
|
||||
return header, fragNext, fragNum, totalFrags, plaintext, nil
|
||||
}
|
||||
|
||||
func decryptPayloadsAny(
|
||||
packet []byte,
|
||||
fragments [][]byte,
|
||||
suite negotiatedSuite,
|
||||
encryptionKey []byte,
|
||||
integrityKey []byte,
|
||||
) (ikeHeader, []payload, error) {
|
||||
if len(fragments) > 0 {
|
||||
var (
|
||||
firstHeader ikeHeader
|
||||
firstNext uint8
|
||||
totalExpected uint16
|
||||
plaintexts = make(map[uint16][]byte)
|
||||
)
|
||||
for _, fragPacket := range fragments {
|
||||
hdr, next, num, total, plain, err := decryptSingleFragment(fragPacket, suite, encryptionKey, integrityKey)
|
||||
if err != nil {
|
||||
return ikeHeader{}, nil, err
|
||||
}
|
||||
if totalExpected == 0 {
|
||||
firstHeader = hdr
|
||||
totalExpected = total
|
||||
} else if total != totalExpected || hdr.MessageID != firstHeader.MessageID || hdr.Exchange != firstHeader.Exchange {
|
||||
return ikeHeader{}, nil, fmt.Errorf("%w: inconsistent fragment headers", errMalformedPacket)
|
||||
}
|
||||
if num == 1 {
|
||||
firstNext = next
|
||||
}
|
||||
plaintexts[num] = plain
|
||||
}
|
||||
if uint16(len(plaintexts)) != totalExpected {
|
||||
return ikeHeader{}, nil, fmt.Errorf("%w: missing fragments: received %d of %d", errMalformedPacket, len(plaintexts), totalExpected)
|
||||
}
|
||||
var fullPlaintext []byte
|
||||
for i := uint16(1); i <= totalExpected; i++ {
|
||||
chunk, ok := plaintexts[i]
|
||||
if !ok {
|
||||
return ikeHeader{}, nil, fmt.Errorf("%w: missing fragment %d", errMalformedPacket, i)
|
||||
}
|
||||
fullPlaintext = append(fullPlaintext, chunk...)
|
||||
}
|
||||
payloads, err := parsePayloadChain(firstNext, fullPlaintext)
|
||||
if err != nil {
|
||||
return ikeHeader{}, nil, err
|
||||
}
|
||||
return firstHeader, payloads, nil
|
||||
}
|
||||
|
||||
header, _, err := parseIKEPacket(packet)
|
||||
if err != nil {
|
||||
return ikeHeader{}, nil, err
|
||||
}
|
||||
if header.NextPayload == payloadEncryptedFragment {
|
||||
hdr, next, num, total, plain, err := decryptSingleFragment(packet, suite, encryptionKey, integrityKey)
|
||||
if err != nil {
|
||||
return ikeHeader{}, nil, err
|
||||
}
|
||||
if num != 1 || total != 1 {
|
||||
return ikeHeader{}, nil, fmt.Errorf("%w: standalone fragment with total=%d", errMalformedPacket, total)
|
||||
}
|
||||
payloads, err := parsePayloadChain(next, plain)
|
||||
if err != nil {
|
||||
return ikeHeader{}, nil, err
|
||||
}
|
||||
return hdr, payloads, nil
|
||||
}
|
||||
return decryptPayloads(packet, suite, encryptionKey, integrityKey)
|
||||
}
|
||||
|
||||
var modpPrimes = map[uint16]string{
|
||||
dhMODP1024: "FFFFFFFFFFFFFFFFC90FDAA22168C234C4C6628B80DC1CD1" +
|
||||
"29024E088A67CC74020BBEA63B139B22514A08798E3404DD" +
|
||||
|
||||
@@ -113,3 +113,80 @@ func TestIKEKeyDerivationSeparatesDirections(t *testing.T) {
|
||||
t.Fatal("initiator and responder keys were not separated")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRFC7383FragmentationAndReassembly(t *testing.T) {
|
||||
suite := legacyTestSuite()
|
||||
encryptionKey := bytes.Repeat([]byte{0x11}, 16)
|
||||
integrityKey := bytes.Repeat([]byte{0x22}, 20)
|
||||
header := ikeHeader{
|
||||
InitiatorSPI: [8]byte{1, 2, 3, 4, 5, 6, 7, 8},
|
||||
ResponderSPI: [8]byte{8, 7, 6, 5, 4, 3, 2, 1},
|
||||
Exchange: exchangeIKEAuth,
|
||||
Flags: flagInitiator,
|
||||
MessageID: 9,
|
||||
}
|
||||
|
||||
largeCertData := bytes.Repeat([]byte{0xAB, 0xCD, 0xEF, 0x01}, 400) // 1600 bytes
|
||||
inner := []payload{
|
||||
{Type: payloadIDi, Body: []byte{3, 0, 0, 0, 'u', 's', 'e', 'r'}},
|
||||
{Type: payloadCert, Body: largeCertData},
|
||||
{Type: payloadAuth, Body: bytes.Repeat([]byte{0x55}, 64)},
|
||||
}
|
||||
|
||||
// Fragment into chunks with max fragment size 600 bytes
|
||||
packets, err := encryptPayloadsFragmented(
|
||||
header,
|
||||
inner,
|
||||
suite,
|
||||
encryptionKey,
|
||||
integrityKey,
|
||||
600,
|
||||
bytes.NewReader(bytes.Repeat([]byte{0x77}, 1024)),
|
||||
)
|
||||
if err != nil {
|
||||
t.Fatalf("encryptPayloadsFragmented() error = %v", err)
|
||||
}
|
||||
|
||||
if len(packets) < 3 {
|
||||
t.Fatalf("expected at least 3 fragments for large payload, got %d", len(packets))
|
||||
}
|
||||
|
||||
for i, pkt := range packets {
|
||||
hdr, body, parseErr := parseIKEPacket(pkt)
|
||||
if parseErr != nil {
|
||||
t.Fatalf("fragment %d parse error: %v", i+1, parseErr)
|
||||
}
|
||||
if hdr.NextPayload != payloadEncryptedFragment {
|
||||
t.Fatalf("fragment %d NextPayload = %d, want %d (payloadEncryptedFragment)", i+1, hdr.NextPayload, payloadEncryptedFragment)
|
||||
}
|
||||
if len(body) < 8 {
|
||||
t.Fatalf("fragment %d body too short", i+1)
|
||||
}
|
||||
}
|
||||
|
||||
// Decrypt and reassemble
|
||||
decodedHeader, decoded, err := decryptPayloadsAny(nil, packets, suite, encryptionKey, integrityKey)
|
||||
if err != nil {
|
||||
t.Fatalf("decryptPayloadsAny() error = %v", err)
|
||||
}
|
||||
|
||||
if decodedHeader.MessageID != header.MessageID || len(decoded) != len(inner) {
|
||||
t.Fatalf("reassembled payload mismatch: header=%#v, count=%d, want=%d", decodedHeader, len(decoded), len(inner))
|
||||
}
|
||||
|
||||
for index := range inner {
|
||||
if decoded[index].Type != inner[index].Type || !bytes.Equal(decoded[index].Body, inner[index].Body) {
|
||||
t.Fatalf("decoded payload %d = %#v, want %#v", index, decoded[index], inner[index])
|
||||
}
|
||||
}
|
||||
|
||||
// Test tamper detection on second fragment
|
||||
tamperedPackets := make([][]byte, len(packets))
|
||||
for i := range packets {
|
||||
tamperedPackets[i] = append([]byte(nil), packets[i]...)
|
||||
}
|
||||
tamperedPackets[1][len(tamperedPackets[1])-1] ^= 0x55
|
||||
if _, _, err := decryptPayloadsAny(nil, tamperedPackets, suite, encryptionKey, integrityKey); !errors.Is(err, errIntegrityMismatch) {
|
||||
t.Fatalf("tampered fragment decrypt error = %v, want errIntegrityMismatch", err)
|
||||
}
|
||||
}
|
||||
|
||||
+103
-32
@@ -187,6 +187,7 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
|
||||
{Type: payloadNonce, Body: initiatorNonce},
|
||||
makeNotify(notifyNATSource, sourceHash),
|
||||
makeNotify(notifyNATDestination, destinationHash),
|
||||
makeNotify(notifyFragmentationSupported, nil),
|
||||
}
|
||||
var (
|
||||
initRequest []byte
|
||||
@@ -243,6 +244,7 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
|
||||
}
|
||||
break
|
||||
}
|
||||
peerSupportsFragmentation := hasNotifyType(initResponsePayloads, notifyFragmentationSupported)
|
||||
saPayload, err := onePayload(initResponsePayloads, payloadSA)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -345,21 +347,19 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
|
||||
Flags: flagInitiator,
|
||||
MessageID: 1,
|
||||
}
|
||||
authRequest, err := encryptPayloads(authHeader, firstAuthPayloads, ikeSuite, keys.SKei, keys.SKai, provider.config.Random)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
authResponse, err := transport.RoundTrip(ctx, authRequest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
authResponseHeader, authResponsePayloads, err := decryptAndValidate(
|
||||
authResponse, initiatorSPI, responseHeader.ResponderSPI, exchangeIKEAuth, 1, ikeSuite, keys,
|
||||
_, authResponsePayloads, err := sendAndReceiveIKEPayloads(
|
||||
ctx,
|
||||
transport,
|
||||
authHeader,
|
||||
firstAuthPayloads,
|
||||
ikeSuite,
|
||||
keys,
|
||||
peerSupportsFragmentation,
|
||||
provider.config.Random,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_ = authResponseHeader
|
||||
serverName := strings.TrimSpace(provider.config.ServerName)
|
||||
if serverName == "" {
|
||||
serverName = epdg
|
||||
@@ -409,27 +409,27 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
|
||||
requestPayloads = append(requestPayloads, deviceIdentity)
|
||||
}
|
||||
}
|
||||
eapRequest, err := encryptPayloads(ikeHeader{
|
||||
eapHeader := ikeHeader{
|
||||
InitiatorSPI: initiatorSPI,
|
||||
ResponderSPI: responseHeader.ResponderSPI,
|
||||
Exchange: exchangeIKEAuth,
|
||||
Flags: flagInitiator,
|
||||
MessageID: messageID,
|
||||
}, requestPayloads, ikeSuite, keys.SKei, keys.SKai, provider.config.Random)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if requested, notifyErr := deviceIdentityRequested(currentPayloads); notifyErr != nil {
|
||||
return nil, notifyErr
|
||||
} else if requested {
|
||||
deviceIdentityPending = true
|
||||
}
|
||||
eapResponse, err := transport.RoundTrip(ctx, eapRequest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_, currentPayloads, err = decryptAndValidate(
|
||||
eapResponse, initiatorSPI, responseHeader.ResponderSPI, exchangeIKEAuth, messageID, ikeSuite, keys,
|
||||
_, currentPayloads, err = sendAndReceiveIKEPayloads(
|
||||
ctx,
|
||||
transport,
|
||||
eapHeader,
|
||||
requestPayloads,
|
||||
ikeSuite,
|
||||
keys,
|
||||
peerSupportsFragmentation,
|
||||
provider.config.Random,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -454,22 +454,22 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
|
||||
}
|
||||
messageID++
|
||||
cleanupMessageID = messageID + 1
|
||||
finalRequest, err := encryptPayloads(ikeHeader{
|
||||
finalHeader := ikeHeader{
|
||||
InitiatorSPI: initiatorSPI,
|
||||
ResponderSPI: responseHeader.ResponderSPI,
|
||||
Exchange: exchangeIKEAuth,
|
||||
Flags: flagInitiator,
|
||||
MessageID: messageID,
|
||||
}, []payload{initiatorAUTH}, ikeSuite, keys.SKei, keys.SKai, provider.config.Random)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
finalResponse, err := transport.RoundTrip(ctx, finalRequest)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
_, finalPayloads, err := decryptAndValidate(
|
||||
finalResponse, initiatorSPI, responseHeader.ResponderSPI, exchangeIKEAuth, messageID, ikeSuite, keys,
|
||||
_, finalPayloads, err := sendAndReceiveIKEPayloads(
|
||||
ctx,
|
||||
transport,
|
||||
finalHeader,
|
||||
[]payload{initiatorAUTH},
|
||||
ikeSuite,
|
||||
keys,
|
||||
peerSupportsFragmentation,
|
||||
provider.config.Random,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
@@ -763,7 +763,7 @@ func decryptAndValidate(
|
||||
suite negotiatedSuite,
|
||||
keys ikeKeys,
|
||||
) (ikeHeader, []payload, error) {
|
||||
header, payloads, err := decryptPayloads(packet, suite, keys.SKer, keys.SKar)
|
||||
header, payloads, err := decryptPayloadsAny(packet, nil, suite, keys.SKer, keys.SKar)
|
||||
if err != nil {
|
||||
return ikeHeader{}, nil, err
|
||||
}
|
||||
@@ -778,6 +778,77 @@ func decryptAndValidate(
|
||||
return header, payloads, nil
|
||||
}
|
||||
|
||||
func decryptAndValidateFragments(
|
||||
packets [][]byte,
|
||||
initiatorSPI [8]byte,
|
||||
responderSPI [8]byte,
|
||||
exchange uint8,
|
||||
messageID uint32,
|
||||
suite negotiatedSuite,
|
||||
keys ikeKeys,
|
||||
) (ikeHeader, []payload, error) {
|
||||
if len(packets) == 0 {
|
||||
return ikeHeader{}, nil, errors.New("ike: empty exchange response")
|
||||
}
|
||||
if len(packets) == 1 {
|
||||
return decryptAndValidate(packets[0], initiatorSPI, responderSPI, exchange, messageID, suite, keys)
|
||||
}
|
||||
header, payloads, err := decryptPayloadsAny(nil, packets, suite, keys.SKer, keys.SKar)
|
||||
if err != nil {
|
||||
return ikeHeader{}, nil, err
|
||||
}
|
||||
if header.InitiatorSPI != initiatorSPI ||
|
||||
header.ResponderSPI != responderSPI ||
|
||||
header.Exchange != exchange ||
|
||||
header.MessageID != messageID ||
|
||||
header.Flags&flagResponse == 0 ||
|
||||
header.Flags&flagInitiator != 0 {
|
||||
return ikeHeader{}, nil, fmt.Errorf("%w: encrypted response header does not match the request", errUnexpectedPacket)
|
||||
}
|
||||
return header, payloads, nil
|
||||
}
|
||||
|
||||
func sendAndReceiveIKEPayloads(
|
||||
ctx context.Context,
|
||||
transport datagramTransport,
|
||||
header ikeHeader,
|
||||
payloads []payload,
|
||||
suite negotiatedSuite,
|
||||
keys ikeKeys,
|
||||
peerSupportsFragmentation bool,
|
||||
random io.Reader,
|
||||
) (ikeHeader, []payload, error) {
|
||||
var outboundPackets [][]byte
|
||||
var err error
|
||||
if peerSupportsFragmentation {
|
||||
outboundPackets, err = encryptPayloadsFragmented(header, payloads, suite, keys.SKei, keys.SKai, defaultIKEFragmentSize, random)
|
||||
} else {
|
||||
pkt, encryptErr := encryptPayloads(header, payloads, suite, keys.SKei, keys.SKai, random)
|
||||
if encryptErr != nil {
|
||||
return ikeHeader{}, nil, encryptErr
|
||||
}
|
||||
outboundPackets = [][]byte{pkt}
|
||||
}
|
||||
if err != nil {
|
||||
return ikeHeader{}, nil, err
|
||||
}
|
||||
inboundPackets, err := transport.RoundTripExchange(ctx, outboundPackets)
|
||||
if err != nil {
|
||||
return ikeHeader{}, nil, err
|
||||
}
|
||||
return decryptAndValidateFragments(inboundPackets, header.InitiatorSPI, header.ResponderSPI, header.Exchange, header.MessageID, suite, keys)
|
||||
}
|
||||
|
||||
func hasNotifyType(payloads []payload, notifyType uint16) bool {
|
||||
for _, item := range payloadsOfType(payloads, payloadNotify) {
|
||||
kind, _, err := parseNotify(item)
|
||||
if err == nil && kind == notifyType {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
var errNoProposalChosen = errors.New("ike: responder reported NO_PROPOSAL_CHOSEN")
|
||||
|
||||
type invalidKEPayloadError struct {
|
||||
|
||||
@@ -76,7 +76,7 @@ func (transport *firstAuthCaptureTransport) Float(context.Context) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (transport *firstAuthCaptureTransport) RoundTrip(_ context.Context, packet []byte) ([]byte, error) {
|
||||
func (transport *firstAuthCaptureTransport) RoundTrip(ctx context.Context, packet []byte) ([]byte, error) {
|
||||
transport.calls++
|
||||
if len(transport.cookieChallenge) > 0 {
|
||||
switch transport.calls {
|
||||
@@ -106,6 +106,17 @@ func (transport *firstAuthCaptureTransport) RoundTrip(_ context.Context, packet
|
||||
}
|
||||
}
|
||||
|
||||
func (transport *firstAuthCaptureTransport) RoundTripExchange(ctx context.Context, packets [][]byte) ([][]byte, error) {
|
||||
if len(packets) == 0 {
|
||||
return nil, errors.New("test: empty outbound packets")
|
||||
}
|
||||
resp, err := transport.RoundTrip(ctx, packets[0])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return [][]byte{resp}, nil
|
||||
}
|
||||
|
||||
func (transport *firstAuthCaptureTransport) answerIKECookie(packet []byte) ([]byte, error) {
|
||||
header, _, err := parseIKEPacket(packet)
|
||||
if err != nil {
|
||||
@@ -146,8 +157,8 @@ func (transport *firstAuthCaptureTransport) verifyIKECookie(packet []byte) error
|
||||
return errors.New("test: first retried IKE_SA_INIT payload is not the expected COOKIE")
|
||||
}
|
||||
cookies := payloadsOfType(payloads, payloadNotify)
|
||||
if len(cookies) != 3 {
|
||||
return fmt.Errorf("test: retried IKE_SA_INIT has %d notify payloads, want 3", len(cookies))
|
||||
if len(cookies) != 4 {
|
||||
return fmt.Errorf("test: retried IKE_SA_INIT has %d notify payloads, want 4", len(cookies))
|
||||
}
|
||||
found := false
|
||||
for _, item := range cookies {
|
||||
|
||||
@@ -39,7 +39,7 @@ func newSessionRelay(
|
||||
keepalive time.Duration,
|
||||
) *sessionRelay {
|
||||
if keepalive <= 0 {
|
||||
keepalive = 20 * time.Second
|
||||
keepalive = 15 * time.Second
|
||||
}
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
relay := &sessionRelay{
|
||||
|
||||
@@ -3,6 +3,7 @@ package ike
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"errors"
|
||||
"net"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
@@ -49,6 +50,16 @@ func (*fakeSessionTransport) Float(context.Context) error { return nil }
|
||||
func (*fakeSessionTransport) RoundTrip(context.Context, []byte) ([]byte, error) {
|
||||
return nil, context.DeadlineExceeded
|
||||
}
|
||||
func (t *fakeSessionTransport) RoundTripExchange(ctx context.Context, packets [][]byte) ([][]byte, error) {
|
||||
if len(packets) == 0 {
|
||||
return nil, errors.New("empty outbound packets")
|
||||
}
|
||||
resp, err := t.RoundTrip(ctx, packets[0])
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return [][]byte{resp}, nil
|
||||
}
|
||||
func (transport *fakeSessionTransport) SendESP(ctx context.Context, packet []byte) error {
|
||||
return transport.SendSessionPacket(ctx, packet, false)
|
||||
}
|
||||
|
||||
@@ -20,6 +20,7 @@ type datagramTransport interface {
|
||||
RemoteAddr() *net.UDPAddr
|
||||
Float(context.Context) error
|
||||
RoundTrip(context.Context, []byte) ([]byte, error)
|
||||
RoundTripExchange(context.Context, [][]byte) ([][]byte, error)
|
||||
SendESP(context.Context, []byte) error
|
||||
ReceiveESP(context.Context, []byte) (int, error)
|
||||
SendSessionPacket(context.Context, []byte, bool) error
|
||||
@@ -96,6 +97,29 @@ func roundTripDatagram(
|
||||
read func([]byte, time.Time) (int, error),
|
||||
packet []byte,
|
||||
) ([]byte, error) {
|
||||
writeAll := func(values [][]byte) error {
|
||||
if len(values) > 0 {
|
||||
return write(values[0])
|
||||
}
|
||||
return nil
|
||||
}
|
||||
responses, err := roundTripFragments(ctx, timeout, writeAll, read, [][]byte{packet})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(responses) == 0 {
|
||||
return nil, errors.New("ike: empty datagram response")
|
||||
}
|
||||
return responses[0], nil
|
||||
}
|
||||
|
||||
func roundTripFragments(
|
||||
ctx context.Context,
|
||||
timeout time.Duration,
|
||||
writeAll func([][]byte) error,
|
||||
read func([]byte, time.Time) (int, error),
|
||||
packets [][]byte,
|
||||
) ([][]byte, error) {
|
||||
if ctx == nil {
|
||||
ctx = context.Background()
|
||||
}
|
||||
@@ -110,20 +134,44 @@ func roundTripDatagram(
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := write(packet); err != nil {
|
||||
if err := writeAll(packets); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
attemptDeadline := time.Now().Add(interval)
|
||||
if deadline.Before(attemptDeadline) {
|
||||
attemptDeadline = deadline
|
||||
}
|
||||
var (
|
||||
totalExpected uint16
|
||||
fragments = make(map[uint16][]byte)
|
||||
)
|
||||
for time.Now().Before(attemptDeadline) {
|
||||
if err := ctx.Err(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
n, err := read(buffer, attemptDeadline)
|
||||
if err == nil {
|
||||
return append([]byte(nil), buffer[:n]...), nil
|
||||
pkt := append([]byte(nil), buffer[:n]...)
|
||||
header, body, parseErr := parseIKEPacket(pkt)
|
||||
if parseErr == nil && header.NextPayload == payloadEncryptedFragment && len(body) >= 8 {
|
||||
fragNum := binary.BigEndian.Uint16(body[4:6])
|
||||
total := binary.BigEndian.Uint16(body[6:8])
|
||||
if total > 1 {
|
||||
if totalExpected == 0 {
|
||||
totalExpected = total
|
||||
}
|
||||
fragments[fragNum] = pkt
|
||||
if uint16(len(fragments)) == totalExpected {
|
||||
res := make([][]byte, 0, totalExpected)
|
||||
for i := uint16(1); i <= totalExpected; i++ {
|
||||
res = append(res, fragments[i])
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
return [][]byte{pkt}, nil
|
||||
}
|
||||
if timeoutError, ok := err.(net.Error); ok && timeoutError.Timeout() {
|
||||
lastErr = err
|
||||
@@ -222,25 +270,47 @@ func (transport *directUDP) Float(ctx context.Context) error {
|
||||
}
|
||||
|
||||
func (transport *directUDP) RoundTrip(ctx context.Context, packet []byte) ([]byte, error) {
|
||||
responses, err := transport.RoundTripExchange(ctx, [][]byte{packet})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(responses) == 0 {
|
||||
return nil, errors.New("ike: empty exchange response")
|
||||
}
|
||||
return responses[0], nil
|
||||
}
|
||||
|
||||
func (transport *directUDP) RoundTripExchange(ctx context.Context, packets [][]byte) ([][]byte, error) {
|
||||
transport.mu.Lock()
|
||||
defer transport.mu.Unlock()
|
||||
if transport.conn == nil {
|
||||
return nil, errors.New("ike: UDP transport is closed")
|
||||
}
|
||||
requestHeader, _, err := parseIKEPacket(packet)
|
||||
if len(packets) == 0 {
|
||||
return nil, errors.New("ike: outbound packet list is empty")
|
||||
}
|
||||
requestHeader, _, err := parseIKEPacket(packets[0])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ike: invalid outbound packet: %w", err)
|
||||
}
|
||||
wirePacket := packet
|
||||
if transport.floated {
|
||||
wirePacket = append([]byte{0, 0, 0, 0}, packet...)
|
||||
}
|
||||
write := func(value []byte) error {
|
||||
if err := transport.conn.SetWriteDeadline(deadlineFor(ctx, transport.config.Timeout)); err != nil {
|
||||
return err
|
||||
var wirePackets [][]byte
|
||||
for _, pkt := range packets {
|
||||
wire := pkt
|
||||
if transport.floated {
|
||||
wire = append([]byte{0, 0, 0, 0}, pkt...)
|
||||
}
|
||||
_, err := transport.conn.Write(value)
|
||||
return err
|
||||
wirePackets = append(wirePackets, wire)
|
||||
}
|
||||
writeAll := func(values [][]byte) error {
|
||||
for _, value := range values {
|
||||
if err := transport.conn.SetWriteDeadline(deadlineFor(ctx, transport.config.Timeout)); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := transport.conn.Write(value); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
read := func(buffer []byte, attemptDeadline time.Time) (int, error) {
|
||||
for {
|
||||
@@ -252,10 +322,6 @@ func (transport *directUDP) RoundTrip(ctx context.Context, packet []byte) ([]byt
|
||||
return 0, err
|
||||
}
|
||||
if transport.floated {
|
||||
// IKE and ESP legitimately share UDP/4500. An ESP packet can
|
||||
// arrive immediately before the IKE response that completes
|
||||
// CHILD_SA setup; discard it here and keep the same absolute
|
||||
// attempt deadline while waiting for marked IKE.
|
||||
if !hasNonESPMarker(buffer[:n]) {
|
||||
continue
|
||||
}
|
||||
@@ -268,7 +334,7 @@ func (transport *directUDP) RoundTrip(ctx context.Context, packet []byte) ([]byt
|
||||
return n, nil
|
||||
}
|
||||
}
|
||||
return roundTripDatagram(ctx, transport.config.Timeout, write, read, wirePacket)
|
||||
return roundTripFragments(ctx, transport.config.Timeout, writeAll, read, wirePackets)
|
||||
}
|
||||
|
||||
func (transport *directUDP) SendESP(ctx context.Context, packet []byte) error {
|
||||
@@ -561,12 +627,26 @@ func (transport *socks5UDP) Float(_ context.Context) error {
|
||||
}
|
||||
|
||||
func (transport *socks5UDP) RoundTrip(ctx context.Context, packet []byte) ([]byte, error) {
|
||||
responses, err := transport.RoundTripExchange(ctx, [][]byte{packet})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(responses) == 0 {
|
||||
return nil, errors.New("ike: empty exchange response")
|
||||
}
|
||||
return responses[0], nil
|
||||
}
|
||||
|
||||
func (transport *socks5UDP) RoundTripExchange(ctx context.Context, packets [][]byte) ([][]byte, error) {
|
||||
transport.mu.Lock()
|
||||
defer transport.mu.Unlock()
|
||||
if transport.udp == nil {
|
||||
return nil, errors.New("ike: SOCKS5 UDP transport is closed")
|
||||
}
|
||||
requestHeader, _, err := parseIKEPacket(packet)
|
||||
if len(packets) == 0 {
|
||||
return nil, errors.New("ike: outbound packet list is empty")
|
||||
}
|
||||
requestHeader, _, err := parseIKEPacket(packets[0])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ike: invalid outbound packet: %w", err)
|
||||
}
|
||||
@@ -576,18 +656,18 @@ func (transport *socks5UDP) RoundTrip(ctx context.Context, packet []byte) ([]byt
|
||||
// Once a gateway answers, keep it pinned for the lifetime of the IKE SA.
|
||||
if !transport.floated && requestHeader.Exchange == exchangeIKEInit && requestHeader.MessageID == 0 && len(transport.remotes) > 1 {
|
||||
var lastErr error
|
||||
var cookieResponse []byte
|
||||
var cookieResponse [][]byte
|
||||
for _, candidate := range transport.remotes {
|
||||
transport.remote = cloneUDPAddr(candidate)
|
||||
response, attemptErr := transport.roundTripLocked(ctx, packet, requestHeader)
|
||||
responses, attemptErr := transport.roundTripFragmentsLocked(ctx, packets, requestHeader)
|
||||
if attemptErr == nil {
|
||||
if ikeInitResponseHasCookie(response) {
|
||||
if len(responses) > 0 && ikeInitResponseHasCookie(responses[0]) {
|
||||
if cookieResponse == nil {
|
||||
cookieResponse = append([]byte(nil), response...)
|
||||
cookieResponse = responses
|
||||
}
|
||||
continue
|
||||
}
|
||||
return response, nil
|
||||
return responses, nil
|
||||
}
|
||||
lastErr = attemptErr
|
||||
if ctx.Err() != nil || !isNetworkTimeout(attemptErr) {
|
||||
@@ -599,24 +679,32 @@ func (transport *socks5UDP) RoundTrip(ctx context.Context, packet []byte) ([]byt
|
||||
}
|
||||
return nil, fmt.Errorf("ike: all %d resolved ePDG addresses timed out: %w", len(transport.remotes), lastErr)
|
||||
}
|
||||
return transport.roundTripLocked(ctx, packet, requestHeader)
|
||||
return transport.roundTripFragmentsLocked(ctx, packets, requestHeader)
|
||||
}
|
||||
|
||||
func (transport *socks5UDP) roundTripLocked(ctx context.Context, packet []byte, requestHeader ikeHeader) ([]byte, error) {
|
||||
wireIKE := packet
|
||||
if transport.floated {
|
||||
wireIKE = append([]byte{0, 0, 0, 0}, packet...)
|
||||
}
|
||||
datagram, err := marshalSOCKS5Datagram(transport.remote, wireIKE)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
write := func(value []byte) error {
|
||||
if err := transport.udp.SetWriteDeadline(deadlineFor(ctx, transport.config.Timeout)); err != nil {
|
||||
return err
|
||||
func (transport *socks5UDP) roundTripFragmentsLocked(ctx context.Context, packets [][]byte, requestHeader ikeHeader) ([][]byte, error) {
|
||||
var datagrams [][]byte
|
||||
for _, pkt := range packets {
|
||||
wireIKE := pkt
|
||||
if transport.floated {
|
||||
wireIKE = append([]byte{0, 0, 0, 0}, pkt...)
|
||||
}
|
||||
_, err := transport.udp.Write(value)
|
||||
return err
|
||||
datagram, err := marshalSOCKS5Datagram(transport.remote, wireIKE)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
datagrams = append(datagrams, datagram)
|
||||
}
|
||||
writeAll := func(values [][]byte) error {
|
||||
for _, value := range values {
|
||||
if err := transport.udp.SetWriteDeadline(deadlineFor(ctx, transport.config.Timeout)); err != nil {
|
||||
return err
|
||||
}
|
||||
if _, err := transport.udp.Write(value); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
read := func(buffer []byte, attemptDeadline time.Time) (int, error) {
|
||||
for {
|
||||
@@ -630,10 +718,6 @@ func (transport *socks5UDP) roundTripLocked(ctx context.Context, packet []byte,
|
||||
return 0, err
|
||||
}
|
||||
if transport.floated {
|
||||
// The relay can deliver ESP before the marked IKE response on
|
||||
// the same UDP/4500 association. Do not accept it as IKE, and
|
||||
// do not abort the exchange; keep waiting within the original
|
||||
// deadline.
|
||||
if !hasNonESPMarker(payload) {
|
||||
continue
|
||||
}
|
||||
@@ -646,7 +730,7 @@ func (transport *socks5UDP) roundTripLocked(ctx context.Context, packet []byte,
|
||||
return len(payload), nil
|
||||
}
|
||||
}
|
||||
return roundTripDatagram(ctx, transport.config.Timeout, write, read, datagram)
|
||||
return roundTripFragments(ctx, transport.config.Timeout, writeAll, read, datagrams)
|
||||
}
|
||||
|
||||
func isNetworkTimeout(err error) bool {
|
||||
|
||||
+14
-12
@@ -31,9 +31,10 @@ const (
|
||||
payloadDelete = 42
|
||||
payloadTSi = 44
|
||||
payloadTSr = 45
|
||||
payloadEncrypted = 46
|
||||
payloadCP = 47
|
||||
payloadEAP = 48
|
||||
payloadEncrypted = 46
|
||||
payloadCP = 47
|
||||
payloadEAP = 48
|
||||
payloadEncryptedFragment = 53
|
||||
|
||||
protocolIKE = 1
|
||||
protocolESP = 3
|
||||
@@ -55,15 +56,16 @@ const (
|
||||
dhMODP2048 = 14
|
||||
transformAttributeKeyLen = 14
|
||||
|
||||
notifyInitialContact = 16384
|
||||
notifyMOBIKESupported = 16396
|
||||
notifyNATSource = 16388
|
||||
notifyNATDestination = 16389
|
||||
notifyCookie = 16390
|
||||
notifyEAPOnlyAuth = 16417
|
||||
notifyDeviceIdentity = 41101
|
||||
notifyInvalidKE = 17
|
||||
notifyNoProposal = 14
|
||||
notifyInitialContact = 16384
|
||||
notifyMOBIKESupported = 16396
|
||||
notifyNATSource = 16388
|
||||
notifyNATDestination = 16389
|
||||
notifyCookie = 16390
|
||||
notifyEAPOnlyAuth = 16417
|
||||
notifyFragmentationSupported = 16430
|
||||
notifyDeviceIdentity = 41101
|
||||
notifyInvalidKE = 17
|
||||
notifyNoProposal = 14
|
||||
)
|
||||
|
||||
var (
|
||||
|
||||
Reference in New Issue
Block a user