feat: implement IKE session relay and transport layer for ePDG communication

This commit is contained in:
MengMengCode
2026-08-20 13:14:45 +08:00
parent 497cd24c8d
commit b8df7f43f8
8 changed files with 580 additions and 91 deletions
+233
View File
@@ -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" +
+77
View File
@@ -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
View File
@@ -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 {
+14 -3
View File
@@ -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 {
+1 -1
View File
@@ -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{
+11
View File
@@ -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)
}
+127 -43
View File
@@ -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
View File
@@ -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 (