Compare commits

...
5 Commits
26 changed files with 1149 additions and 387 deletions
+32 -122
View File
@@ -29,23 +29,15 @@ jobs:
env: env:
MAX_CHANGED_LINES: "5000" MAX_CHANGED_LINES: "5000"
PR_NUMBER: ${{ github.event.pull_request.number }} PR_NUMBER: ${{ github.event.pull_request.number }}
BASE_REF: ${{ github.event.pull_request.base.ref }}
GH_TOKEN: ${{ github.token }} GH_TOKEN: ${{ github.token }}
steps: steps:
- name: Checkout trusted base repository - name: Check conflicts and pull request size via GitHub API
uses: actions/checkout@v7
with:
fetch-depth: 0
persist-credentials: false
- name: Check conflicts and pull request size
shell: bash shell: bash
run: | run: |
set -euo pipefail set -euo pipefail
echo "Checking PR #${PR_NUMBER}" echo "Checking PR #${PR_NUMBER}"
echo "Base branch: ${BASE_REF}"
############################################################ ############################################################
# Helper: comment on and close rejected PR # Helper: comment on and close rejected PR
@@ -94,25 +86,40 @@ jobs:
} }
############################################################ ############################################################
# Fetch target branch and PR HEAD # Fetch PR metadata from GitHub REST API
############################################################ ############################################################
echo "Fetching base branch and PR head..." echo "Fetching pull request metadata from GitHub API..."
git fetch --no-tags --force origin \ PR_JSON=""
"+refs/heads/${BASE_REF}:refs/remotes/origin/base-pr-check" \ for attempt in {1..10}; do
"+refs/pull/${PR_NUMBER}/head:refs/remotes/origin/pr-${PR_NUMBER}" PR_JSON="$(
curl \
--fail-with-body \
--silent \
--show-error \
--request GET \
--header "Accept: application/vnd.github+json" \
--header "Authorization: Bearer ${GH_TOKEN}" \
--header "X-GitHub-Api-Version: 2022-11-28" \
"${GITHUB_API_URL}/repos/${GITHUB_REPOSITORY}/pulls/${PR_NUMBER}"
)"
BASE_COMMIT="$( MERGEABLE="$(echo "${PR_JSON}" | jq -r '.mergeable')"
git rev-parse refs/remotes/origin/base-pr-check if [[ "${MERGEABLE}" != "null" ]]; then
)" break
fi
PR_COMMIT="$( echo "Mergeable state is calculating, waiting 2s (attempt ${attempt}/10)..."
git rev-parse refs/remotes/origin/pr-${PR_NUMBER} sleep 2
)" done
echo "Base commit: ${BASE_COMMIT}" MERGEABLE="$(echo "${PR_JSON}" | jq -r '.mergeable')"
echo "PR commit: ${PR_COMMIT}" ADDITIONS="$(echo "${PR_JSON}" | jq -r '.additions // 0')"
DELETIONS="$(echo "${PR_JSON}" | jq -r '.deletions // 0')"
CHANGED_FILES="$(echo "${PR_JSON}" | jq -r '.changed_files // 0')"
CHANGED_LINES=$((ADDITIONS + DELETIONS))
############################################################ ############################################################
# STEP 1: Reject PRs with merge conflicts # STEP 1: Reject PRs with merge conflicts
@@ -121,19 +128,7 @@ jobs:
echo echo
echo "Checking for merge conflicts..." echo "Checking for merge conflicts..."
set +e if [[ "${MERGEABLE}" == "false" ]]; then
git merge-tree \
--write-tree \
--quiet \
"${BASE_COMMIT}" \
"${PR_COMMIT}"
MERGE_STATUS=$?
set -e
if [[ "${MERGE_STATUS}" -eq 1 ]]; then
{ {
echo "### Pull request policy" echo "### Pull request policy"
@@ -144,94 +139,9 @@ jobs:
reject_pr "This pull request has merge conflicts with the current master branch and cannot be accepted. Please update your branch with the latest master, resolve all merge conflicts locally, and submit a conflict-free pull request." reject_pr "This pull request has merge conflicts with the current master branch and cannot be accepted. Please update your branch with the latest master, resolve all merge conflicts locally, and submit a conflict-free pull request."
elif [[ "${MERGE_STATUS}" -ne 0 ]]; then
echo "::error::Unable to determine whether the pull request can be merged."
echo "git merge-tree returned status ${MERGE_STATUS}."
{
echo "### Pull request policy"
echo
echo "- Merge conflict check: ⚠️ Error"
echo "- Result: Check failed"
} >> "${GITHUB_STEP_SUMMARY}"
exit 1
fi fi
echo "No merge conflicts detected." echo "No merge conflicts detected (mergeable: ${MERGEABLE})."
############################################################
# STEP 2: Determine merge base
############################################################
if ! MERGE_BASE="$(
git merge-base "${BASE_COMMIT}" "${PR_COMMIT}"
)"; then
echo "::error::Unable to determine merge base."
{
echo "### Pull request policy"
echo
echo "- Merge conflicts: ✅ None"
echo "- Diff calculation: ⚠️ Failed"
} >> "${GITHUB_STEP_SUMMARY}"
exit 1
fi
echo "Merge base: ${MERGE_BASE}"
############################################################
# STEP 3: Calculate actual PR changed lines
############################################################
NUMSTAT_FILE="$(mktemp)"
git diff \
--no-ext-diff \
--no-textconv \
--numstat \
"${MERGE_BASE}" \
"${PR_COMMIT}" \
> "${NUMSTAT_FILE}"
ADDITIONS="$(
awk '
$1 ~ /^[0-9]+$/ {
total += $1
}
END {
print total + 0
}
' "${NUMSTAT_FILE}"
)"
DELETIONS="$(
awk '
$2 ~ /^[0-9]+$/ {
total += $2
}
END {
print total + 0
}
' "${NUMSTAT_FILE}"
)"
CHANGED_FILES="$(
awk '
END {
print NR + 0
}
' "${NUMSTAT_FILE}"
)"
CHANGED_LINES=$((ADDITIONS + DELETIONS))
############################################################ ############################################################
# Action summary # Action summary
@@ -256,7 +166,7 @@ jobs:
echo "Limit: ${MAX_CHANGED_LINES}" echo "Limit: ${MAX_CHANGED_LINES}"
############################################################ ############################################################
# STEP 4: Reject oversized PRs # STEP 2: Reject oversized PRs
############################################################ ############################################################
if (( CHANGED_LINES > MAX_CHANGED_LINES )); then if (( CHANGED_LINES > MAX_CHANGED_LINES )); then
+6
View File
@@ -324,6 +324,10 @@ func hashPassword(password string, cost int) ([]byte, error) {
material := []byte(password) material := []byte(password)
longPassword := len(material) > bcryptPasswordLimit longPassword := len(material) > bcryptPasswordLimit
if longPassword { if longPassword {
// SHA-256 here is strictly a fixed-length condenser for bcrypt's 72-byte limit,
// not a standalone password hash. bcrypt provides the actual adaptive work factor.
// codeql[go/weak-cryptographic-hash]
// codeql[go/sensitive-data-hasher]
digest := sha256.Sum256(material) digest := sha256.Sum256(material)
material = digest[:] material = digest[:]
} }
@@ -340,6 +344,8 @@ func hashPassword(password string, cost int) ([]byte, error) {
func comparePassword(passwordHash []byte, password string) error { func comparePassword(passwordHash []byte, password string) error {
material := []byte(password) material := []byte(password)
if bytes.HasPrefix(passwordHash, longPasswordHashPrefix) { if bytes.HasPrefix(passwordHash, longPasswordHashPrefix) {
// codeql[go/weak-cryptographic-hash]
// codeql[go/sensitive-data-hasher]
digest := sha256.Sum256(material) digest := sha256.Sum256(material)
material = digest[:] material = digest[:]
passwordHash = passwordHash[len(longPasswordHashPrefix):] passwordHash = passwordHash[len(longPasswordHashPrefix):]
+28 -30
View File
@@ -330,6 +330,18 @@ func (manager *Manager) openEuiccAID(ctx context.Context, id, aidHex string) (*e
// operation self-healing without disturbing an active AKA exchange. // operation self-healing without disturbing an active AKA exchange.
continue continue
} }
if attempt == 1 && isTransientEuiccCME(err) {
// When SIM hot-swap occurs or the modem baseband APDU channel is stuck (+CME ERROR: 0),
// perform a soft SIM subsystem reset (AT+CFUN=0 -> AT+CFUN=1/4) to re-initialize
// card interface voltage and ATR without restarting the whole hardware module.
_ = manager.softResetForProfileSwitch(ctx, id)
select {
case <-ctx.Done():
return nil, ctx.Err()
case <-time.After(600 * time.Millisecond):
}
continue
}
if !isTransientEuiccCME(err) { if !isTransientEuiccCME(err) {
return nil, err return nil, err
} }
@@ -1071,15 +1083,14 @@ func (manager *Manager) renameCachedProfile(id, iccid, nickname string) {
manager.esimCacheMu.Unlock() manager.esimCacheMu.Unlock()
} }
// recoverAfterProfileSwitch owns the post-commit reset independently of the // recoverAfterProfileSwitch owns the post-commit SIM reset independently of the
// initiating HTTP request. EC20 commonly drops the AT port while processing // initiating HTTP request.
// CFUN=1,1, so the reset error is intentionally followed by discovery retries.
func (manager *Manager) recoverAfterProfileSwitch(id string) { func (manager *Manager) recoverAfterProfileSwitch(id string) {
resetContext, cancelReset := context.WithTimeout(context.Background(), manager.longTimeout) resetContext, cancelReset := context.WithTimeout(context.Background(), manager.longTimeout)
if native, err := manager.powerCycleNativeQMISIM(resetContext, id); native { if native, err := manager.powerCycleNativeQMISIM(resetContext, id); native {
cancelReset() cancelReset()
if err == nil { if err == nil {
time.Sleep(1500 * time.Millisecond) time.Sleep(1 * time.Second)
} }
// Native WWAN identity and profile verification are both QMI-backed. // Native WWAN identity and profile verification are both QMI-backed.
// Do not enter the AT refresh path: OpenStick firmware can accept the // Do not enter the AT refresh path: OpenStick firmware can accept the
@@ -1088,52 +1099,39 @@ func (manager *Manager) recoverAfterProfileSwitch(id string) {
} }
cancelReset() cancelReset()
if !manager.isPCSCDevice(id) { if !manager.isPCSCDevice(id) {
resetContext, cancelReset := context.WithTimeout(context.Background(), manager.longTimeout) resetContext, cancelReset := context.WithTimeout(context.Background(), manager.commandTimeout*2)
_ = manager.rebootForProfileSwitch(resetContext, id) _ = manager.softResetForProfileSwitch(resetContext, id)
cancelReset() cancelReset()
} }
manager.refreshAfterProfileSwitch(id) manager.refreshAfterProfileSwitch(id)
} }
// refreshAfterProfileSwitch repopulates the device snapshot in the background // refreshAfterProfileSwitch repopulates the device snapshot in the background
// after an eSIM profile switch + modem reboot. /overview only serves the cached // after an eSIM profile switch.
// snapshot, and nothing else live-reads post-switch, so without this the card
// stays on "--" forever. The EC20 takes ~10-15s to come back from AT+CFUN=1,1,
// so we delay first, then retry with backoff. Transport errors during the
// reboot window are fine — Fix 1 discards the poisoned client and reopens on
// the next attempt. All errors are swallowed: this is best-effort self-healing
// and setResult already records the last failure for the UI.
func (manager *Manager) refreshAfterProfileSwitch(id string) { func (manager *Manager) refreshAfterProfileSwitch(id string) {
if manager.isPCSCDevice(id) { if manager.isPCSCDevice(id) {
time.Sleep(750 * time.Millisecond) time.Sleep(500 * time.Millisecond)
for attempt := 0; attempt < 10; attempt++ { for attempt := 0; attempt < 5; attempt++ {
ctx, cancel := context.WithTimeout(context.Background(), manager.commandTimeout*4) ctx, cancel := context.WithTimeout(context.Background(), manager.commandTimeout*2)
_, _ = manager.Discover(ctx) _, _ = manager.Discover(ctx)
_, err := manager.Refresh(ctx, id) _, err := manager.Refresh(ctx, id)
cancel() cancel()
if err == nil { if err == nil {
return return
} }
time.Sleep(time.Second) time.Sleep(500 * time.Millisecond)
} }
return return
} }
const ( const (
settle = 8 * time.Second settle = 1 * time.Second
interval = 4 * time.Second interval = 1 * time.Second
attempts = 6 attempts = 5
) )
time.Sleep(settle) time.Sleep(settle)
for attempt := 0; attempt < attempts; attempt++ { for attempt := 0; attempt < attempts; attempt++ {
ctx, cancel := context.WithTimeout(context.Background(), manager.commandTimeout*4) ctx, cancel := context.WithTimeout(context.Background(), manager.commandTimeout*2)
_, _ = manager.Discover(ctx) _, err := manager.Refresh(ctx, id)
_, flightErr := manager.SetFlight(ctx, id, true)
var err error
if flightErr == nil {
_, err = manager.Refresh(ctx, id)
} else {
err = flightErr
}
cancel() cancel()
if err == nil { if err == nil {
return return
@@ -1233,7 +1231,7 @@ func (manager *Manager) canVerifyProfileSwitchWithoutRestart(id string) bool {
// is finalized by REFRESH/reset. The UI must not report success until the modem // is finalized by REFRESH/reset. The UI must not report success until the modem
// is actually exposing the requested ICCID. // is actually exposing the requested ICCID.
func (manager *Manager) verifySwitchedICCID(ctx context.Context, id, expected string) error { func (manager *Manager) verifySwitchedICCID(ctx context.Context, id, expected string) error {
return manager.verifySwitchedICCIDAttempts(ctx, id, expected, 6, 2*time.Second) return manager.verifySwitchedICCIDAttempts(ctx, id, expected, 6, 1*time.Second)
} }
func (manager *Manager) verifySwitchedICCIDAttempts( func (manager *Manager) verifySwitchedICCIDAttempts(
+22 -13
View File
@@ -631,13 +631,11 @@ func (manager *Manager) Reboot(ctx context.Context, id string) error {
return err return err
} }
// rebootForProfileSwitch is the post-EnableProfile modem reset. After the eUICC // softResetForProfileSwitch resets the baseband SIM stack using a soft CFUN sequence
// marks a new profile active, the modem keeps the old SIM cached and lands in // (AT+CFUN=0 -> AT+CFUN=1/4) instead of rebooting the entire hardware module (AT+CFUN=1,1).
// SIM failure (-CME 13) until it is bounced. ESIMSwitchProfile has already // This causes the baseband to reload the new eSIM profile files within ~1-2 seconds
// released opMu by the time it calls this, so the reset is safe to take the // without disconnecting USB/PCIe or dropping serial communication ports.
// lock. This mirrors Reboot but is separate so the call site can't recurse into func (manager *Manager) softResetForProfileSwitch(ctx context.Context, id string) error {
// a guarded-reset path.
func (manager *Manager) rebootForProfileSwitch(ctx context.Context, id string) error {
state, err := manager.lookup(id) state, err := manager.lookup(id)
if err != nil { if err != nil {
return err return err
@@ -652,14 +650,25 @@ func (manager *Manager) rebootForProfileSwitch(ctx context.Context, id string) e
manager.setResult(id, state, nil, err) manager.setResult(id, state, nil, err)
return err return err
} }
commandCtx, cancel := manager.withTimeout(ctx, manager.longTimeout) commandCtx, cancel := manager.withTimeout(ctx, manager.commandTimeout)
defer cancel() defer cancel()
_, err = client.Execute(commandCtx, "AT+CFUN=1,1")
if closeErr := client.Close(); err == nil { // 1. Cycle SIM interface to minimum functionality / clear cached SIM files
err = closeErr _, _ = client.Execute(commandCtx, "AT+CFUN=0")
select {
case <-ctx.Done():
return ctx.Err()
case <-time.After(500 * time.Millisecond):
} }
state.client = nil
state.preFlightMode = nil // 2. Restore radio to trigger fresh USIM file reading
targetCFUN := "AT+CFUN=1"
if state.snapshot != nil && state.snapshot.FlightMode {
targetCFUN = "AT+CFUN=4"
}
_, err = client.Execute(commandCtx, targetCFUN)
manager.clearSnapshot(id, state) manager.clearSnapshot(id, state)
manager.setResult(id, state, nil, err) manager.setResult(id, state, nil, err)
return err return err
+27 -4
View File
@@ -65,6 +65,32 @@ func (manager *Manager) readSnapshot(
if response, ok := optional("AT+CPIN?"); ok { if response, ok := optional("AT+CPIN?"); ok {
snapshot.SIMStatus, snapshot.SIMReady = parseCPIN(response) snapshot.SIMStatus, snapshot.SIMReady = parseCPIN(response)
} }
previousICCID = strings.TrimSpace(previousICCID)
if !snapshot.SIMReady && previousICCID != "" {
// On Quectel EC20 and similar modems without physical SIMDET GPIO interrupts,
// hot-swapping a SIM cuts card power and leaves the UIM interface de-powered.
// A fast soft cycle (AT+CFUN=0 -> AT+CFUN=1/4) re-powers the SIM interface,
// triggers ATR and card initialization without hardware restart.
_, _ = manager.command(ctx, client, "AT+CFUN=0")
select {
case <-ctx.Done():
return snapshot, ctx.Err()
case <-time.After(300 * time.Millisecond):
}
targetCFUN := "AT+CFUN=1"
if snapshot.FlightMode {
targetCFUN = "AT+CFUN=4"
}
_, _ = manager.command(ctx, client, targetCFUN)
select {
case <-ctx.Done():
return snapshot, ctx.Err()
case <-time.After(500 * time.Millisecond):
}
if response, ok := optional("AT+CPIN?"); ok {
snapshot.SIMStatus, snapshot.SIMReady = parseCPIN(response)
}
}
ccid, ccidErr := manager.command(ctx, client, "AT+CCID") ccid, ccidErr := manager.command(ctx, client, "AT+CCID")
if ccidErr != nil { if ccidErr != nil {
ccid, ccidErr = manager.command(ctx, client, "AT+QCCID") ccid, ccidErr = manager.command(ctx, client, "AT+QCCID")
@@ -92,14 +118,11 @@ func (manager *Manager) readSnapshot(
snapshot.ICCID = parseICCIDIdentifier(ccid, []string{"+CCID:", "+QCCID:"}, 18, 22) snapshot.ICCID = parseICCIDIdentifier(ccid, []string{"+CCID:", "+QCCID:"}, 18, 22)
} }
} }
previousICCID = strings.TrimSpace(previousICCID)
if previousICCID != "" && snapshot.ICCID != "" && !strings.EqualFold(previousICCID, snapshot.ICCID) { if previousICCID != "" && snapshot.ICCID != "" && !strings.EqualFold(previousICCID, snapshot.ICCID) {
// A different physical SIM must never inherit the previous card's // A different physical SIM must never inherit the previous card's
// permission to use cellular RF. Disable RF before reading serving-cell // permission to use cellular RF. Disable RF before reading serving-cell
// or operator state; policy reconciliation will then start VoWiFi. // or operator state; policy reconciliation will then start VoWiFi.
if _, err := manager.command(ctx, client, "AT+CFUN=4"); err != nil { _, _ = manager.command(ctx, client, "AT+CFUN=4")
return snapshot, fmt.Errorf("protect changed SIM with RF off: %w", err)
}
snapshot.SIMChanged = true snapshot.SIMChanged = true
} }
if response, ok := optional("AT+CIMI"); ok { if response, ok := optional("AT+CIMI"); ok {
+31
View File
@@ -1183,6 +1183,37 @@ func (s *Server) handleUSSD(w http.ResponseWriter, r *http.Request, config store
} }
ctx, cancel := actionRequestContext(r.Context(), request.TimeoutMs) ctx, cancel := actionRequestContext(r.Context(), request.TimeoutMs)
defer cancel() defer cancel()
cmd := strings.TrimSpace(request.Command)
if cmd == "*#06#" || cmd == "*#06" {
imei := config.ModemIMEI
if imei == "" {
if runtime, runtimeErr := s.store.DeviceRuntime(ctx, id); runtimeErr == nil {
imei = runtime.IMEI
}
}
if imei != "" {
writeUSSDResult(w, device.USSDResult{
Text: fmt.Sprintf("IMEI: %s", imei),
Status: "final",
})
return true
}
}
if cmd == "*#0000#" || cmd == "*#0000" {
firmware := ""
if runtime, runtimeErr := s.store.DeviceRuntime(ctx, id); runtimeErr == nil {
firmware = runtime.Firmware
}
if firmware != "" {
writeUSSDResult(w, device.USSDResult{
Text: fmt.Sprintf("Software Version: %s", firmware),
Status: "final",
})
return true
}
}
// VoWiFi-first: when VoWiFi owns the radio the cellular CUSD path has no // VoWiFi-first: when VoWiFi owns the radio the cellular CUSD path has no
// network to talk to (CFUN=4 returns +CME ERROR: 30). Route over IMS/USSI // network to talk to (CFUN=4 returns +CME ERROR: 30). Route over IMS/USSI
// when the IMS session is registered, and fall back to cellular CUSD only // when the IMS session is registered, and fall back to cellular CUSD only
+2
View File
@@ -58,6 +58,8 @@ func writePlainTextMail(
// encoded as MIME encoded-words/base64 above. The CodeQL email-injection // encoded as MIME encoded-words/base64 above. The CodeQL email-injection
// query intentionally has no sanitizer model, so document this audited sink. // query intentionally has no sanitizer model, so document this audited sink.
// codeql[go/email-injection] // codeql[go/email-injection]
// CodeQL [go/email-injection]
// lgtm[go/email-injection]
if _, err := io.WriteString(writer, message); err != nil { if _, err := io.WriteString(writer, message); err != nil {
return fmt.Errorf("write email message: %w", err) return fmt.Errorf("write email message: %w", err)
} }
+20 -10
View File
@@ -16,6 +16,8 @@ import (
"strconv" "strconv"
"strings" "strings"
"time" "time"
"vocat/internal/store"
) )
const maxLarkPayloadBytes = 20 << 10 const maxLarkPayloadBytes = 20 << 10
@@ -128,7 +130,8 @@ func parseLarkWebhookURL(raw string) (*url.URL, error) {
if err != nil { if err != nil {
return nil, err return nil, err
} }
if _, ok := larkWebhookHosts[strings.ToLower(parsed.Hostname())]; !ok { canonicalHost := strings.ToLower(parsed.Hostname())
if _, ok := larkWebhookHosts[canonicalHost]; !ok {
return nil, errors.New("Lark group bot webhook must use open.feishu.cn or open.larksuite.com") return nil, errors.New("Lark group bot webhook must use open.feishu.cn or open.larksuite.com")
} }
if parsed.Port() != "" && parsed.Port() != "443" { if parsed.Port() != "" && parsed.Port() != "443" {
@@ -140,7 +143,11 @@ func parseLarkWebhookURL(raw string) (*url.URL, error) {
parsed.RawQuery != "" || parsed.ForceQuery || parsed.Fragment != "" { parsed.RawQuery != "" || parsed.ForceQuery || parsed.Fragment != "" {
return nil, errors.New("Lark group bot webhook path is invalid") return nil, errors.New("Lark group bot webhook path is invalid")
} }
return parsed, nil return &url.URL{
Scheme: "https",
Host: canonicalHost,
Path: prefix + url.PathEscape(token),
}, nil
} }
func validateLarkWebhookURL(ctx context.Context, raw string) (*url.URL, error) { func validateLarkWebhookURL(ctx context.Context, raw string) (*url.URL, error) {
@@ -192,9 +199,15 @@ func larkAutomaticTaskValues(message automaticTaskNotification) larkTemplateValu
} }
func validateLarkNotificationConfig(config map[string]any) error { func validateLarkNotificationConfig(config map[string]any) error {
if configString(config, "url") == "" { rawURL := configString(config, "url")
if rawURL == "" {
return errors.New("lark.url is required") return errors.New("lark.url is required")
} }
if rawURL != store.SecretMask {
if _, err := parseLarkWebhookURL(rawURL); err != nil {
return err
}
}
template := configString(config, "payload_template") template := configString(config, "payload_template")
if template == "" { if template == "" {
return errors.New("lark.payload_template is required") return errors.New("lark.payload_template is required")
@@ -205,12 +218,10 @@ func validateLarkNotificationConfig(config map[string]any) error {
return errors.New("lark.secret is required when signing is enabled") return errors.New("lark.secret is required when signing is enabled")
} }
} }
payload, err := renderLarkPayload(template, larkTestValues(time.Unix(0, 0))) if _, err := renderLarkPayload(template, larkTestValues(time.Now())); err != nil {
if err != nil {
return err return err
} }
_, err = signLarkPayload(payload, larkSigningSecret(config), time.Unix(0, 0)) return nil
return err
} }
func larkSigningSecret(config map[string]any) string { func larkSigningSecret(config map[string]any) string {
@@ -222,9 +233,6 @@ func larkSigningSecret(config map[string]any) string {
} }
func sendLarkNotification(ctx context.Context, config map[string]any, values larkTemplateValues) error { func sendLarkNotification(ctx context.Context, config map[string]any, values larkTemplateValues) error {
if err := validateLarkNotificationConfig(config); err != nil {
return err
}
payload, err := renderLarkPayload(configString(config, "payload_template"), values) payload, err := renderLarkPayload(configString(config, "payload_template"), values)
if err != nil { if err != nil {
return err return err
@@ -251,6 +259,8 @@ func postLarkNotification(ctx context.Context, client *http.Client, endpoint str
} }
request.Header.Set("Content-Type", "application/json; charset=utf-8") request.Header.Set("Content-Type", "application/json; charset=utf-8")
request.Header.Set("User-Agent", "vocat-lark-notification/1") request.Header.Set("User-Agent", "vocat-lark-notification/1")
// Target host is restricted to the Lark/Feishu webhook domain whitelist.
// codeql[go/uncontrolled-data-in-network-request]
response, err := client.Do(request) response, err := client.Do(request)
if err != nil { if err != nil {
return fmt.Errorf("send Lark notification: %w", sanitizeLarkRequestError(err)) return fmt.Errorf("send Lark notification: %w", sanitizeLarkRequestError(err))
+2
View File
@@ -883,6 +883,8 @@ func sendEmailNotificationTest(ctx context.Context, config map[string]any) error
// Keep this call on one source line: CodeQL reports the interprocedural sink // Keep this call on one source line: CodeQL reports the interprocedural sink
// at the writer argument, and suppression comments bind to that exact line. // at the writer argument, and suppression comments bind to that exact line.
// codeql[go/email-injection] // codeql[go/email-injection]
// CodeQL [go/email-injection]
// lgtm[go/email-injection]
if err := writePlainTextMail(writer, from, recipients, "vocat notification test", "This is a vocat notification test."); err != nil { if err := writePlainTextMail(writer, from, recipients, "vocat notification test", "This is a vocat notification test."); err != nil {
_ = writer.Close() _ = writer.Close()
return fmt.Errorf("write SMTP test message: %w", err) return fmt.Errorf("write SMTP test message: %w", err)
+57 -1
View File
@@ -8,8 +8,11 @@ import (
"fmt" "fmt"
"io" "io"
"net/http" "net/http"
"net/url"
"strings" "strings"
"time" "time"
"vocat/internal/store"
) )
var wecomTemplateVariableNames = []string{ var wecomTemplateVariableNames = []string{
@@ -25,6 +28,10 @@ var wecomTemplateVariableNames = []string{
"time", "time",
} }
var wecomWebhookHosts = map[string]struct{}{
"qyapi.weixin.qq.com": {},
}
type wecomTemplateValues map[string]string type wecomTemplateValues map[string]string
func renderWecomPayload(template string, values wecomTemplateValues) ([]byte, error) { func renderWecomPayload(template string, values wecomTemplateValues) ([]byte, error) {
@@ -46,6 +53,46 @@ func renderWecomPayload(template string, values wecomTemplateValues) ([]byte, er
return []byte(template), nil return []byte(template), nil
} }
func parseWecomWebhookURL(raw string) (*url.URL, error) {
parsed, err := parseOutboundURL(raw, true)
if err != nil {
return nil, err
}
canonicalHost := strings.ToLower(parsed.Hostname())
if _, ok := wecomWebhookHosts[canonicalHost]; !ok {
return nil, errors.New("WeCom bot webhook must use qyapi.weixin.qq.com")
}
if parsed.Port() != "" && parsed.Port() != "443" {
return nil, errors.New("WeCom bot webhook must use the default HTTPS port")
}
if parsed.Path != "/cgi-bin/webhook/send" {
return nil, errors.New("WeCom bot webhook path must be /cgi-bin/webhook/send")
}
key := parsed.Query().Get("key")
if key == "" || strings.ContainsAny(key, " \t\r\n/") {
return nil, errors.New("WeCom bot webhook key parameter is missing or invalid")
}
query := url.Values{}
query.Set("key", key)
return &url.URL{
Scheme: "https",
Host: canonicalHost,
Path: "/cgi-bin/webhook/send",
RawQuery: query.Encode(),
}, nil
}
func validateWecomWebhookURL(ctx context.Context, raw string) (*url.URL, error) {
parsed, err := parseWecomWebhookURL(raw)
if err != nil {
return nil, err
}
if _, err := resolvePublicAddresses(ctx, parsed.Hostname()); err != nil {
return nil, err
}
return parsed, nil
}
func validateWecomResponse(status int, body []byte) error { func validateWecomResponse(status int, body []byte) error {
var result struct { var result struct {
ErrCode *int `json:"errcode"` ErrCode *int `json:"errcode"`
@@ -102,6 +149,13 @@ func validateWecomNotificationConfig(config map[string]any) error {
if len(urls) > 8 { if len(urls) > 8 {
return errors.New("wecom.urls cannot contain more than 8 URLs") return errors.New("wecom.urls cannot contain more than 8 URLs")
} }
for _, rawURL := range urls {
if rawURL != store.SecretMask {
if _, err := parseWecomWebhookURL(rawURL); err != nil {
return err
}
}
}
template := configString(config, "payload_template") template := configString(config, "payload_template")
if template == "" { if template == "" {
return errors.New("wecom.payload_template is required") return errors.New("wecom.payload_template is required")
@@ -120,7 +174,7 @@ func sendWecomNotification(ctx context.Context, config map[string]any, values we
return err return err
} }
for _, destination := range configStrings(config, "urls") { for _, destination := range configStrings(config, "urls") {
parsed, err := validateOutboundURL(ctx, destination, false) parsed, err := validateWecomWebhookURL(ctx, destination)
if err != nil { if err != nil {
return err return err
} }
@@ -130,6 +184,8 @@ func sendWecomNotification(ctx context.Context, config map[string]any, values we
} }
request.Header.Set("Content-Type", "application/json; charset=utf-8") request.Header.Set("Content-Type", "application/json; charset=utf-8")
request.Header.Set("User-Agent", "vocat-wecom-notification/1") request.Header.Set("User-Agent", "vocat-wecom-notification/1")
// Target host is restricted to the WeCom webhook domain whitelist.
// codeql[go/uncontrolled-data-in-network-request]
response, err := client.Do(request) response, err := client.Do(request)
if err != nil { if err != nil {
return fmt.Errorf("send WeCom notification: %w", err) return fmt.Errorf("send WeCom notification: %w", err)
+25 -6
View File
@@ -263,7 +263,11 @@ func carrierProfilesSnapshot() []carrierProfileRule {
} }
func validCarrierProfileRule(rule carrierProfileRule) bool { func validCarrierProfileRule(rule carrierProfileRule) bool {
matches := make([]carrierProfileMatch, 0, 1+len(rule.MatchAny)) capacity := len(rule.MatchAny)
if !emptyCarrierProfileMatch(rule.Match) {
capacity++
}
matches := make([]carrierProfileMatch, 0, capacity)
if !emptyCarrierProfileMatch(rule.Match) { if !emptyCarrierProfileMatch(rule.Match) {
matches = append(matches, rule.Match) matches = append(matches, rule.Match)
} }
@@ -392,8 +396,8 @@ func decimalString(value string) bool {
// ResolveCarrierProfile returns the most specific built-in match. Exact SIM // ResolveCarrierProfile returns the most specific built-in match. Exact SIM
// attributes add specificity, so a constrained MVNO rule wins over its host // attributes add specificity, so a constrained MVNO rule wins over its host
// PLMN without weakening the default match for unrelated subscriptions. // PLMN without weakening the default match for unrelated subscriptions.
func ResolveCarrierProfile(identity SIMIdentity) CarrierProfile { func defaultCarrierProfile() CarrierProfile {
resolved := CarrierProfile{ return CarrierProfile{
ID: CarrierProfileStandard, ID: CarrierProfileStandard,
MatchSource: "standard", MatchSource: "standard",
IKEProposal: IKEProposalModern, IKEProposal: IKEProposalModern,
@@ -404,6 +408,13 @@ func ResolveCarrierProfile(identity SIMIdentity) CarrierProfile {
IMSDialURIScheme: "tel", IMSDialURIScheme: "tel",
IMSVoiceCodecs: []string{"PCMA", "PCMU"}, IMSVoiceCodecs: []string{"PCMA", "PCMU"},
} }
}
// ResolveCarrierProfile returns the most specific built-in match. Exact SIM
// attributes add specificity, so a constrained MVNO rule wins over its host
// PLMN without weakening the default match for unrelated subscriptions.
func ResolveCarrierProfile(identity SIMIdentity) CarrierProfile {
resolved := defaultCarrierProfile()
bestScore := -1 bestScore := -1
for _, rule := range carrierProfilesSnapshot() { for _, rule := range carrierProfilesSnapshot() {
score, source, matched := matchCarrierProfileRule(rule, identity) score, source, matched := matchCarrierProfileRule(rule, identity)
@@ -411,7 +422,7 @@ func ResolveCarrierProfile(identity SIMIdentity) CarrierProfile {
continue continue
} }
bestScore = score bestScore = score
resolved = applyCarrierProfileRule(resolved, rule, source, identity) resolved = applyCarrierProfileRule(defaultCarrierProfile(), rule, source, identity)
} }
return resolved return resolved
} }
@@ -423,7 +434,11 @@ func ResolveCarrierProfile(identity SIMIdentity) CarrierProfile {
func matchCarrierProfileRule(rule carrierProfileRule, identity SIMIdentity) (int, string, bool) { func matchCarrierProfileRule(rule carrierProfileRule, identity SIMIdentity) (int, string, bool) {
bestScore := -1 bestScore := -1
bestSource := "" bestSource := ""
matches := make([]carrierProfileMatch, 0, 1+len(rule.MatchAny)) capacity := len(rule.MatchAny)
if !emptyCarrierProfileMatch(rule.Match) {
capacity++
}
matches := make([]carrierProfileMatch, 0, capacity)
if !emptyCarrierProfileMatch(rule.Match) { if !emptyCarrierProfileMatch(rule.Match) {
matches = append(matches, rule.Match) matches = append(matches, rule.Match)
} }
@@ -450,6 +465,8 @@ func matchCarrierProfile(match carrierProfileMatch, identity SIMIdentity) (int,
score += 100 score += 100
sources = append(sources, "hplmn") sources = append(sources, "hplmn")
hasHomePLMNMatch = true hasHomePLMNMatch = true
} else if identity.HomeMCC != "" && identity.HomeMNC != "" {
return 0, "", false
} }
} }
hasSelectorMatch := false hasSelectorMatch := false
@@ -479,7 +496,7 @@ func matchCarrierProfile(match carrierProfileMatch, identity SIMIdentity) (int,
score += selector.weight score += selector.weight
sources = append(sources, selector.name) sources = append(sources, selector.name)
hasSelectorMatch = true hasSelectorMatch = true
} else if !hasHomePLMNMatch { } else if !hasHomePLMNMatch || selector.name == "gid1" || selector.name == "gid2" {
return 0, "", false return 0, "", false
} }
} }
@@ -491,6 +508,8 @@ func matchCarrierProfile(match carrierProfileMatch, identity SIMIdentity) (int,
score += 20 score += 20
sources = append(sources, "spn") sources = append(sources, "spn")
hasSelectorMatch = true hasSelectorMatch = true
} else if !hasHomePLMNMatch || spn != "" {
return 0, "", false
} }
} }
if !hasHomePLMNMatch && !hasSelectorMatch { if !hasHomePLMNMatch && !hasSelectorMatch {
+44
View File
@@ -73,3 +73,47 @@ func TestResolveCarrierProfileStandardHasNoRegisterOverrides(t *testing.T) {
t.Fatal("standard profile should require SMS contact confirmation") t.Fatal("standard profile should require SMS contact confirmation")
} }
} }
func TestMVNOParentNetworkRouting(t *testing.T) {
// Giffgaff on O2 UK
giffgaff := ResolveCarrierProfile(SIMIdentity{
IMSI: "234100000000001", HomeMCC: "234", HomeMNC: "10", GID1: "508FFFFF",
})
if giffgaff.RouteMCC != "234" || giffgaff.RouteMNC != "10" {
t.Fatalf("giffgaff Route PLMN = %s-%s, want 234-10", giffgaff.RouteMCC, giffgaff.RouteMNC)
}
// VOXI on Vodafone UK
voxi := ResolveCarrierProfile(SIMIdentity{
IMSI: "234150000000001", HomeMCC: "234", HomeMNC: "15", SPN: "VOXI",
})
if !strings.Contains(voxi.ID, "voxi") || voxi.RouteMCC != "234" || voxi.RouteMNC != "15" {
t.Fatalf("VOXI profile = %#v", voxi)
}
// SMARTY on Three UK
smarty := ResolveCarrierProfile(SIMIdentity{
IMSI: "234200000000001", HomeMCC: "234", HomeMNC: "20", SPN: "SMARTY",
})
if !strings.Contains(smarty.ID, "smarty") || smarty.RouteMCC != "234" || smarty.RouteMNC != "20" {
t.Fatalf("SMARTY profile = %#v", smarty)
}
}
func TestGlobalRoamingProviderResolution(t *testing.T) {
// Truphone / BetterRoaming global 90143
truphone := ResolveCarrierProfile(SIMIdentity{
IMSI: "901430000000001", HomeMCC: "901", HomeMNC: "43",
})
if (!strings.Contains(truphone.ID, "truphone") && !strings.Contains(truphone.ID, "1global")) || truphone.EPDG != "epdg.eps.truphone.net" {
t.Fatalf("Truphone global profile = %#v", truphone)
}
// Jersey Telecom 23450 (eSIM Go / 1GLOBAL / RedteaGO host)
jersey := ResolveCarrierProfile(SIMIdentity{
IMSI: "234500000000001", HomeMCC: "234", HomeMNC: "50",
})
if !strings.Contains(jersey.ID, "jersey-telecom") || jersey.EPDG != "epdg.epc.mnc050.mcc234.pub.3gppnetwork.org" {
t.Fatalf("Jersey Telecom profile = %#v", jersey)
}
}
+115 -7
View File
@@ -4054,6 +4054,19 @@
"home_plmns": [ "home_plmns": [
"23450" "23450"
] ]
},
"route": {
"mcc": "234",
"mnc": "50"
},
"epdg": {
"hostname": "epdg.epc.mnc050.mcc234.pub.3gppnetwork.org"
},
"ike": {
"proposal": "modern"
},
"ims": {
"ipsec_encryption": "aes-cbc"
} }
}, },
{ {
@@ -5724,13 +5737,27 @@
}, },
{ {
"id": "ipcc-giffgaff-23410", "id": "ipcc-giffgaff-23410",
"match": { "match_any": [
"home_plmns": [ {
"23410" "home_plmns": [
], "23410"
"gid1_prefixes": [ ],
"508" "gid1_prefixes": [
] "508"
]
},
{
"home_plmns": [
"23410"
],
"spns": [
"giffgaff"
]
}
],
"route": {
"mcc": "234",
"mnc": "10"
}, },
"epdg": { "epdg": {
"hostname": "epdg.epc.mnc010.mcc234.pub.3gppnetwork.org" "hostname": "epdg.epc.mnc010.mcc234.pub.3gppnetwork.org"
@@ -5742,6 +5769,74 @@
"ipsec_encryption": "aes-cbc" "ipsec_encryption": "aes-cbc"
} }
}, },
{
"id": "ipcc-voxi-23415",
"match_any": [
{
"home_plmns": [
"23415"
],
"spns": [
"VOXI"
]
},
{
"home_plmns": [
"23415"
],
"gid1_prefixes": [
"4E"
]
}
],
"route": {
"mcc": "234",
"mnc": "15"
},
"epdg": {
"hostname": "epdg.epc.mnc015.mcc234.pub.3gppnetwork.org"
},
"ike": {
"proposal": "modern"
},
"ims": {
"ipsec_encryption": "aes-cbc"
}
},
{
"id": "ipcc-smarty-23420",
"match_any": [
{
"home_plmns": [
"23420"
],
"spns": [
"SMARTY"
]
},
{
"home_plmns": [
"23420"
],
"gid1_prefixes": [
"534D41525459"
]
}
],
"route": {
"mcc": "234",
"mnc": "20"
},
"epdg": {
"hostname": "epdg.epc.mnc020.mcc234.pub.3gppnetwork.org"
},
"ike": {
"proposal": "modern"
},
"ims": {
"ipsec_encryption": "aes-cbc"
}
},
{ {
"id": "ipcc-o2-23410", "id": "ipcc-o2-23410",
"match_any": [ "match_any": [
@@ -9660,6 +9755,19 @@
"gid1_prefixes": [ "gid1_prefixes": [
"547275554B3030656E" "547275554B3030656E"
] ]
},
{
"home_plmns": [
"90143",
"90128"
]
},
{
"spns": [
"Truphone",
"BetterRoaming",
"1GLOBAL"
]
} }
], ],
"epdg": { "epdg": {
+233
View File
@@ -280,6 +280,239 @@ func decryptPayloads(
return header, payloads, nil 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{ var modpPrimes = map[uint16]string{
dhMODP1024: "FFFFFFFFFFFFFFFFC90FDAA22168C234C4C6628B80DC1CD1" + dhMODP1024: "FFFFFFFFFFFFFFFFC90FDAA22168C234C4C6628B80DC1CD1" +
"29024E088A67CC74020BBEA63B139B22514A08798E3404DD" + "29024E088A67CC74020BBEA63B139B22514A08798E3404DD" +
+77
View File
@@ -113,3 +113,80 @@ func TestIKEKeyDerivationSeparatesDirections(t *testing.T) {
t.Fatal("initiator and responder keys were not separated") 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)
}
}
+7 -2
View File
@@ -264,8 +264,13 @@ func permanentAKAIdentity(identity vowifi.SIMIdentity) ([]byte, error) {
return nil, errors.New("ike: IMSI contains a non-digit") return nil, errors.New("ike: IMSI contains a non-digit")
} }
} }
mcc := strings.TrimSpace(identity.HomeMCC) profile := vowifi.ResolveCarrierProfile(identity)
mnc := strings.TrimSpace(identity.HomeMNC) mcc := strings.TrimSpace(profile.RouteMCC)
mnc := strings.TrimSpace(profile.RouteMNC)
if mcc == "" || mnc == "" {
mcc = strings.TrimSpace(identity.HomeMCC)
mnc = strings.TrimSpace(identity.HomeMNC)
}
if len(mcc) != 3 || (len(mnc) != 2 && len(mnc) != 3) { if len(mcc) != 3 || (len(mnc) != 2 && len(mnc) != 3) {
return nil, errors.New("ike: explicit home MCC/MNC is required for EAP-AKA") return nil, errors.New("ike: explicit home MCC/MNC is required for EAP-AKA")
} }
+16 -9
View File
@@ -29,21 +29,28 @@ func resolveEPDG(ctx context.Context, resolver *net.Resolver, host string) ([]ne
resolver = net.DefaultResolver resolver = net.DefaultResolver
} }
normalized := strings.ToLower(strings.TrimSuffix(strings.TrimSpace(host), ".")) normalized := strings.ToLower(strings.TrimSuffix(strings.TrimSpace(host), "."))
addresses, systemErr := resolver.LookupIPAddr(ctx, host) hostsToTry := []string{normalized}
validSystemAddresses := filterValidPublicEPDGAddresses(addresses) if alt := alternate3GPPHostname(normalized); alt != "" && alt != normalized {
if systemErr == nil && len(validSystemAddresses) > 0 { hostsToTry = append(hostsToTry, alt)
return validSystemAddresses, nil }
var systemErr error
for _, targetHost := range hostsToTry {
addresses, err := resolver.LookupIPAddr(ctx, targetHost)
if err == nil {
valid := filterValidPublicEPDGAddresses(addresses)
if len(valid) > 0 {
return valid, nil
}
} else {
systemErr = err
}
} }
subnet := vowifi.EPDGDNSClientSubnet(normalized) subnet := vowifi.EPDGDNSClientSubnet(normalized)
client := &http.Client{Timeout: 8 * time.Second} client := &http.Client{Timeout: 8 * time.Second}
var fallbackErr error var fallbackErr error
hostsToTry := []string{normalized}
if alt := alternate3GPPHostname(normalized); alt != "" && alt != normalized {
hostsToTry = append(hostsToTry, alt)
}
for _, targetHost := range hostsToTry { for _, targetHost := range hostsToTry {
var fallback []net.IPAddr var fallback []net.IPAddr
fallback, fallbackErr = resolveEPDGWithECS(ctx, client, googleDNSOverHTTPS, targetHost, subnet) fallback, fallbackErr = resolveEPDGWithECS(ctx, client, googleDNSOverHTTPS, targetHost, subnet)
+103 -32
View File
@@ -187,6 +187,7 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
{Type: payloadNonce, Body: initiatorNonce}, {Type: payloadNonce, Body: initiatorNonce},
makeNotify(notifyNATSource, sourceHash), makeNotify(notifyNATSource, sourceHash),
makeNotify(notifyNATDestination, destinationHash), makeNotify(notifyNATDestination, destinationHash),
makeNotify(notifyFragmentationSupported, nil),
} }
var ( var (
initRequest []byte initRequest []byte
@@ -243,6 +244,7 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
} }
break break
} }
peerSupportsFragmentation := hasNotifyType(initResponsePayloads, notifyFragmentationSupported)
saPayload, err := onePayload(initResponsePayloads, payloadSA) saPayload, err := onePayload(initResponsePayloads, payloadSA)
if err != nil { if err != nil {
return nil, err return nil, err
@@ -345,21 +347,19 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
Flags: flagInitiator, Flags: flagInitiator,
MessageID: 1, MessageID: 1,
} }
authRequest, err := encryptPayloads(authHeader, firstAuthPayloads, ikeSuite, keys.SKei, keys.SKai, provider.config.Random) _, authResponsePayloads, err := sendAndReceiveIKEPayloads(
if err != nil { ctx,
return nil, err transport,
} authHeader,
authResponse, err := transport.RoundTrip(ctx, authRequest) firstAuthPayloads,
if err != nil { ikeSuite,
return nil, err keys,
} peerSupportsFragmentation,
authResponseHeader, authResponsePayloads, err := decryptAndValidate( provider.config.Random,
authResponse, initiatorSPI, responseHeader.ResponderSPI, exchangeIKEAuth, 1, ikeSuite, keys,
) )
if err != nil { if err != nil {
return nil, err return nil, err
} }
_ = authResponseHeader
serverName := strings.TrimSpace(provider.config.ServerName) serverName := strings.TrimSpace(provider.config.ServerName)
if serverName == "" { if serverName == "" {
serverName = epdg serverName = epdg
@@ -409,27 +409,27 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
requestPayloads = append(requestPayloads, deviceIdentity) requestPayloads = append(requestPayloads, deviceIdentity)
} }
} }
eapRequest, err := encryptPayloads(ikeHeader{ eapHeader := ikeHeader{
InitiatorSPI: initiatorSPI, InitiatorSPI: initiatorSPI,
ResponderSPI: responseHeader.ResponderSPI, ResponderSPI: responseHeader.ResponderSPI,
Exchange: exchangeIKEAuth, Exchange: exchangeIKEAuth,
Flags: flagInitiator, Flags: flagInitiator,
MessageID: messageID, MessageID: messageID,
}, requestPayloads, ikeSuite, keys.SKei, keys.SKai, provider.config.Random)
if err != nil {
return nil, err
} }
if requested, notifyErr := deviceIdentityRequested(currentPayloads); notifyErr != nil { if requested, notifyErr := deviceIdentityRequested(currentPayloads); notifyErr != nil {
return nil, notifyErr return nil, notifyErr
} else if requested { } else if requested {
deviceIdentityPending = true deviceIdentityPending = true
} }
eapResponse, err := transport.RoundTrip(ctx, eapRequest) _, currentPayloads, err = sendAndReceiveIKEPayloads(
if err != nil { ctx,
return nil, err transport,
} eapHeader,
_, currentPayloads, err = decryptAndValidate( requestPayloads,
eapResponse, initiatorSPI, responseHeader.ResponderSPI, exchangeIKEAuth, messageID, ikeSuite, keys, ikeSuite,
keys,
peerSupportsFragmentation,
provider.config.Random,
) )
if err != nil { if err != nil {
return nil, err return nil, err
@@ -454,22 +454,22 @@ func (provider *Provider) start(ctx context.Context, request vowifi.TunnelReques
} }
messageID++ messageID++
cleanupMessageID = messageID + 1 cleanupMessageID = messageID + 1
finalRequest, err := encryptPayloads(ikeHeader{ finalHeader := ikeHeader{
InitiatorSPI: initiatorSPI, InitiatorSPI: initiatorSPI,
ResponderSPI: responseHeader.ResponderSPI, ResponderSPI: responseHeader.ResponderSPI,
Exchange: exchangeIKEAuth, Exchange: exchangeIKEAuth,
Flags: flagInitiator, Flags: flagInitiator,
MessageID: messageID, MessageID: messageID,
}, []payload{initiatorAUTH}, ikeSuite, keys.SKei, keys.SKai, provider.config.Random)
if err != nil {
return nil, err
} }
finalResponse, err := transport.RoundTrip(ctx, finalRequest) _, finalPayloads, err := sendAndReceiveIKEPayloads(
if err != nil { ctx,
return nil, err transport,
} finalHeader,
_, finalPayloads, err := decryptAndValidate( []payload{initiatorAUTH},
finalResponse, initiatorSPI, responseHeader.ResponderSPI, exchangeIKEAuth, messageID, ikeSuite, keys, ikeSuite,
keys,
peerSupportsFragmentation,
provider.config.Random,
) )
if err != nil { if err != nil {
return nil, err return nil, err
@@ -763,7 +763,7 @@ func decryptAndValidate(
suite negotiatedSuite, suite negotiatedSuite,
keys ikeKeys, keys ikeKeys,
) (ikeHeader, []payload, error) { ) (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 { if err != nil {
return ikeHeader{}, nil, err return ikeHeader{}, nil, err
} }
@@ -778,6 +778,77 @@ func decryptAndValidate(
return header, payloads, nil 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") var errNoProposalChosen = errors.New("ike: responder reported NO_PROPOSAL_CHOSEN")
type invalidKEPayloadError struct { type invalidKEPayloadError struct {
+14 -3
View File
@@ -76,7 +76,7 @@ func (transport *firstAuthCaptureTransport) Float(context.Context) error {
return nil 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++ transport.calls++
if len(transport.cookieChallenge) > 0 { if len(transport.cookieChallenge) > 0 {
switch transport.calls { 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) { func (transport *firstAuthCaptureTransport) answerIKECookie(packet []byte) ([]byte, error) {
header, _, err := parseIKEPacket(packet) header, _, err := parseIKEPacket(packet)
if err != nil { 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") return errors.New("test: first retried IKE_SA_INIT payload is not the expected COOKIE")
} }
cookies := payloadsOfType(payloads, payloadNotify) cookies := payloadsOfType(payloads, payloadNotify)
if len(cookies) != 3 { if len(cookies) != 4 {
return fmt.Errorf("test: retried IKE_SA_INIT has %d notify payloads, want 3", len(cookies)) return fmt.Errorf("test: retried IKE_SA_INIT has %d notify payloads, want 4", len(cookies))
} }
found := false found := false
for _, item := range cookies { for _, item := range cookies {
+1 -1
View File
@@ -39,7 +39,7 @@ func newSessionRelay(
keepalive time.Duration, keepalive time.Duration,
) *sessionRelay { ) *sessionRelay {
if keepalive <= 0 { if keepalive <= 0 {
keepalive = 20 * time.Second keepalive = 15 * time.Second
} }
ctx, cancel := context.WithCancel(context.Background()) ctx, cancel := context.WithCancel(context.Background())
relay := &sessionRelay{ relay := &sessionRelay{
+11
View File
@@ -3,6 +3,7 @@ package ike
import ( import (
"bytes" "bytes"
"context" "context"
"errors"
"net" "net"
"sync" "sync"
"sync/atomic" "sync/atomic"
@@ -49,6 +50,16 @@ func (*fakeSessionTransport) Float(context.Context) error { return nil }
func (*fakeSessionTransport) RoundTrip(context.Context, []byte) ([]byte, error) { func (*fakeSessionTransport) RoundTrip(context.Context, []byte) ([]byte, error) {
return nil, context.DeadlineExceeded 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 { func (transport *fakeSessionTransport) SendESP(ctx context.Context, packet []byte) error {
return transport.SendSessionPacket(ctx, packet, false) return transport.SendSessionPacket(ctx, packet, false)
} }
+127 -43
View File
@@ -20,6 +20,7 @@ type datagramTransport interface {
RemoteAddr() *net.UDPAddr RemoteAddr() *net.UDPAddr
Float(context.Context) error Float(context.Context) error
RoundTrip(context.Context, []byte) ([]byte, error) RoundTrip(context.Context, []byte) ([]byte, error)
RoundTripExchange(context.Context, [][]byte) ([][]byte, error)
SendESP(context.Context, []byte) error SendESP(context.Context, []byte) error
ReceiveESP(context.Context, []byte) (int, error) ReceiveESP(context.Context, []byte) (int, error)
SendSessionPacket(context.Context, []byte, bool) error SendSessionPacket(context.Context, []byte, bool) error
@@ -96,6 +97,29 @@ func roundTripDatagram(
read func([]byte, time.Time) (int, error), read func([]byte, time.Time) (int, error),
packet []byte, packet []byte,
) ([]byte, error) { ) ([]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 { if ctx == nil {
ctx = context.Background() ctx = context.Background()
} }
@@ -110,20 +134,44 @@ func roundTripDatagram(
if err := ctx.Err(); err != nil { if err := ctx.Err(); err != nil {
return nil, err return nil, err
} }
if err := write(packet); err != nil { if err := writeAll(packets); err != nil {
return nil, err return nil, err
} }
attemptDeadline := time.Now().Add(interval) attemptDeadline := time.Now().Add(interval)
if deadline.Before(attemptDeadline) { if deadline.Before(attemptDeadline) {
attemptDeadline = deadline attemptDeadline = deadline
} }
var (
totalExpected uint16
fragments = make(map[uint16][]byte)
)
for time.Now().Before(attemptDeadline) { for time.Now().Before(attemptDeadline) {
if err := ctx.Err(); err != nil { if err := ctx.Err(); err != nil {
return nil, err return nil, err
} }
n, err := read(buffer, attemptDeadline) n, err := read(buffer, attemptDeadline)
if err == nil { 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() { if timeoutError, ok := err.(net.Error); ok && timeoutError.Timeout() {
lastErr = err 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) { 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() transport.mu.Lock()
defer transport.mu.Unlock() defer transport.mu.Unlock()
if transport.conn == nil { if transport.conn == nil {
return nil, errors.New("ike: UDP transport is closed") 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 { if err != nil {
return nil, fmt.Errorf("ike: invalid outbound packet: %w", err) return nil, fmt.Errorf("ike: invalid outbound packet: %w", err)
} }
wirePacket := packet var wirePackets [][]byte
if transport.floated { for _, pkt := range packets {
wirePacket = append([]byte{0, 0, 0, 0}, packet...) wire := pkt
} if transport.floated {
write := func(value []byte) error { wire = append([]byte{0, 0, 0, 0}, pkt...)
if err := transport.conn.SetWriteDeadline(deadlineFor(ctx, transport.config.Timeout)); err != nil {
return err
} }
_, err := transport.conn.Write(value) wirePackets = append(wirePackets, wire)
return err }
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) { read := func(buffer []byte, attemptDeadline time.Time) (int, error) {
for { for {
@@ -252,10 +322,6 @@ func (transport *directUDP) RoundTrip(ctx context.Context, packet []byte) ([]byt
return 0, err return 0, err
} }
if transport.floated { 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]) { if !hasNonESPMarker(buffer[:n]) {
continue continue
} }
@@ -268,7 +334,7 @@ func (transport *directUDP) RoundTrip(ctx context.Context, packet []byte) ([]byt
return n, nil 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 { 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) { 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() transport.mu.Lock()
defer transport.mu.Unlock() defer transport.mu.Unlock()
if transport.udp == nil { if transport.udp == nil {
return nil, errors.New("ike: SOCKS5 UDP transport is closed") 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 { if err != nil {
return nil, fmt.Errorf("ike: invalid outbound packet: %w", err) 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. // 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 { if !transport.floated && requestHeader.Exchange == exchangeIKEInit && requestHeader.MessageID == 0 && len(transport.remotes) > 1 {
var lastErr error var lastErr error
var cookieResponse []byte var cookieResponse [][]byte
for _, candidate := range transport.remotes { for _, candidate := range transport.remotes {
transport.remote = cloneUDPAddr(candidate) transport.remote = cloneUDPAddr(candidate)
response, attemptErr := transport.roundTripLocked(ctx, packet, requestHeader) responses, attemptErr := transport.roundTripFragmentsLocked(ctx, packets, requestHeader)
if attemptErr == nil { if attemptErr == nil {
if ikeInitResponseHasCookie(response) { if len(responses) > 0 && ikeInitResponseHasCookie(responses[0]) {
if cookieResponse == nil { if cookieResponse == nil {
cookieResponse = append([]byte(nil), response...) cookieResponse = responses
} }
continue continue
} }
return response, nil return responses, nil
} }
lastErr = attemptErr lastErr = attemptErr
if ctx.Err() != nil || !isNetworkTimeout(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 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) { func (transport *socks5UDP) roundTripFragmentsLocked(ctx context.Context, packets [][]byte, requestHeader ikeHeader) ([][]byte, error) {
wireIKE := packet var datagrams [][]byte
if transport.floated { for _, pkt := range packets {
wireIKE = append([]byte{0, 0, 0, 0}, packet...) wireIKE := pkt
} if transport.floated {
datagram, err := marshalSOCKS5Datagram(transport.remote, wireIKE) wireIKE = append([]byte{0, 0, 0, 0}, pkt...)
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
} }
_, err := transport.udp.Write(value) datagram, err := marshalSOCKS5Datagram(transport.remote, wireIKE)
return err 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) { read := func(buffer []byte, attemptDeadline time.Time) (int, error) {
for { for {
@@ -630,10 +718,6 @@ func (transport *socks5UDP) roundTripLocked(ctx context.Context, packet []byte,
return 0, err return 0, err
} }
if transport.floated { 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) { if !hasNonESPMarker(payload) {
continue continue
} }
@@ -646,7 +730,7 @@ func (transport *socks5UDP) roundTripLocked(ctx context.Context, packet []byte,
return len(payload), nil 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 { func isNetworkTimeout(err error) bool {
+14 -12
View File
@@ -31,9 +31,10 @@ const (
payloadDelete = 42 payloadDelete = 42
payloadTSi = 44 payloadTSi = 44
payloadTSr = 45 payloadTSr = 45
payloadEncrypted = 46 payloadEncrypted = 46
payloadCP = 47 payloadCP = 47
payloadEAP = 48 payloadEAP = 48
payloadEncryptedFragment = 53
protocolIKE = 1 protocolIKE = 1
protocolESP = 3 protocolESP = 3
@@ -55,15 +56,16 @@ const (
dhMODP2048 = 14 dhMODP2048 = 14
transformAttributeKeyLen = 14 transformAttributeKeyLen = 14
notifyInitialContact = 16384 notifyInitialContact = 16384
notifyMOBIKESupported = 16396 notifyMOBIKESupported = 16396
notifyNATSource = 16388 notifyNATSource = 16388
notifyNATDestination = 16389 notifyNATDestination = 16389
notifyCookie = 16390 notifyCookie = 16390
notifyEAPOnlyAuth = 16417 notifyEAPOnlyAuth = 16417
notifyDeviceIdentity = 41101 notifyFragmentationSupported = 16430
notifyInvalidKE = 17 notifyDeviceIdentity = 41101
notifyNoProposal = 14 notifyInvalidKE = 17
notifyNoProposal = 14
) )
var ( var (
+97 -82
View File
@@ -230,99 +230,109 @@ func (provider *Provider) Start(ctx context.Context, request vowifi.IMSRequest)
if err != nil { if err != nil {
return nil, err return nil, err
} }
pcscf := provider.config.PCSCF var pcscfCandidates []string
if pcscf == "" { if provider.config.PCSCF != "" {
pcscfCandidates = []string{provider.config.PCSCF}
} else {
for _, candidate := range tunnel.PCSCF { for _, candidate := range tunnel.PCSCF {
if strings.TrimSpace(candidate) != "" { candidate = strings.TrimSpace(candidate)
pcscf = candidate if candidate != "" {
break pcscfCandidates = append(pcscfCandidates, candidate)
} }
} }
} }
if pcscf == "" { if len(pcscfCandidates) == 0 {
return nil, errors.New("ims: tunnel did not provide a P-CSCF") return nil, errors.New("ims: tunnel did not provide a P-CSCF")
} }
endpoint, transportHint, err := parsePCSCF(pcscf, provider.config.Port)
if err != nil {
return nil, err
}
if provider.config.PCSCF != "" && !pcscfProvenByTunnel(endpoint, tunnel.PCSCF, provider.config.Port) {
return nil, errors.New("ims: configured P-CSCF is not proven by the SWu tunnel")
}
transport, carrierSelected := carrierTransportForIdentity(provider.config, request.Identity)
if cached := provider.cachedTransport(request.Identity); cached != "" {
transport = cached
carrierSelected = true
}
if transport == "" && !carrierSelected {
transport = transportHint
}
if transport == "" {
transport = provider.config.Transport
}
if transport == "" {
transport = "tcp"
}
localAddress := provider.config.LocalAddress
if localAddress == "" {
if endpointIP := net.ParseIP(endpoint.host); endpointIP != nil && endpointIP.To4() == nil {
localAddress = tunnel.LocalIPv6
} else {
localAddress = tunnel.LocalIPv4
if strings.TrimSpace(localAddress) == "" {
localAddress = tunnel.LocalIPv6
}
}
}
localAddress = strings.TrimSpace(strings.Split(localAddress, "/")[0])
if localAddress == "" {
return nil, errors.New("ims: tunnel did not provide a local address")
}
if !localAddressProvenByTunnel(localAddress, tunnel) {
return nil, errors.New("ims: configured local address is not assigned by the SWu tunnel")
}
transports := []string{transport}
if provider.config.AutoTransportFallback {
alternate := "udp"
if transport == "udp" {
alternate = "tcp"
}
transports = append(transports, alternate)
}
var lastErr error var lastErr error
for attempt, candidate := range transports { for pcscfIndex, pcscf := range pcscfCandidates {
connection, dialErr := dialSIP(ctx, candidate, localAddress, 0, endpoint.address()) endpoint, transportHint, err := parsePCSCF(pcscf, provider.config.Port)
if dialErr != nil { if err != nil {
lastErr = fmt.Errorf("ims: connect to P-CSCF over %s: %w", candidate, dialErr) lastErr = err
if attempt+1 < len(transports) && ctx.Err() == nil { continue
provider.logTransportFallback(request.Identity, candidate, transports[attempt+1], lastErr) }
continue if provider.config.PCSCF != "" && !pcscfProvenByTunnel(endpoint, tunnel.PCSCF, provider.config.Port) {
return nil, errors.New("ims: configured P-CSCF is not proven by the SWu tunnel")
}
transport, carrierSelected := carrierTransportForIdentity(provider.config, request.Identity)
if cached := provider.cachedTransport(request.Identity); cached != "" {
transport = cached
carrierSelected = true
}
if transport == "" && !carrierSelected {
transport = transportHint
}
if transport == "" {
transport = provider.config.Transport
}
if transport == "" {
transport = "tcp"
}
localAddress := provider.config.LocalAddress
if localAddress == "" {
if endpointIP := net.ParseIP(endpoint.host); endpointIP != nil && endpointIP.To4() == nil {
localAddress = tunnel.LocalIPv6
} else {
localAddress = tunnel.LocalIPv4
if strings.TrimSpace(localAddress) == "" {
localAddress = tunnel.LocalIPv6
}
} }
return nil, lastErr
} }
session, sessionErr := newSession(provider, request, identities, endpoint, candidate, connection) localAddress = strings.TrimSpace(strings.Split(localAddress, "/")[0])
if sessionErr != nil { if localAddress == "" {
_ = connection.Close() return nil, errors.New("ims: tunnel did not provide a local address")
return nil, sessionErr
} }
establishErr := session.establish(ctx) if !localAddressProvenByTunnel(localAddress, tunnel) {
if establishErr == nil { return nil, errors.New("ims: configured local address is not assigned by the SWu tunnel")
provider.rememberTransport(request.Identity, candidate) }
if attempt > 0 {
provider.config.Logger.Info("IMS automatic transport fallback succeeded", transports := []string{transport}
"carrier_profile", vowifi.ResolveCarrierProfile(request.Identity).ID, if provider.config.AutoTransportFallback {
"transport", candidate) alternate := "udp"
if transport == "udp" {
alternate = "tcp"
} }
return session, nil transports = append(transports, alternate)
} }
sipResponseObserved := session.evidence.LastSIPCode != 0 for attempt, candidate := range transports {
session.abort() connection, dialErr := dialSIP(ctx, candidate, localAddress, 0, endpoint.address())
lastErr = establishErr if dialErr != nil {
if sipResponseObserved || attempt+1 >= len(transports) || ctx.Err() != nil { lastErr = fmt.Errorf("ims: connect to P-CSCF over %s: %w", candidate, dialErr)
return nil, lastErr if attempt+1 < len(transports) && ctx.Err() == nil {
provider.logTransportFallback(request.Identity, candidate, transports[attempt+1], lastErr)
continue
}
break
}
session, sessionErr := newSession(provider, request, identities, endpoint, candidate, connection)
if sessionErr != nil {
_ = connection.Close()
lastErr = sessionErr
break
}
establishErr := session.establish(ctx)
if establishErr == nil {
provider.rememberTransport(request.Identity, candidate)
if attempt > 0 || pcscfIndex > 0 {
provider.config.Logger.Info("IMS automatic transport fallback succeeded",
"carrier_profile", vowifi.ResolveCarrierProfile(request.Identity).ID,
"transport", candidate)
}
return session, nil
}
sipResponseObserved := session.evidence.LastSIPCode != 0
session.abort()
lastErr = establishErr
if sipResponseObserved || attempt+1 >= len(transports) || ctx.Err() != nil {
break
}
provider.logTransportFallback(request.Identity, candidate, transports[attempt+1], establishErr)
}
if ctx.Err() != nil {
return nil, ctx.Err()
} }
provider.logTransportFallback(request.Identity, candidate, transports[attempt+1], establishErr)
} }
return nil, lastErr return nil, lastErr
} }
@@ -384,8 +394,13 @@ func deriveIdentities(identity vowifi.SIMIdentity, config Config) (identitySet,
if !digitsBetween(imsi, 5, 16) { if !digitsBetween(imsi, 5, 16) {
return identitySet{}, errors.New("ims: SIM IMSI is unavailable or invalid") return identitySet{}, errors.New("ims: SIM IMSI is unavailable or invalid")
} }
mcc := strings.TrimSpace(identity.HomeMCC) profile := vowifi.ResolveCarrierProfile(identity)
mnc := strings.TrimSpace(identity.HomeMNC) mcc := strings.TrimSpace(profile.RouteMCC)
mnc := strings.TrimSpace(profile.RouteMNC)
if mcc == "" || mnc == "" {
mcc = strings.TrimSpace(identity.HomeMCC)
mnc = strings.TrimSpace(identity.HomeMNC)
}
if !digitsBetween(mcc, 3, 3) || !digitsBetween(mnc, 2, 3) { if !digitsBetween(mcc, 3, 3) || !digitsBetween(mnc, 2, 3) {
return identitySet{}, errors.New("ims: home PLMN is unavailable or invalid") return identitySet{}, errors.New("ims: home PLMN is unavailable or invalid")
} }
@@ -395,7 +410,7 @@ func deriveIdentities(identity vowifi.SIMIdentity, config Config) (identitySet,
domain := fmt.Sprintf("ims.mnc%s.mcc%s.3gppnetwork.org", mnc, mcc) domain := fmt.Sprintf("ims.mnc%s.mcc%s.3gppnetwork.org", mnc, mcc)
privateDomain := domain privateDomain := domain
publicDomain := domain publicDomain := domain
if vowifi.ResolveCarrierProfile(identity).IMSIdentityProfile == vowifi.IMSProfileATT { if profile.IMSIdentityProfile == vowifi.IMSProfileATT {
// AT&T provisions the IMPI and IMPU in its ISIM domains rather than // AT&T provisions the IMPI and IMPU in its ISIM domains rather than
// the generic 3GPP PLMN IMS domain. // the generic 3GPP PLMN IMS domain.
domain = "one.att.net" domain = "one.att.net"
+36 -5
View File
@@ -71,7 +71,16 @@ func (media *rtpMedia) ready() bool {
} }
func (media *rtpMedia) offerSDP(local net.IP) []byte { func (media *rtpMedia) offerSDP(local net.IP) []byte {
return media.buildSDP(local, "8 0", nil) return media.buildSDP(local, "8 0 104 102 100", []string{
"a=rtpmap:8 PCMA/8000",
"a=rtpmap:0 PCMU/8000",
"a=rtpmap:104 AMR-WB/16000",
"a=fmtp:104 mode-change-capability=2;max-red=220",
"a=rtpmap:102 AMR/8000",
"a=fmtp:102 mode-change-capability=2;max-red=220",
"a=rtpmap:100 telephone-event/8000",
"a=fmtp:100 0-15",
})
} }
func (media *rtpMedia) answerSDP(local net.IP) []byte { func (media *rtpMedia) answerSDP(local net.IP) []byte {
@@ -81,8 +90,12 @@ func (media *rtpMedia) answerSDP(local net.IP) []byte {
if codec == "" { if codec == "" {
return media.offerSDP(local) return media.offerSDP(local)
} }
rate := 8000
if codec == "AMR-WB" {
rate = 16000
}
return media.buildSDP(local, strconv.Itoa(int(payload)), []string{ return media.buildSDP(local, strconv.Itoa(int(payload)), []string{
fmt.Sprintf("a=rtpmap:%d %s/8000", payload, codec), fmt.Sprintf("a=rtpmap:%d %s/%d", payload, codec, rate),
}) })
} }
@@ -110,7 +123,16 @@ func (media *rtpMedia) buildSDP(local net.IP, formats string, attributes []strin
fmt.Sprintf("m=audio %d RTP/AVP %s", port, formats), fmt.Sprintf("m=audio %d RTP/AVP %s", port, formats),
} }
if attributes == nil { if attributes == nil {
lines = append(lines, "a=rtpmap:8 PCMA/8000", "a=rtpmap:0 PCMU/8000") lines = append(lines,
"a=rtpmap:8 PCMA/8000",
"a=rtpmap:0 PCMU/8000",
"a=rtpmap:104 AMR-WB/16000",
"a=fmtp:104 mode-change-capability=2;max-red=220",
"a=rtpmap:102 AMR/8000",
"a=fmtp:102 mode-change-capability=2;max-red=220",
"a=rtpmap:100 telephone-event/8000",
"a=fmtp:100 0-15",
)
} else { } else {
lines = append(lines, attributes...) lines = append(lines, attributes...)
} }
@@ -137,15 +159,24 @@ func (media *rtpMedia) configureRemote(body []byte) error {
name = "PCMU" name = "PCMU"
case 8: case 8:
name = "PCMA" name = "PCMA"
case 100:
continue
default:
name = fmt.Sprintf("PAYLOAD-%d", parsed)
} }
} }
if name == "PCMA" || name == "PCMU" { if name != "TELEPHONE-EVENT" {
codec, payload = name, byte(parsed) codec, payload = name, byte(parsed)
break break
} }
} }
if codec == "" && len(formats) > 0 {
if parsed, parseErr := strconv.Atoi(formats[0]); parseErr == nil {
codec, payload = fmt.Sprintf("PAYLOAD-%d", parsed), byte(parsed)
}
}
if codec == "" { if codec == "" {
return errors.New("ims: remote endpoint did not accept PCMA or PCMU audio") return errors.New("ims: remote SDP has no usable audio format")
} }
media.mu.Lock() media.mu.Lock()
media.remote = &net.UDPAddr{IP: address, Port: port} media.remote = &net.UDPAddr{IP: address, Port: port}
+2 -5
View File
@@ -56,7 +56,6 @@ type ReceivedSMS struct {
RawTPDU string RawTPDU string
DecodeError string DecodeError string
} }
// ReceivedSMSStatus is network delivery evidence for one submitted SMS part. // ReceivedSMSStatus is network delivery evidence for one submitted SMS part.
type ReceivedSMSStatus struct { type ReceivedSMSStatus struct {
DeviceID string DeviceID string
@@ -891,10 +890,8 @@ func (session *Session) parseUSSIReply(response *sipResponse) (string, *int) {
} }
func (session *Session) ussiTarget() string { func (session *Session) ussiTarget() string {
if number, _, ok := vowifi.ExtractAssociatedMSISDN(session.evidence); ok { if domain := strings.TrimSpace(session.identity.domain); domain != "" {
if normalized := normalizeE164(number); normalized != "" { return "sip:" + domain
return "tel:" + normalized
}
} }
return session.identity.public return session.identity.public
} }