diff --git a/internal/vowifi/ike/crypto.go b/internal/vowifi/ike/crypto.go index f91a43d..1ada5f6 100644 --- a/internal/vowifi/ike/crypto.go +++ b/internal/vowifi/ike/crypto.go @@ -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" + diff --git a/internal/vowifi/ike/crypto_test.go b/internal/vowifi/ike/crypto_test.go index 00cb71b..3179bf5 100644 --- a/internal/vowifi/ike/crypto_test.go +++ b/internal/vowifi/ike/crypto_test.go @@ -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) + } +} diff --git a/internal/vowifi/ike/provider.go b/internal/vowifi/ike/provider.go index fa18969..480a065 100644 --- a/internal/vowifi/ike/provider.go +++ b/internal/vowifi/ike/provider.go @@ -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 { diff --git a/internal/vowifi/ike/provider_test.go b/internal/vowifi/ike/provider_test.go index f4b9ace..9ee2d67 100644 --- a/internal/vowifi/ike/provider_test.go +++ b/internal/vowifi/ike/provider_test.go @@ -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 { diff --git a/internal/vowifi/ike/relay.go b/internal/vowifi/ike/relay.go index b4bb0d3..943e87d 100644 --- a/internal/vowifi/ike/relay.go +++ b/internal/vowifi/ike/relay.go @@ -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{ diff --git a/internal/vowifi/ike/relay_test.go b/internal/vowifi/ike/relay_test.go index dae6da9..fcfee35 100644 --- a/internal/vowifi/ike/relay_test.go +++ b/internal/vowifi/ike/relay_test.go @@ -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) } diff --git a/internal/vowifi/ike/transport.go b/internal/vowifi/ike/transport.go index 1b14665..b0f892c 100644 --- a/internal/vowifi/ike/transport.go +++ b/internal/vowifi/ike/transport.go @@ -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 { diff --git a/internal/vowifi/ike/wire.go b/internal/vowifi/ike/wire.go index dc16f8a..a58b522 100644 --- a/internal/vowifi/ike/wire.go +++ b/internal/vowifi/ike/wire.go @@ -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 (