Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ require (
github.com/casbin/casbin/v2 v2.127.0
github.com/convox/convox v0.0.0-20251023182947-1ddac03d0705
github.com/convox/stdcli v0.0.0-20240813092220-8beeb2dc2420
github.com/convox/stdsdk v0.0.3
github.com/coreos/go-oidc/v3 v3.15.0
github.com/fxamacker/cbor/v2 v2.9.0
github.com/getsentry/sentry-go v0.29.0
Expand Down Expand Up @@ -108,7 +109,6 @@ require (
github.com/convox/inotify v0.0.0-20170313035821-b56f5149b5c6 // indirect
github.com/convox/logger v0.0.0-20180522214415-e39179955b52 // indirect
github.com/convox/stdapi v1.1.3-0.20221110171947-8d98f61e61ed // indirect
github.com/convox/stdsdk v0.0.3 // indirect
github.com/convox/version v0.0.0-20160822184233-ffefa0d565d2 // indirect
github.com/cpuguy83/go-md2man/v2 v2.0.6 // indirect
github.com/creack/pty v1.1.18 // indirect
Expand Down
32 changes: 25 additions & 7 deletions internal/cli/common.go
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
package cli

import (
"fmt"
"net/url"
"os"
"reflect"
"strings"
Expand All @@ -21,15 +21,29 @@ import (
// - No MFA: abc123def456...
// - TOTP: abc123def456....totp.123456
// - WebAuthn: abc123def456....webauthn.base64_assertion
//
// The auth is percent-encoded so inline WebAuthn data (standard base64 with '/', '+', '=') can't break URL
// parsing, which would otherwise fail with an error message containing the session token.
func buildRackURL(gatewayURL, auth string) string {
// Add /api/v1/rack-proxy prefix to the gateway URL
base := strings.TrimSuffix(gatewayURL, "/") + "/api/v1/rack-proxy"
scheme := "https"
host := strings.TrimSuffix(gatewayURL, "/")
if rest, ok := strings.CutPrefix(host, "http://"); ok {
scheme, host = "http", rest
} else {
host = strings.TrimPrefix(host, "https://")
}
host, basePath, _ := strings.Cut(host, "/")
if basePath != "" {
basePath = "/" + basePath
}

// Inject auth as basic auth password
if strings.HasPrefix(base, "http://") {
return fmt.Sprintf("http://convox:%s@%s", auth, strings.TrimPrefix(base, "http://"))
u := url.URL{
Scheme: scheme,
User: url.UserPassword("convox", auth),
Host: host,
Path: basePath + "/api/v1/rack-proxy",
}
return fmt.Sprintf("https://convox:%s@%s", auth, strings.TrimPrefix(base, "https://"))
return u.String()
}

// Global flags that should NEVER be forwarded to the Convox SDK
Expand Down Expand Up @@ -63,7 +77,11 @@ func SetupConvoxCommandWithMFA(
if err != nil {
return nil, nil, err
}
if err := requireSecureGatewayURL(gatewayURL); err != nil {
return nil, nil, err
}

secureConvoxSDK()
client, err := sdk.New(buildRackURL(gatewayURL, auth))
if err != nil {
return nil, nil, err
Expand Down
47 changes: 4 additions & 43 deletions internal/cli/mfa_helpers.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,6 @@ import (
"encoding/base64"
"encoding/json"
"fmt"
"net/http"
"os"
"strings"
"time"
Expand Down Expand Up @@ -261,55 +260,17 @@ func CollectMFAAuthWithPIN(
// collectWebAuthnAssertionWithPIN collects a WebAuthn assertion, optionally using a cached PIN.
// Returns the assertion data, the PIN used (for caching), and any error.
func collectWebAuthnAssertionWithPIN(baseURL, bearer, cachedPIN string) (string, string, error) {
endpoint := fmt.Sprintf("%s/api/v1/auth/mfa/webauthn/assertion/start", baseURL)
req, err := http.NewRequest(http.MethodPost, endpoint, http.NoBody)
start, err := startWebAuthnAssertion(baseURL, bearer)
if err != nil {
return "", "", err
}
req.Header.Set("Authorization", "Bearer "+bearer)

resp, err := HTTPClient.Do(req)
options, err := buildAssertionOptions(baseURL, extractAllowedCredentialIDs(start), start)
if err != nil {
return "", "", err
}
defer func() { _ = resp.Body.Close() }()

if resp.StatusCode != http.StatusOK {
return "", "", fmt.Errorf("failed to start WebAuthn assertion")
}

var startResp struct {
Options struct {
PublicKey struct {
Challenge string `json:"challenge"`
RPID string `json:"rpId"`
AllowCredentials []struct {
ID string `json:"id"`
} `json:"allowCredentials"`
Timeout int `json:"timeout"`
UserVerification string `json:"userVerification"`
} `json:"publicKey"`
} `json:"options"`
SessionData string `json:"session_data"`
}

if err := json.NewDecoder(resp.Body).Decode(&startResp); err != nil {
return "", "", err
}

allowedCreds := make([]string, 0, len(startResp.Options.PublicKey.AllowCredentials))
for _, cred := range startResp.Options.PublicKey.AllowCredentials {
allowedCreds = append(allowedCreds, cred.ID)
}

assertion, pinUsed, err := webauthn.GetAssertionWithCachedPIN(webauthn.AssertionOptions{
Challenge: startResp.Options.PublicKey.Challenge,
RPID: startResp.Options.PublicKey.RPID,
AllowCredentials: allowedCreds,
Timeout: startResp.Options.PublicKey.Timeout,
UserVerification: startResp.Options.PublicKey.UserVerification,
Origin: baseURL,
}, cachedPIN)
assertion, pinUsed, err := webauthn.GetAssertionWithCachedPIN(options, cachedPIN)
if err != nil {
return "", "", err
}
Expand All @@ -320,7 +281,7 @@ func collectWebAuthnAssertionWithPIN(baseURL, bearer, cachedPIN string) (string,
}

inlineData := map[string]any{
"session_data": startResp.SessionData,
"session_data": start.SessionData,
"assertion_response": assertionJSON,
}

Expand Down
50 changes: 50 additions & 0 deletions internal/cli/mfa_rpid_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,50 @@
package cli

import (
"testing"

"github.com/stretchr/testify/require"
)

func TestResolveRPID(t *testing.T) {
cases := []struct {
host, server, want string
ok bool
}{
{"gateway-us.example.ts.net", "", "gateway-us.example.ts.net", true},
{"gateway-us.example.ts.net", "gateway-us.example.ts.net", "gateway-us.example.ts.net", true},
{"Gateway.Example.com", "gateway.example.com", "gateway.example.com", true},
{"gateway.example.com", "example.com", "example.com", true},
{"gateway.example.com", "google.com", "", false},
{"gateway.example.com", "evil-example.com", "", false},
{"gateway.example.com", "com", "", false},
{"gateway-us.example.ts.net", "gateway-eu.example.ts.net", "", false},
{"localhost", "localhost", "localhost", true},
}
for _, tc := range cases {
got, err := resolveRPID(tc.host, tc.server)
if !tc.ok {
require.Errorf(t, err, "host=%s server=%s", tc.host, tc.server)
continue
}
require.NoError(t, err)
require.Equal(t, tc.want, got)
}
}

func TestBuildAssertionOptionsUsesGatewayOriginAndRejectsForeignRPID(t *testing.T) {
start := &webAuthnStartResponse{}
start.Options.PublicKey.Challenge = "abc"
start.Options.PublicKey.RPID = "gateway-us.example.ts.net"
start.Options.PublicKey.UserVerification = "required"

opts, err := buildAssertionOptions("https://gateway-us.example.ts.net/", []string{"cred"}, start)
require.NoError(t, err)
require.Equal(t, "gateway-us.example.ts.net", opts.RPID)
require.Equal(t, "https://gateway-us.example.ts.net", opts.Origin)
require.Equal(t, "required", opts.UserVerification)

start.Options.PublicKey.RPID = "accounts.google.com"
_, err = buildAssertionOptions("https://gateway-us.example.ts.net", []string{"cred"}, start)
require.ErrorContains(t, err, "refusing")
}
23 changes: 22 additions & 1 deletion internal/cli/mfa_verify.go
Original file line number Diff line number Diff line change
Expand Up @@ -205,7 +205,10 @@ func buildAssertionOptions(
}

origin := fmt.Sprintf("%s://%s", parsedURL.Scheme, parsedURL.Host)
rpID := parsedURL.Hostname()
rpID, err := resolveRPID(parsedURL.Hostname(), start.Options.PublicKey.RPID)
if err != nil {
return webauthn.AssertionOptions{}, err
}

return webauthn.AssertionOptions{
Challenge: start.Options.PublicKey.Challenge,
Expand All @@ -217,6 +220,24 @@ func buildAssertionOptions(
}, nil
}

// resolveRPID returns the relying party ID to sign for. It is derived from the configured gateway host;
// a server-sent RP ID is only accepted if it is that host or a parent domain of it, so a malicious or
// spoofed gateway can't get an assertion (or a credential listing) for an unrelated site.
func resolveRPID(gatewayHost, serverRPID string) (string, error) {
host := strings.ToLower(strings.TrimSpace(gatewayHost))
requested := strings.ToLower(strings.TrimSpace(serverRPID))
if requested == "" || requested == host {
return host, nil
}
if strings.Contains(requested, ".") && strings.HasSuffix(host, "."+requested) {
return requested, nil
}
return "", fmt.Errorf(
"gateway asked for a security key assertion for %q, which doesn't match the gateway host %q; refusing",
serverRPID, gatewayHost,
)
}

func submitWebAuthnAssertion(baseURL, sessionToken, sessionData, assertionJSON string) error {
verifyEndpoint := fmt.Sprintf("%s/api/v1/auth/mfa/webauthn/assertion/verify", strings.TrimSuffix(baseURL, "/"))
payload := map[string]any{
Expand Down
88 changes: 88 additions & 0 deletions internal/cli/sdk_tls.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
package cli

import (
"context"
"crypto/tls"
"crypto/x509"
"fmt"
"net"
"net/http"
"net/url"
"strings"
"time"

"github.com/convox/stdsdk"
"github.com/gorilla/websocket"
)

// sdkRootCAs overrides the trusted roots for proxied Convox requests (tests only; nil = system roots).
var sdkRootCAs *x509.CertPool

// secureConvoxSDK makes the Convox SDK verify gateway TLS certificates.
//
// stdsdk ships an http.Client and websocket dialer with InsecureSkipVerify, and it re-applies the
// insecure websocket TLS config on every websocket call. Proxied commands carry the session token and
// MFA proof in Basic auth, so a man-in-the-middle with any certificate could capture them. HTTP requests
// get a verifying client; websockets get a NetDialTLSContext hook, which gorilla uses instead of
// TLSClientConfig, so stdsdk's per-call override has no effect.
func secureConvoxSDK() {
stdsdk.DefaultClient = newVerifyingSDKClient()
websocket.DefaultDialer.NetDialTLSContext = dialVerifiedTLS
}

func verifyingTLSConfig(serverName string) *tls.Config {
return &tls.Config{
MinVersion: tls.VersionTLS12,
RootCAs: sdkRootCAs,
ServerName: serverName,
}
}

func newVerifyingSDKClient() *http.Client {
transport := http.DefaultTransport.(*http.Transport).Clone()
transport.TLSClientConfig = verifyingTLSConfig("")
transport.TLSHandshakeTimeout = 10 * time.Second
transport.IdleConnTimeout = 90 * time.Second
return &http.Client{Transport: transport}
}

func dialVerifiedTLS(ctx context.Context, network, addr string) (net.Conn, error) {
host, _, err := net.SplitHostPort(addr)
if err != nil {
return nil, err
}
dialer := &tls.Dialer{
NetDialer: &net.Dialer{Timeout: 30 * time.Second, KeepAlive: 10 * time.Second},
Config: verifyingTLSConfig(host),
}
return dialer.DialContext(ctx, network, addr)
}

// requireSecureGatewayURL refuses to send credentials over plain HTTP unless the gateway is on this machine.
func requireSecureGatewayURL(gatewayURL string) error {
u, err := url.Parse(strings.TrimSpace(gatewayURL))
if err != nil {
return fmt.Errorf("invalid gateway URL %q: %w", gatewayURL, err)
}
switch u.Scheme {
case "https":
return nil
case "http":
if isLoopbackHost(u.Hostname()) {
return nil
}
return fmt.Errorf(
"refusing to send credentials to %s over plain HTTP; use an https:// gateway URL", u.Host,
)
default:
return fmt.Errorf("unsupported gateway URL scheme %q", u.Scheme)
}
}

func isLoopbackHost(host string) bool {
if strings.EqualFold(host, "localhost") {
return true
}
ip := net.ParseIP(host)
return ip != nil && ip.IsLoopback()
}
Loading
Loading