From cd8e7df7b793c8c6455f7a82527559164460fc58 Mon Sep 17 00:00:00 2001 From: Nathan Broadbent Date: Fri, 9 Oct 2026 18:02:47 +1300 Subject: [PATCH] Store WebAuthn challenges server-side; verify gateway TLS in the CLI WebAuthn (server): - Challenges used to round-trip through the client as session_data and were trusted on the way back. One captured assertion could be replayed as a step-up proof indefinitely, and a zeroed Expires skipped go-webauthn's expiry check. - Challenges now live in a webauthn_challenges table, bound to the user and the session that started the ceremony. They are single-use (DELETE ... RETURNING), expire after a few minutes, and the client only gets an opaque ID back in the same session_data field. This covers login MFA, step-up, inline WebAuthn on the proxy, and enrollment. - The credential sign counter is stored and enforced (clone detection). CLI: - Every proxied Convox command went through stdsdk's http.Client and websocket dialer, which set InsecureSkipVerify. A man-in-the-middle with any certificate got the session token and MFA proof from the Basic auth header. Both now verify certificates against the system roots, and plain http:// is refused except for loopback gateways. - The rack-proxy URL's credentials are URL-encoded, so inline WebAuthn data can't break parsing (the parse error used to echo the token). - Inline step-up WebAuthn derives the RP ID from the configured gateway host and rejects a different RP ID sent by the server. --- go.mod | 2 +- internal/cli/common.go | 32 +- internal/cli/mfa_helpers.go | 47 +-- internal/cli/mfa_rpid_test.go | 50 +++ internal/cli/mfa_verify.go | 23 +- internal/cli/sdk_tls.go | 88 +++++ internal/cli/sdk_tls_test.go | 127 +++++++ internal/gateway/auth/mfa/webauthn.go | 355 ++++-------------- .../gateway/auth/mfa/webauthn_assertion.go | 221 +++++++++++ .../auth/mfa/webauthn_challenge_test.go | 175 +++++++++ internal/gateway/auth/mfa/webauthn_test.go | 35 +- internal/gateway/auth/mfa/webauthn_user.go | 1 + internal/gateway/db/database.go | 1 + internal/gateway/db/mfa.go | 4 +- internal/gateway/db/mfa_scan.go | 2 +- .../20261009120000_webauthn_challenges.sql | 21 ++ internal/gateway/db/sessions_mutations.go | 39 -- internal/gateway/db/sessions_test.go | 250 ------------ internal/gateway/db/types.go | 1 + internal/gateway/db/webauthn_challenges.go | 139 +++++++ .../gateway/handlers/auth_mfa_enrollment.go | 90 +---- .../gateway/handlers/auth_mfa_verification.go | 12 +- internal/gateway/handlers/dto.go | 6 +- .../middleware/mfa_stepup_webauthn_test.go | 6 +- .../gateway/openapi/generated/swagger.json | 3 +- .../gateway/testutil/webauthntest/webauthn.go | 62 +-- internal/integration/integration_test.go | 6 +- 27 files changed, 1034 insertions(+), 764 deletions(-) create mode 100644 internal/cli/mfa_rpid_test.go create mode 100644 internal/cli/sdk_tls.go create mode 100644 internal/cli/sdk_tls_test.go create mode 100644 internal/gateway/auth/mfa/webauthn_assertion.go create mode 100644 internal/gateway/auth/mfa/webauthn_challenge_test.go create mode 100644 internal/gateway/db/migrations/20261009120000_webauthn_challenges.sql create mode 100644 internal/gateway/db/webauthn_challenges.go diff --git a/go.mod b/go.mod index 96315f88..b8d57015 100644 --- a/go.mod +++ b/go.mod @@ -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 @@ -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 diff --git a/internal/cli/common.go b/internal/cli/common.go index dbaaa7f8..485331f8 100644 --- a/internal/cli/common.go +++ b/internal/cli/common.go @@ -1,7 +1,7 @@ package cli import ( - "fmt" + "net/url" "os" "reflect" "strings" @@ -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 @@ -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 diff --git a/internal/cli/mfa_helpers.go b/internal/cli/mfa_helpers.go index 7c902f68..5bc290ec 100644 --- a/internal/cli/mfa_helpers.go +++ b/internal/cli/mfa_helpers.go @@ -4,7 +4,6 @@ import ( "encoding/base64" "encoding/json" "fmt" - "net/http" "os" "strings" "time" @@ -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 } @@ -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, } diff --git a/internal/cli/mfa_rpid_test.go b/internal/cli/mfa_rpid_test.go new file mode 100644 index 00000000..ee04c4c8 --- /dev/null +++ b/internal/cli/mfa_rpid_test.go @@ -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") +} diff --git a/internal/cli/mfa_verify.go b/internal/cli/mfa_verify.go index 3bb238b2..56b083a9 100644 --- a/internal/cli/mfa_verify.go +++ b/internal/cli/mfa_verify.go @@ -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, @@ -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{ diff --git a/internal/cli/sdk_tls.go b/internal/cli/sdk_tls.go new file mode 100644 index 00000000..26e601ad --- /dev/null +++ b/internal/cli/sdk_tls.go @@ -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() +} diff --git a/internal/cli/sdk_tls_test.go b/internal/cli/sdk_tls_test.go new file mode 100644 index 00000000..a4eb1b6c --- /dev/null +++ b/internal/cli/sdk_tls_test.go @@ -0,0 +1,127 @@ +package cli + +import ( + "context" + "crypto/x509" + "net/http" + "net/http/httptest" + "strings" + "sync" + "testing" + "time" + + "github.com/convox/convox/sdk" + "github.com/spf13/cobra" + "github.com/stretchr/testify/require" +) + +type authRecorder struct { + mu sync.Mutex + auth []string +} + +func (a *authRecorder) handler(w http.ResponseWriter, r *http.Request) { + a.mu.Lock() + a.auth = append(a.auth, r.Header.Get("Authorization")) + a.mu.Unlock() + w.Header().Set("Content-Type", "application/json") + _, _ = w.Write([]byte(`[]`)) +} + +func (a *authRecorder) seen() []string { + a.mu.Lock() + defer a.mu.Unlock() + return append([]string(nil), a.auth...) +} + +func configureRackForTLSTest(t *testing.T, gatewayURL string) { + t.Helper() + ConfigPath = t.TempDir() + RackFlag = "" + t.Setenv("RACK_GATEWAY_API_TOKEN", "") + t.Setenv("RACK_GATEWAY_URL", "") + require.NoError(t, SaveConfig(&Config{ + Current: "us", + Gateways: map[string]GatewayConfig{ + "us": {URL: gatewayURL, Token: "SESSION-TOKEN-SECRET", ExpiresAt: time.Now().Add(time.Hour)}, + }, + })) +} + +func TestConvoxCommandsRejectUntrustedTLSCertificate(t *testing.T) { + recorder := &authRecorder{} + srv := httptest.NewTLSServer(http.HandlerFunc(recorder.handler)) + defer srv.Close() + configureRackForTLSTest(t, srv.URL) + + client, _, err := SetupConvoxCommandWithMFA(&cobra.Command{}, nil, "totp.123456") + require.NoError(t, err) + _, err = client.AppList() + require.Error(t, err) + require.Contains(t, err.Error(), "certificate") + require.Empty(t, recorder.seen(), "credentials must not reach a server with an untrusted certificate") +} + +func TestConvoxCommandsAcceptTrustedTLSCertificate(t *testing.T) { + recorder := &authRecorder{} + srv := httptest.NewTLSServer(http.HandlerFunc(recorder.handler)) + defer srv.Close() + configureRackForTLSTest(t, srv.URL) + + pool := x509.NewCertPool() + pool.AddCert(srv.Certificate()) + sdkRootCAs = pool + t.Cleanup(func() { sdkRootCAs = nil }) + + client, _, err := SetupConvoxCommandWithMFA(&cobra.Command{}, nil, "totp.123456") + require.NoError(t, err) + _, err = client.AppList() + require.NoError(t, err) + require.Len(t, recorder.seen(), 1) +} + +func TestWebsocketTLSDialVerifiesCertificate(t *testing.T) { + srv := httptest.NewTLSServer(http.NotFoundHandler()) + defer srv.Close() + addr := strings.TrimPrefix(srv.URL, "https://") + + _, err := dialVerifiedTLS(context.Background(), "tcp", addr) + require.Error(t, err, "untrusted certificate must be rejected") + + pool := x509.NewCertPool() + pool.AddCert(srv.Certificate()) + sdkRootCAs = pool + t.Cleanup(func() { sdkRootCAs = nil }) + + conn, err := dialVerifiedTLS(context.Background(), "tcp", addr) + require.NoError(t, err) + require.NoError(t, conn.Close()) +} + +func TestRequireSecureGatewayURL(t *testing.T) { + for _, ok := range []string{ + "https://gateway.example.com", + "http://localhost:8447", + "http://127.0.0.1:9447", + "http://[::1]:8447", + } { + require.NoError(t, requireSecureGatewayURL(ok), ok) + } + for _, bad := range []string{ + "http://gateway.example.com", + "http://10.0.0.5:8447", + "ftp://gateway.example.com", + } { + require.Error(t, requireSecureGatewayURL(bad), bad) + } +} + +func TestBuildRackURLEscapesInlineMFA(t *testing.T) { + // Inline WebAuthn data is standard base64 and can contain '/', '+' and '='. + auth := "SESSIONTOKEN123.webauthn.eyJzZXNzaW9uX2RhdGEiOiJ7fSJ9/+abc==" + client, err := sdk.New(buildRackURL("https://gateway.example.com", auth)) + require.NoError(t, err) + password, _ := client.Client.Endpoint.User.Password() + require.Equal(t, auth, password) + require.Equal(t, "/api/v1/rack-proxy", client.Client.Endpoint.Path) +} diff --git a/internal/gateway/auth/mfa/webauthn.go b/internal/gateway/auth/mfa/webauthn.go index a997fe1d..227d0fdf 100644 --- a/internal/gateway/auth/mfa/webauthn.go +++ b/internal/gateway/auth/mfa/webauthn.go @@ -2,8 +2,10 @@ package mfa import ( "encoding/json" + "errors" "fmt" "strings" + "time" "github.com/go-webauthn/webauthn/protocol" "github.com/go-webauthn/webauthn/webauthn" @@ -11,69 +13,74 @@ import ( "github.com/DocSpring/rack-gateway/internal/gateway/db" ) -// StartWebAuthnEnrollment begins WebAuthn credential registration. -// Returns the challenge options and a session identifier for the frontend. -func (s *Service) StartWebAuthnEnrollment(user *db.User) (*StartWebAuthnEnrollmentResult, string, error) { +// webAuthnChallengeTTL bounds how long a WebAuthn ceremony may take. The challenge is stored +// server-side and deleted when used, so it can't be replayed. +const webAuthnChallengeTTL = 5 * time.Minute + +// ErrWebAuthnChallenge is returned when the challenge is unknown, expired, already used or belongs +// to another user or session. +var ErrWebAuthnChallenge = errors.New("WebAuthn challenge expired or already used; please try again") + +// StartWebAuthnEnrollment begins WebAuthn credential registration for the given session. +// The registration challenge is stored server-side, bound to the user and session. +func (s *Service) StartWebAuthnEnrollment(user *db.User, sessionID int64) (*StartWebAuthnEnrollmentResult, error) { if user == nil { - return nil, "", fmt.Errorf("user required") + return nil, fmt.Errorf("user required") } if s.webAuthn == nil { - return nil, "", fmt.Errorf("WebAuthn not configured") + return nil, fmt.Errorf("WebAuthn not configured") } if err := s.prepareEnrollment(user.ID); err != nil { - return nil, "", err + return nil, err } // Get existing WebAuthn credentials for exclusion methods, err := s.db.ListMFAMethods(user.ID) if err != nil { - return nil, "", err + return nil, err } waUser := &webAuthnUser{user: user, methods: methods} - options, session, err := s.webAuthn.BeginRegistration(waUser) + options, session, err := s.webAuthn.BeginRegistration( + waUser, + webauthn.WithAuthenticatorSelection(protocol.AuthenticatorSelection{ + UserVerification: protocol.VerificationRequired, + }), + ) if err != nil { - return nil, "", fmt.Errorf("failed to begin WebAuthn registration: %w", err) + return nil, fmt.Errorf("failed to begin WebAuthn registration: %w", err) } - // Generate a unique session ID for this enrollment - sessionID := fmt.Sprintf("webauthn_enroll_%d_%d", user.ID, s.now().UnixNano()) - - // Store session data - caller must persist this (typically in user session metadata) - sessionData, err := json.Marshal(session) - if err != nil { - return nil, "", fmt.Errorf("failed to marshal session: %w", err) + if _, err := s.storeChallenge(user.ID, &sessionID, db.WebAuthnChallengeRegistration, session); err != nil { + return nil, err } // Store a minimal placeholder method to get an ID - method, err := s.db.CreateMFAMethod(user.ID, "webauthn_pending", "Security Key", sessionID, nil, nil, nil, nil) + placeholder := fmt.Sprintf("webauthn_enroll_%d_%d", user.ID, s.now().UnixNano()) + method, err := s.db.CreateMFAMethod(user.ID, "webauthn_pending", "Security Key", placeholder, nil, nil, nil, nil) if err != nil { - return nil, "", err + return nil, err } backupCodes, err := s.backupCodesForEnrollment(user) if err != nil { - return nil, "", err + return nil, err } - result := &StartWebAuthnEnrollmentResult{ + return &StartWebAuthnEnrollmentResult{ MethodID: method.ID, PublicKeyOptions: options, BackupCodes: backupCodes, - } - - // Return session data so caller can store it in their session - return result, string(sessionData), nil + }, nil } -// ConfirmWebAuthnEnrollment finalizes WebAuthn registration. -// sessionDataJSON is the WebAuthn session data returned from StartWebAuthnEnrollment. -// methodID is the placeholder method ID returned from StartWebAuthnEnrollment. +// ConfirmWebAuthnEnrollment finalizes WebAuthn registration using the registration challenge stored +// for this session. methodID is the placeholder method ID returned from StartWebAuthnEnrollment. func (s *Service) ConfirmWebAuthnEnrollment( user *db.User, + sessionID int64, methodID int64, - sessionDataJSON []byte, credentialJSON []byte, label string, ) (int64, error) { @@ -88,7 +95,12 @@ func (s *Service) ConfirmWebAuthnEnrollment( return 0, err } - credential, err := s.createWebAuthnCredential(user, sessionDataJSON, credentialJSON) + challenge, err := s.db.ConsumeSessionWebAuthnChallenge(user.ID, sessionID, db.WebAuthnChallengeRegistration) + if err != nil { + return 0, challengeError(err) + } + + credential, err := s.createWebAuthnCredential(user, challenge.SessionData, credentialJSON) if err != nil { return 0, err } @@ -104,6 +116,28 @@ func (s *Service) ConfirmWebAuthnEnrollment( return methodID, nil } +// storeChallenge persists go-webauthn session data server-side and returns the opaque challenge ID. +func (s *Service) storeChallenge( + userID int64, + sessionID *int64, + purpose string, + session *webauthn.SessionData, +) (string, error) { + session.Expires = s.now().Add(webAuthnChallengeTTL) + sessionData, err := json.Marshal(session) + if err != nil { + return "", fmt.Errorf("failed to marshal WebAuthn session: %w", err) + } + return s.db.CreateWebAuthnChallenge(userID, sessionID, purpose, sessionData, webAuthnChallengeTTL) +} + +func challengeError(err error) error { + if errors.Is(err, db.ErrWebAuthnChallengeNotFound) { + return ErrWebAuthnChallenge + } + return err +} + func (s *Service) validatePendingMethod(userID, methodID int64) error { method, err := s.db.GetMFAMethodByID(methodID) if err != nil || method == nil || method.UserID != userID { @@ -166,7 +200,7 @@ func (s *Service) storeParsedCredential( return fmt.Errorf("failed to marshal metadata: %w", err) } - return s.db.UpdateMFAMethodCredential( + if err := s.db.UpdateMFAMethodCredential( methodID, "webauthn", label, @@ -174,263 +208,8 @@ func (s *Service) storeParsedCredential( credential.PublicKey, transports, metadataJSON, - ) -} - -// StartWebAuthnAssertion begins a WebAuthn assertion (login) ceremony. -// Returns the challenge options and session data that must be stored for verification. -func (s *Service) StartWebAuthnAssertion(user *db.User) (*protocol.CredentialAssertion, []byte, error) { - if user == nil { - return nil, nil, fmt.Errorf("user required") - } - if s.webAuthn == nil { - return nil, nil, fmt.Errorf("WebAuthn not configured") - } - - methods, err := s.db.ListMFAMethods(user.ID) - if err != nil { - return nil, nil, err - } - - waUser := &webAuthnUser{user: user, methods: methods} - - options, session, err := s.webAuthn.BeginLogin(waUser) - if err != nil { - return nil, nil, fmt.Errorf("failed to begin login: %w", err) - } - - sessionJSON, err := json.Marshal(session) - if err != nil { - return nil, nil, fmt.Errorf("failed to marshal session: %w", err) - } - - return options, sessionJSON, nil -} - -// VerifyWebAuthnAssertion validates a WebAuthn assertion response using stored session data. -// Includes rate limiting and automatic account locking. -func (s *Service) VerifyWebAuthnAssertion( - user *db.User, - sessionJSON []byte, - credentialJSON []byte, - ipAddress string, - userAgent string, - sessionID *int64, -) (*VerificationResult, error) { - if user == nil { - return nil, fmt.Errorf("user required") - } - if s.webAuthn == nil { - return nil, fmt.Errorf("WebAuthn not configured") - } - - if err := s.ensureUserUnlocked(user.ID); err != nil { - return nil, err - } - - if err := s.checkWebAuthnRateLimit(user.ID, ipAddress, userAgent, sessionID); err != nil { - return nil, err - } - - methods, err := s.db.ListMFAMethods(user.ID) - if err != nil { - return nil, err - } - - if result, handled, err := s.handleE2EAssertion( - user.ID, - methods, - ipAddress, - userAgent, - sessionID, - ); handled { - return result, err - } - - credential, err := s.validateAssertion( - user, - methods, - sessionJSON, - credentialJSON, - ipAddress, - userAgent, - sessionID, - ) - if err != nil { - return nil, err - } - - return s.findAndConfirmCredential(user.ID, methods, credential, ipAddress, userAgent, sessionID) -} - -func (s *Service) checkWebAuthnRateLimit( - userID int64, - ipAddress string, - userAgent string, - sessionID *int64, -) error { - return s.enforceAttemptLimit( - func(id int64, window int) (int, error) { - return s.db.CountRecentWebAuthnAttempts(id, window) - }, - userID, - 5, - func() error { - _ = s.db.LogWebAuthnAttempt(userID, nil, false, "rate_limited", ipAddress, userAgent, sessionID) - return nil - }, - ) -} - -func (s *Service) handleE2EAssertion( - userID int64, - methods []*db.MFAMethod, - ipAddress string, - userAgent string, - sessionID *int64, -) (*VerificationResult, bool, error) { - result, handled, err := s.maybeHandleE2EWebAuthn(methods, func(method *db.MFAMethod) error { - _ = s.db.LogWebAuthnAttempt(userID, &method.ID, true, "e2e_test", ipAddress, userAgent, sessionID) - return nil - }) - if !handled { - return nil, false, nil - } - if err != nil { - _ = s.db.LogWebAuthnAttempt(userID, nil, false, "no_method_enrolled", ipAddress, userAgent, sessionID) - } - return result, true, err -} - -func (s *Service) validateAssertion( - user *db.User, - methods []*db.MFAMethod, - sessionJSON []byte, - credentialJSON []byte, - ipAddress string, - userAgent string, - sessionID *int64, -) ([]byte, error) { - var session webauthn.SessionData - if err := json.Unmarshal(sessionJSON, &session); err != nil { - _ = s.db.LogWebAuthnAttempt(user.ID, nil, false, "invalid_session", ipAddress, userAgent, sessionID) - return nil, fmt.Errorf("failed to unmarshal session: %w", err) - } - - waUser := &webAuthnUser{user: user, methods: methods} - - parsedResponse, err := protocol.ParseCredentialRequestResponseBody( - strings.NewReader(string(credentialJSON)), - ) - if err != nil { - _ = s.db.LogWebAuthnAttempt(user.ID, nil, false, "invalid_credential", ipAddress, userAgent, sessionID) - return nil, fmt.Errorf("failed to parse assertion: %w", err) - } - - credential, err := s.webAuthn.ValidateLogin(waUser, session, parsedResponse) - if err != nil { - _ = s.db.LogWebAuthnAttempt(user.ID, nil, false, "validation_failed", ipAddress, userAgent, sessionID) - if lockErr := s.checkAndLockAccount(user.ID); lockErr != nil { - return nil, lockErr - } - return nil, fmt.Errorf("failed to validate assertion: %w", err) - } - - return credential.ID, nil -} - -func (s *Service) findAndConfirmCredential( - userID int64, - methods []*db.MFAMethod, - credentialID []byte, - ipAddress string, - userAgent string, - sessionID *int64, -) (*VerificationResult, error) { - for _, method := range methods { - if method.Type == "webauthn" && string(method.CredentialID) == string(credentialID) { - if err := s.touchMFAMethod(method); err != nil { - return nil, err - } - _ = s.db.LogWebAuthnAttempt(userID, &method.ID, true, "", ipAddress, userAgent, sessionID) - return &VerificationResult{MethodID: method.ID}, nil - } - } - - _ = s.db.LogWebAuthnAttempt(userID, nil, false, "credential_not_found", ipAddress, userAgent, sessionID) - if err := s.checkAndLockAccount(userID); err != nil { - return nil, err - } - return nil, fmt.Errorf("credential not found") -} - -// VerifyWebAuthn validates a WebAuthn assertion during login or step-up (legacy, for web UI). -// -// Deprecated: Use StartWebAuthnAssertion + VerifyWebAuthnAssertion for better session management. -func (s *Service) VerifyWebAuthn(user *db.User, credentialJSON []byte) (*VerificationResult, error) { - if user == nil { - return nil, fmt.Errorf("user required") - } - if s.webAuthn == nil { - return nil, fmt.Errorf("WebAuthn not configured") - } - - methods, err := s.db.ListMFAMethods(user.ID) - if err != nil { - return nil, err - } - - if result, handled, err := s.maybeHandleE2EWebAuthn(methods, nil); handled { - return result, err - } - - credentialID, err := s.performLegacyWebAuthnValidation(user, methods, credentialJSON) - if err != nil { - return nil, err - } - - return s.matchCredentialToMethod(methods, credentialID) -} - -func (s *Service) performLegacyWebAuthnValidation( - user *db.User, - methods []*db.MFAMethod, - credentialJSON []byte, -) ([]byte, error) { - waUser := &webAuthnUser{user: user, methods: methods} - options, session, err := s.webAuthn.BeginLogin(waUser) - if err != nil { - return nil, fmt.Errorf("failed to begin login: %w", err) - } - - _ = options - - parsedResponse, err := protocol.ParseCredentialRequestResponseBody( - strings.NewReader(string(credentialJSON)), - ) - if err != nil { - return nil, fmt.Errorf("failed to parse assertion: %w", err) - } - - credential, err := s.webAuthn.ValidateLogin(waUser, *session, parsedResponse) - if err != nil { - return nil, fmt.Errorf("failed to validate assertion: %w", err) - } - - return credential.ID, nil -} - -func (s *Service) matchCredentialToMethod( - methods []*db.MFAMethod, - credentialID []byte, -) (*VerificationResult, error) { - for _, method := range methods { - if method.Type == "webauthn" && string(method.CredentialID) == string(credentialID) { - if err := s.touchMFAMethod(method); err != nil { - return nil, err - } - return &VerificationResult{MethodID: method.ID}, nil - } + ); err != nil { + return err } - return nil, fmt.Errorf("credential not found") + return s.db.UpdateMFAMethodSignCount(methodID, credential.Authenticator.SignCount) } diff --git a/internal/gateway/auth/mfa/webauthn_assertion.go b/internal/gateway/auth/mfa/webauthn_assertion.go new file mode 100644 index 00000000..5190d091 --- /dev/null +++ b/internal/gateway/auth/mfa/webauthn_assertion.go @@ -0,0 +1,221 @@ +package mfa + +import ( + "encoding/json" + "fmt" + "strings" + + "github.com/go-webauthn/webauthn/protocol" + "github.com/go-webauthn/webauthn/webauthn" + + "github.com/DocSpring/rack-gateway/internal/gateway/db" +) + +// attemptContext carries request details recorded with every WebAuthn attempt. +type attemptContext struct { + userID int64 + ipAddress string + userAgent string + sessionID *int64 +} + +func (a attemptContext) log(s *Service, methodID *int64, success bool, reason string) { + _ = s.db.LogWebAuthnAttempt(a.userID, methodID, success, reason, a.ipAddress, a.userAgent, a.sessionID) +} + +// StartWebAuthnAssertion begins a WebAuthn assertion (login / step-up) ceremony. +// +// The challenge is stored server-side, bound to the user and (when given) the session that started it, +// and the caller receives only an opaque challenge ID to send back with the assertion. User +// verification (PIN or biometric) is required. +func (s *Service) StartWebAuthnAssertion( + user *db.User, + sessionID *int64, +) (*protocol.CredentialAssertion, string, error) { + if user == nil { + return nil, "", fmt.Errorf("user required") + } + if s.webAuthn == nil { + return nil, "", fmt.Errorf("WebAuthn not configured") + } + + methods, err := s.db.ListMFAMethods(user.ID) + if err != nil { + return nil, "", err + } + + waUser := &webAuthnUser{user: user, methods: methods} + options, session, err := s.webAuthn.BeginLogin( + waUser, + webauthn.WithUserVerification(protocol.VerificationRequired), + ) + if err != nil { + return nil, "", fmt.Errorf("failed to begin login: %w", err) + } + + challengeID, err := s.storeChallenge(user.ID, sessionID, db.WebAuthnChallengeAssertion, session) + if err != nil { + return nil, "", err + } + return options, challengeID, nil +} + +// VerifyWebAuthnAssertion validates a WebAuthn assertion against the server-side challenge identified +// by challengeID. The challenge is consumed whether or not the assertion is valid. When the challenge was +// started from a session and sessionID is given, the two must match. Includes rate limiting, automatic +// account locking, and signature-counter clone detection. +func (s *Service) VerifyWebAuthnAssertion( + user *db.User, + challengeID []byte, + credentialJSON []byte, + ipAddress string, + userAgent string, + sessionID *int64, +) (*VerificationResult, error) { + if user == nil { + return nil, fmt.Errorf("user required") + } + if s.webAuthn == nil { + return nil, fmt.Errorf("WebAuthn not configured") + } + attempt := attemptContext{userID: user.ID, ipAddress: ipAddress, userAgent: userAgent, sessionID: sessionID} + + if err := s.ensureUserUnlocked(user.ID); err != nil { + return nil, err + } + if err := s.checkWebAuthnRateLimit(attempt); err != nil { + return nil, err + } + + challenge, err := s.db.ConsumeWebAuthnChallenge( + strings.TrimSpace(string(challengeID)), user.ID, sessionID, db.WebAuthnChallengeAssertion, + ) + if err != nil { + attempt.log(s, nil, false, "invalid_challenge") + return nil, challengeError(err) + } + + methods, err := s.db.ListMFAMethods(user.ID) + if err != nil { + return nil, err + } + + if result, handled, err := s.handleE2EAssertion(attempt, methods); handled { + return result, err + } + + credential, err := s.validateAssertion(user, methods, challenge.SessionData, credentialJSON, attempt) + if err != nil { + return nil, err + } + + return s.findAndConfirmCredential(methods, credential, attempt) +} + +func (s *Service) checkWebAuthnRateLimit(attempt attemptContext) error { + return s.enforceAttemptLimit( + func(id int64, window int) (int, error) { + return s.db.CountRecentWebAuthnAttempts(id, window) + }, + attempt.userID, + 5, + func() error { + attempt.log(s, nil, false, "rate_limited") + return nil + }, + ) +} + +func (s *Service) handleE2EAssertion( + attempt attemptContext, + methods []*db.MFAMethod, +) (*VerificationResult, bool, error) { + result, handled, err := s.maybeHandleE2EWebAuthn(methods, func(method *db.MFAMethod) error { + attempt.log(s, &method.ID, true, "e2e_test") + return nil + }) + if !handled { + return nil, false, nil + } + if err != nil { + attempt.log(s, nil, false, "no_method_enrolled") + } + return result, true, err +} + +func (s *Service) validateAssertion( + user *db.User, + methods []*db.MFAMethod, + sessionJSON []byte, + credentialJSON []byte, + attempt attemptContext, +) (*webauthn.Credential, error) { + var session webauthn.SessionData + if err := json.Unmarshal(sessionJSON, &session); err != nil { + attempt.log(s, nil, false, "invalid_session") + return nil, fmt.Errorf("failed to unmarshal session: %w", err) + } + + parsedResponse, err := protocol.ParseCredentialRequestResponseBody( + strings.NewReader(string(credentialJSON)), + ) + if err != nil { + attempt.log(s, nil, false, "invalid_credential") + return nil, fmt.Errorf("failed to parse assertion: %w", err) + } + + waUser := &webAuthnUser{user: user, methods: methods} + credential, err := s.webAuthn.ValidateLogin(waUser, session, parsedResponse) + if err != nil { + attempt.log(s, nil, false, "validation_failed") + if lockErr := s.checkAndLockAccount(user.ID); lockErr != nil { + return nil, lockErr + } + return nil, fmt.Errorf("failed to validate assertion: %w", err) + } + + return credential, nil +} + +// findAndConfirmCredential matches the validated credential to its MFA method, rejects a signature +// counter that didn't increase (possible cloned authenticator), and stores the new counter. +func (s *Service) findAndConfirmCredential( + methods []*db.MFAMethod, + credential *webauthn.Credential, + attempt attemptContext, +) (*VerificationResult, error) { + method := findWebAuthnMethod(methods, credential.ID) + if method == nil { + attempt.log(s, nil, false, "credential_not_found") + if err := s.checkAndLockAccount(attempt.userID); err != nil { + return nil, err + } + return nil, fmt.Errorf("credential not found") + } + + if credential.Authenticator.CloneWarning { + attempt.log(s, &method.ID, false, "sign_count_not_increased") + if err := s.checkAndLockAccount(attempt.userID); err != nil { + return nil, err + } + return nil, fmt.Errorf("security key signature counter did not increase; the key may have been cloned") + } + + if err := s.db.UpdateMFAMethodSignCount(method.ID, credential.Authenticator.SignCount); err != nil { + return nil, err + } + if err := s.touchMFAMethod(method); err != nil { + return nil, err + } + attempt.log(s, &method.ID, true, "") + return &VerificationResult{MethodID: method.ID}, nil +} + +func findWebAuthnMethod(methods []*db.MFAMethod, credentialID []byte) *db.MFAMethod { + for _, method := range methods { + if method.Type == "webauthn" && string(method.CredentialID) == string(credentialID) { + return method + } + } + return nil +} diff --git a/internal/gateway/auth/mfa/webauthn_challenge_test.go b/internal/gateway/auth/mfa/webauthn_challenge_test.go new file mode 100644 index 00000000..7512d468 --- /dev/null +++ b/internal/gateway/auth/mfa/webauthn_challenge_test.go @@ -0,0 +1,175 @@ +package mfa + +import ( + "crypto/rand" + "encoding/hex" + "errors" + "testing" + "time" + + "github.com/stretchr/testify/require" + + "github.com/DocSpring/rack-gateway/internal/gateway/db" + "github.com/DocSpring/rack-gateway/internal/gateway/testutil/dbtest" + "github.com/DocSpring/rack-gateway/internal/gateway/testutil/webauthntest" +) + +func newTestSession(t *testing.T, database *db.Database, userID int64) int64 { + t.Helper() + buf := make([]byte, 32) + _, err := rand.Read(buf) + require.NoError(t, err) + session, err := database.CreateUserSession( + userID, hex.EncodeToString(buf), time.Now().Add(time.Hour), "web", "", "", "127.0.0.1", "test", nil, nil, + ) + require.NoError(t, err) + return session.ID +} + +type webAuthnFixture struct { + service *Service + database *db.Database + user *db.User + method *db.MFAMethod + credential *webauthntest.MockCredential +} + +func newWebAuthnFixture(t *testing.T) *webAuthnFixture { + t.Helper() + database := dbtest.NewDatabase(t) + service, err := NewService( + database, "Test Gateway", 24*time.Hour, 10*time.Minute, []byte("pepper"), + "", "", "localhost", "http://localhost", nil, + ) + require.NoError(t, err) + user, err := database.CreateUser("key@example.com", "Key User", []string{"admin"}) + require.NoError(t, err) + credential, err := webauthntest.GenerateMockCredential() + require.NoError(t, err) + method, err := database.CreateMFAMethod( + user.ID, "webauthn", "Key", "", credential.ID, credential.PublicKey, nil, nil, + ) + require.NoError(t, err) + require.NoError(t, database.ConfirmMFAMethod(method.ID, time.Now())) + return &webAuthnFixture{service: service, database: database, user: user, method: method, credential: credential} +} + +// assert starts a ceremony for startSession and verifies it for verifySession. +func (f *webAuthnFixture) assert(t *testing.T, startSession, verifySession *int64) (string, string, error) { + t.Helper() + options, challengeID, err := f.service.StartWebAuthnAssertion(f.user, startSession) + require.NoError(t, err) + assertionJSON, err := f.credential.GenerateAssertion(options, "http://localhost") + require.NoError(t, err) + _, err = f.service.VerifyWebAuthnAssertion( + f.user, []byte(challengeID), []byte(assertionJSON), "127.0.0.1", "test", verifySession, + ) + return challengeID, assertionJSON, err +} + +func TestWebAuthnAssertionChallengeIsSingleUse(t *testing.T) { + t.Parallel() + f := newWebAuthnFixture(t) + f.credential.Counter = 1 + + challengeID, assertionJSON, err := f.assert(t, nil, nil) + require.NoError(t, err) + + // Replaying the same assertion and challenge must fail. + _, err = f.service.VerifyWebAuthnAssertion( + f.user, []byte(challengeID), []byte(assertionJSON), "127.0.0.1", "test", nil, + ) + require.ErrorIs(t, err, ErrWebAuthnChallenge) +} + +func TestWebAuthnAssertionRejectsForgedOrForeignChallenge(t *testing.T) { + t.Parallel() + f := newWebAuthnFixture(t) + + // A made-up challenge ID is rejected. + _, err := f.service.VerifyWebAuthnAssertion( + f.user, []byte("wac_forged"), []byte(`{}`), "127.0.0.1", "test", nil, + ) + require.ErrorIs(t, err, ErrWebAuthnChallenge) + + // A challenge started by one session can't be used by another session. + sessionA := newTestSession(t, f.database, f.user.ID) + sessionB := newTestSession(t, f.database, f.user.ID) + _, _, err = f.assert(t, &sessionA, &sessionB) + require.ErrorIs(t, err, ErrWebAuthnChallenge) + + // The owning session succeeds. + f.credential.Counter = 1 + _, _, err = f.assert(t, &sessionA, &sessionA) + require.NoError(t, err) + + // A challenge belonging to another user is rejected. + other, err := f.database.CreateUser("other@example.com", "Other", []string{"admin"}) + require.NoError(t, err) + _, challengeID, err := f.service.StartWebAuthnAssertion(f.user, nil) + require.NoError(t, err) + _, err = f.service.VerifyWebAuthnAssertion(other, []byte(challengeID), []byte(`{}`), "127.0.0.1", "test", nil) + require.ErrorIs(t, err, ErrWebAuthnChallenge) +} + +func TestWebAuthnAssertionChallengeExpires(t *testing.T) { + t.Parallel() + f := newWebAuthnFixture(t) + options, challengeID, err := f.service.StartWebAuthnAssertion(f.user, nil) + require.NoError(t, err) + _, err = f.database.DB().Exec( + "UPDATE webauthn_challenges SET expires_at = NOW() - INTERVAL '1 second' WHERE id = $1", challengeID, + ) + require.NoError(t, err) + + assertionJSON, err := f.credential.GenerateAssertion(options, "http://localhost") + require.NoError(t, err) + _, err = f.service.VerifyWebAuthnAssertion( + f.user, []byte(challengeID), []byte(assertionJSON), "127.0.0.1", "test", nil, + ) + require.ErrorIs(t, err, ErrWebAuthnChallenge) +} + +func TestWebAuthnAssertionRequiresUserVerification(t *testing.T) { + t.Parallel() + f := newWebAuthnFixture(t) + f.credential.WithoutUserVerification = true + + _, _, err := f.assert(t, nil, nil) + require.Error(t, err, "assertion without user verification (PIN) must be rejected") + require.False(t, errors.Is(err, ErrWebAuthnChallenge)) +} + +func TestWebAuthnAssertionEnforcesSignatureCounter(t *testing.T) { + t.Parallel() + f := newWebAuthnFixture(t) + + f.credential.Counter = 10 + _, _, err := f.assert(t, nil, nil) + require.NoError(t, err) + stored, err := f.database.GetMFAMethodByID(f.method.ID) + require.NoError(t, err) + require.Equal(t, uint32(10), stored.SignCount, "new counter must be stored") + + // A counter that doesn't increase means the key may have been cloned. + f.credential.Counter = 10 + _, _, err = f.assert(t, nil, nil) + require.ErrorContains(t, err, "signature counter") + f.credential.Counter = 7 + _, _, err = f.assert(t, nil, nil) + require.ErrorContains(t, err, "signature counter") + + f.credential.Counter = 11 + _, _, err = f.assert(t, nil, nil) + require.NoError(t, err) +} + +func TestWebAuthnAssertionAllowsZeroCounterAuthenticators(t *testing.T) { + t.Parallel() + f := newWebAuthnFixture(t) + // Many platform authenticators and passkey managers always report 0. + for i := 0; i < 2; i++ { + _, _, err := f.assert(t, nil, nil) + require.NoError(t, err) + } +} diff --git a/internal/gateway/auth/mfa/webauthn_test.go b/internal/gateway/auth/mfa/webauthn_test.go index 32647a7c..0368fb97 100644 --- a/internal/gateway/auth/mfa/webauthn_test.go +++ b/internal/gateway/auth/mfa/webauthn_test.go @@ -61,19 +61,19 @@ func TestMockCredentialGeneratesValidAssertion(t *testing.T) { t.Fatalf("failed to confirm MFA method: %v", err) } - _, sessionData, err := service.StartWebAuthnAssertion(user) + options, challengeID, err := service.StartWebAuthnAssertion(user, nil) if err != nil { t.Fatalf("failed to start assertion: %v", err) } - assertionJSON, err := credential.GenerateAssertionForSession(sessionData, "http://localhost") + assertionJSON, err := credential.GenerateAssertion(options, "http://localhost") if err != nil { t.Fatalf("failed to generate assertion: %v", err) } result, err := service.VerifyWebAuthnAssertion( user, - sessionData, + []byte(challengeID), []byte(assertionJSON), "127.0.0.1", "test-agent", @@ -111,8 +111,9 @@ func TestWebAuthnEnrollment_Success(t *testing.T) { } user, _ := database.CreateUser("webauthn@example.com", "WebAuthn Test", []string{"admin"}) + sessionID := newTestSession(t, database, user.ID) - result, sessionData, err := service.StartWebAuthnEnrollment(user) + result, err := service.StartWebAuthnEnrollment(user, sessionID) if err != nil { t.Fatalf("expected successful enrollment start, got error: %v", err) } @@ -125,14 +126,20 @@ func TestWebAuthnEnrollment_Success(t *testing.T) { t.Error("expected public key options to be returned") } - if sessionData == "" { - t.Error("expected session data to be returned") + // The registration challenge is stored server-side for this session, with an expiry + challenge, err := database.ConsumeSessionWebAuthnChallenge(user.ID, sessionID, "registration") + if err != nil { + t.Fatalf("expected stored registration challenge: %v", err) } - - // Verify session data is valid JSON var session webauthn.SessionData - if err := json.Unmarshal([]byte(sessionData), &session); err != nil { - t.Errorf("session data should be valid JSON: %v", err) + if err := json.Unmarshal(challenge.SessionData, &session); err != nil { + t.Errorf("stored session data should be valid JSON: %v", err) + } + if session.Expires.IsZero() { + t.Error("stored registration session must have an expiry") + } + if session.UserVerification != protocol.VerificationRequired { + t.Errorf("registration must require user verification, got %q", session.UserVerification) } // Verify backup codes are generated on first enrollment @@ -162,7 +169,7 @@ func TestWebAuthnEnrollment_NotConfigured(t *testing.T) { user, _ := database.CreateUser("webauthn2@example.com", "WebAuthn Test", []string{"admin"}) - _, _, err := serviceWithoutWebAuthn.StartWebAuthnEnrollment(user) + _, err := serviceWithoutWebAuthn.StartWebAuthnEnrollment(user, newTestSession(t, database, user.ID)) if err == nil { t.Error("expected error when WebAuthn not configured") } @@ -192,9 +199,11 @@ func TestWebAuthnEnrollment_CleanupUnconfirmed(t *testing.T) { user, _ := database.CreateUser("webauthn3@example.com", "WebAuthn Test", []string{"admin"}) + sessionID := newTestSession(t, database, user.ID) + // Start enrollment twice to ensure cleanup works - result1, _, _ := service.StartWebAuthnEnrollment(user) - result2, _, _ := service.StartWebAuthnEnrollment(user) + result1, _ := service.StartWebAuthnEnrollment(user, sessionID) + result2, _ := service.StartWebAuthnEnrollment(user, sessionID) // Second enrollment should have cleaned up first if result1.MethodID == result2.MethodID { diff --git a/internal/gateway/auth/mfa/webauthn_user.go b/internal/gateway/auth/mfa/webauthn_user.go index 8ea4375d..5e621755 100644 --- a/internal/gateway/auth/mfa/webauthn_user.go +++ b/internal/gateway/auth/mfa/webauthn_user.go @@ -50,6 +50,7 @@ func (u *webAuthnUser) WebAuthnCredentials() []webauthn.Credential { AttestationType: "", Transport: convertTransports(method.Transports), Flags: extractCredentialFlags(method.Metadata), + Authenticator: webauthn.Authenticator{SignCount: method.SignCount}, }) } return creds diff --git a/internal/gateway/db/database.go b/internal/gateway/db/database.go index 3d1a3aab..220e7bfb 100644 --- a/internal/gateway/db/database.go +++ b/internal/gateway/db/database.go @@ -390,6 +390,7 @@ func (d *Database) dropAllTables() error { // Drop dependent tables first to satisfy foreign keys. tables := []string{ + "webauthn_challenges", "user_resources", "deploy_approval_requests", "api_tokens", diff --git a/internal/gateway/db/mfa.go b/internal/gateway/db/mfa.go index a9746887..ad30ae71 100644 --- a/internal/gateway/db/mfa.go +++ b/internal/gateway/db/mfa.go @@ -216,7 +216,7 @@ func (d *Database) ListAllMFAMethods(userID int64) ([]*MFAMethod, error) { func (d *Database) queryMFAMethods(userID int64, includeUnconfirmed bool) ([]*MFAMethod, error) { query := ` SELECT id, user_id, type, label, secret, credential_id, public_key, transports, metadata, - cli_capable, created_at, confirmed_at, last_used_at + cli_capable, webauthn_sign_count, created_at, confirmed_at, last_used_at FROM mfa_methods WHERE user_id = ?` if !includeUnconfirmed { query += " AND confirmed_at IS NOT NULL" @@ -297,7 +297,7 @@ func (d *Database) UpdateMFAMethodCredential( func (d *Database) GetMFAMethodByID(id int64) (*MFAMethod, error) { query := ` SELECT id, user_id, type, label, secret, credential_id, public_key, transports, metadata, - cli_capable, created_at, confirmed_at, last_used_at + cli_capable, webauthn_sign_count, created_at, confirmed_at, last_used_at FROM mfa_methods WHERE id = ? ` method, err := scanMFAMethod(d.queryRow(query, id)) diff --git a/internal/gateway/db/mfa_scan.go b/internal/gateway/db/mfa_scan.go index 78c08969..ffe0237e 100644 --- a/internal/gateway/db/mfa_scan.go +++ b/internal/gateway/db/mfa_scan.go @@ -23,7 +23,7 @@ func scanMFAMethod(scanner interface { err := scanner.Scan( &method.ID, &method.UserID, &method.Type, &label, &secret, &method.CredentialID, &method.PublicKey, &transports, &metadata, - &method.CLICapable, &method.CreatedAt, &confirmed, &lastUsed, + &method.CLICapable, &method.SignCount, &method.CreatedAt, &confirmed, &lastUsed, ) if err != nil { return nil, err diff --git a/internal/gateway/db/migrations/20261009120000_webauthn_challenges.sql b/internal/gateway/db/migrations/20261009120000_webauthn_challenges.sql new file mode 100644 index 00000000..3fd9cbff --- /dev/null +++ b/internal/gateway/db/migrations/20261009120000_webauthn_challenges.sql @@ -0,0 +1,21 @@ +-- WebAuthn challenges live on the server. Clients only receive an opaque challenge ID, which is +-- bound to the user (and the session that started the ceremony), single-use and short-lived. +-- Previously the full session data (including the challenge) round-tripped through the client +-- and was trusted on the way back, so one captured assertion could be replayed indefinitely. +CREATE TABLE IF NOT EXISTS webauthn_challenges ( + id TEXT PRIMARY KEY, + user_id BIGINT NOT NULL REFERENCES users(id) ON DELETE CASCADE, + session_id BIGINT REFERENCES user_sessions(id) ON DELETE CASCADE, + purpose VARCHAR(20) NOT NULL CHECK (purpose IN ('assertion', 'registration')), + session_data TEXT NOT NULL, + created_at TIMESTAMPTZ NOT NULL DEFAULT NOW(), + expires_at TIMESTAMPTZ NOT NULL +); + +CREATE INDEX IF NOT EXISTS idx_webauthn_challenges_expires ON webauthn_challenges(expires_at); +CREATE INDEX IF NOT EXISTS idx_webauthn_challenges_owner ON webauthn_challenges(user_id, session_id, purpose); + +-- Stored signature counter for cloned-authenticator detection (WebAuthn ยง7.2 step 17). +ALTER TABLE mfa_methods ADD COLUMN IF NOT EXISTS webauthn_sign_count BIGINT NOT NULL DEFAULT 0; + +COMMENT ON COLUMN mfa_methods.webauthn_sign_count IS 'Last WebAuthn signature counter seen for this credential'; diff --git a/internal/gateway/db/sessions_mutations.go b/internal/gateway/db/sessions_mutations.go index e61b04e6..63682e8e 100644 --- a/internal/gateway/db/sessions_mutations.go +++ b/internal/gateway/db/sessions_mutations.go @@ -1,7 +1,6 @@ package db import ( - "encoding/json" "fmt" "time" @@ -147,41 +146,3 @@ func (d *Database) AttachTrustedDeviceToSession(sessionID int64, trustedDeviceID } return nil } - -// UpdateSessionMetadata merges new metadata into an existing session's metadata field. -func (d *Database) UpdateSessionMetadata(sessionID int64, metadata map[string]interface{}) error { - if len(metadata) == 0 { - return nil - } - - session, err := d.GetSessionByID(sessionID) - if err != nil { - return fmt.Errorf("failed to get session: %w", err) - } - if session == nil { - return fmt.Errorf("session not found") - } - - existingMeta := make(map[string]interface{}) - if len(session.Metadata) > 0 { - if err := json.Unmarshal(session.Metadata, &existingMeta); err != nil { - return fmt.Errorf("failed to parse existing metadata: %w", err) - } - } - - for k, v := range metadata { - existingMeta[k] = v - } - - metaJSON := marshalJSONMap(existingMeta) - - _, err = d.exec( - "UPDATE user_sessions SET metadata = ?, updated_at = NOW() WHERE id = ?", - metaJSON, - sessionID, - ) - if err != nil { - return fmt.Errorf("failed to update session metadata: %w", err) - } - return nil -} diff --git a/internal/gateway/db/sessions_test.go b/internal/gateway/db/sessions_test.go index 821a6aef..48e2baaf 100644 --- a/internal/gateway/db/sessions_test.go +++ b/internal/gateway/db/sessions_test.go @@ -1,134 +1,12 @@ package db_test import ( - "encoding/json" "testing" "time" - "github.com/DocSpring/rack-gateway/internal/gateway/db" "github.com/DocSpring/rack-gateway/internal/gateway/testutil/dbtest" ) -func TestUpdateSessionMetadata(t *testing.T) { - t.Parallel() - - database := dbtest.NewDatabase(t) - session := createTestSession(t, database) - - t.Run("merges new metadata with existing", func(t *testing.T) { - testMetadataMerge(t, database, session.ID) - }) - - t.Run("overwrites existing keys", func(t *testing.T) { - testMetadataOverwrite(t, database, session.ID) - }) - - t.Run("handles nil metadata gracefully", func(t *testing.T) { - err := database.UpdateSessionMetadata(session.ID, nil) - if err != nil { - t.Error("should not error on nil metadata") - } - }) - - t.Run("handles empty metadata gracefully", func(t *testing.T) { - err := database.UpdateSessionMetadata(session.ID, map[string]interface{}{}) - if err != nil { - t.Error("should not error on empty metadata") - } - }) - - t.Run("fails for non-existent session", func(t *testing.T) { - err := database.UpdateSessionMetadata(99999, map[string]interface{}{"key": "value"}) - if err == nil { - t.Error("expected error for non-existent session") - } - }) -} - -func createTestSession(t *testing.T, database *db.Database) *db.UserSession { - t.Helper() - user, err := database.CreateUser("session@example.com", "Session Test", []string{"ops"}) - if err != nil { - t.Fatalf("failed to create user: %v", err) - } - - tokenHash := "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef" - expiresAt := time.Now().Add(1 * time.Hour) - initialMeta := map[string]interface{}{ - "initial_key": "initial_value", - } - - session, err := database.CreateUserSession( - user.ID, - tokenHash, - expiresAt, - "web", - "", - "Test Device", - "192.168.1.1", - "Mozilla/5.0", - initialMeta, - nil, - ) - if err != nil { - t.Fatalf("failed to create session: %v", err) - } - return session -} - -func testMetadataMerge(t *testing.T, database *db.Database, sessionID int64) { - t.Helper() - newMeta := map[string]interface{}{ - "new_key": "new_value", - "number": 42, - } - - err := database.UpdateSessionMetadata(sessionID, newMeta) - if err != nil { - t.Fatalf("failed to update metadata: %v", err) - } - - updated, err := database.GetSessionByID(sessionID) - if err != nil { - t.Fatalf("failed to get updated session: %v", err) - } - - var meta map[string]interface{} - if err := json.Unmarshal(updated.Metadata, &meta); err != nil { - t.Fatalf("failed to unmarshal metadata: %v", err) - } - - if meta["initial_key"] != "initial_value" { - t.Error("initial_key should be preserved") - } - if meta["new_key"] != "new_value" { - t.Error("new_key should be added") - } - if meta["number"] != float64(42) { - t.Errorf("number should be 42, got %v", meta["number"]) - } -} - -func testMetadataOverwrite(t *testing.T, database *db.Database, sessionID int64) { - t.Helper() - updateMeta := map[string]interface{}{ - "initial_key": "updated_value", - } - - err := database.UpdateSessionMetadata(sessionID, updateMeta) - if err != nil { - t.Fatalf("failed to update metadata: %v", err) - } - - updated, _ := database.GetSessionByID(sessionID) - var meta map[string]interface{} - _ = json.Unmarshal(updated.Metadata, &meta) - - if meta["initial_key"] != "updated_value" { - t.Errorf("initial_key should be overwritten, got %v", meta["initial_key"]) - } -} - func TestGetSessionByID(t *testing.T) { t.Parallel() @@ -179,131 +57,3 @@ func TestGetSessionByID(t *testing.T) { } }) } - -func TestSessionMetadataWithWebAuthn(t *testing.T) { - t.Parallel() - - database := dbtest.NewDatabase(t) - session := createWebAuthnTestSession(t, database) - - t.Run("stores and retrieves WebAuthn enrollment session", func(t *testing.T) { - testWebAuthnEnrollmentStore(t, database, session.ID) - }) - - t.Run("can overwrite WebAuthn session in metadata", func(t *testing.T) { - testWebAuthnOverwrite(t, database, session.ID) - }) -} - -func createWebAuthnTestSession(t *testing.T, database *db.Database) *db.UserSession { - t.Helper() - user, _ := database.CreateUser("webauthn-session@example.com", "WebAuthn Session Test", []string{"ops"}) - - session, _ := database.CreateUserSession( - user.ID, - "fedcba9876543210fedcba9876543210fedcba9876543210fedcba9876543210", - time.Now().Add(1*time.Hour), - "web", - "", - "", - "", - "", - nil, - nil, - ) - return session -} - -func testWebAuthnEnrollmentStore(t *testing.T, database *db.Database, sessionID int64) { - t.Helper() - webauthnSession := map[string]interface{}{ - "webauthn_enrollment_session": `{"challenge":"abc123","user_id":"1"}`, - "webauthn_enrollment_expires": time.Now().Add(5 * time.Minute).Unix(), - } - - err := database.UpdateSessionMetadata(sessionID, webauthnSession) - if err != nil { - t.Fatalf("failed to store WebAuthn session: %v", err) - } - - retrieved, err := database.GetSessionByID(sessionID) - if err != nil { - t.Fatalf("failed to get session: %v", err) - } - if retrieved == nil { - t.Fatal("session should exist") - } - - var meta map[string]interface{} - if err := json.Unmarshal(retrieved.Metadata, &meta); err != nil { - t.Fatalf("failed to unmarshal metadata: %v", err) - } - - verifyWebAuthnSessionData(t, meta) - verifyWebAuthnExpiration(t, meta) -} - -func verifyWebAuthnSessionData(t *testing.T, meta map[string]interface{}) { - t.Helper() - sessionData, ok := meta["webauthn_enrollment_session"].(string) - if !ok { - t.Fatal("webauthn_enrollment_session should be a string") - } - - if sessionData != `{"challenge":"abc123","user_id":"1"}` { - t.Error("WebAuthn session data mismatch") - } -} - -func verifyWebAuthnExpiration(t *testing.T, meta map[string]interface{}) { - t.Helper() - expiresFloat, ok := meta["webauthn_enrollment_expires"].(float64) - if !ok { - t.Fatal("webauthn_enrollment_expires should be a number") - } - - if expiresFloat <= 0 { - t.Error("WebAuthn expiration should be positive") - } -} - -func testWebAuthnOverwrite(t *testing.T, database *db.Database, sessionID int64) { - t.Helper() - // First store WebAuthn session - webauthnSession := map[string]interface{}{ - "webauthn_enrollment_session": "data", - "webauthn_enrollment_expires": 12345, - } - _ = database.UpdateSessionMetadata(sessionID, webauthnSession) - - // Verify it was stored - retrieved, _ := database.GetSessionByID(sessionID) - var meta map[string]interface{} - _ = json.Unmarshal(retrieved.Metadata, &meta) - - if _, exists := meta["webauthn_enrollment_session"]; !exists { - t.Error("webauthn_enrollment_session should exist") - } - - // Now overwrite with different data - newSession := map[string]interface{}{ - "webauthn_enrollment_session": "new_data", - "webauthn_enrollment_expires": 67890, - } - err := database.UpdateSessionMetadata(sessionID, newSession) - if err != nil { - t.Fatalf("failed to update metadata: %v", err) - } - - // Verify it was overwritten - updated, _ := database.GetSessionByID(sessionID) - var updatedMeta map[string]interface{} - _ = json.Unmarshal(updated.Metadata, &updatedMeta) - - if updatedMeta["webauthn_enrollment_session"] != "new_data" { - t.Error("webauthn_enrollment_session should be overwritten") - } - if updatedMeta["webauthn_enrollment_expires"] != float64(67890) { - t.Error("webauthn_enrollment_expires should be overwritten") - } -} diff --git a/internal/gateway/db/types.go b/internal/gateway/db/types.go index 7c49cbc0..3f3a628f 100644 --- a/internal/gateway/db/types.go +++ b/internal/gateway/db/types.go @@ -176,6 +176,7 @@ type MFAMethod struct { Transports []string `json:"transports,omitempty"` Metadata []byte `json:"metadata,omitempty"` CLICapable bool `json:"cli_capable"` + SignCount uint32 `json:"-"` // last WebAuthn signature counter seen CreatedAt time.Time `json:"created_at"` ConfirmedAt *time.Time `json:"confirmed_at,omitempty"` LastUsedAt *time.Time `json:"last_used_at,omitempty"` diff --git a/internal/gateway/db/webauthn_challenges.go b/internal/gateway/db/webauthn_challenges.go new file mode 100644 index 00000000..386e11cf --- /dev/null +++ b/internal/gateway/db/webauthn_challenges.go @@ -0,0 +1,139 @@ +package db + +import ( + "crypto/rand" + "database/sql" + "encoding/base64" + "errors" + "fmt" + "time" +) + +// WebAuthn challenge purposes. +const ( + WebAuthnChallengeAssertion = "assertion" + WebAuthnChallengeRegistration = "registration" +) + +// ErrWebAuthnChallengeNotFound is returned when a challenge is unknown, expired, already used, +// or belongs to a different user or session. +var ErrWebAuthnChallengeNotFound = errors.New("webauthn challenge not found or expired") + +// WebAuthnChallenge is a server-side WebAuthn ceremony (go-webauthn session data). +type WebAuthnChallenge struct { + ID string + UserID int64 + SessionID *int64 + Purpose string + SessionData []byte +} + +// CreateWebAuthnChallenge stores a challenge and returns its opaque ID. Expired challenges are purged +// opportunistically. Only one registration challenge is kept per user and session. +func (d *Database) CreateWebAuthnChallenge( + userID int64, + sessionID *int64, + purpose string, + sessionData []byte, + ttl time.Duration, +) (string, error) { + id, err := newChallengeID() + if err != nil { + return "", err + } + if _, err := d.exec("DELETE FROM webauthn_challenges WHERE expires_at <= NOW()"); err != nil { + return "", fmt.Errorf("failed to purge expired webauthn challenges: %w", err) + } + if purpose == WebAuthnChallengeRegistration { + if _, err := d.exec(` + DELETE FROM webauthn_challenges + WHERE user_id = ? AND session_id IS NOT DISTINCT FROM ?::BIGINT AND purpose = ?`, + userID, nullableInt64(sessionID), purpose, + ); err != nil { + return "", fmt.Errorf("failed to replace webauthn registration challenge: %w", err) + } + } + _, err = d.exec(` + INSERT INTO webauthn_challenges (id, user_id, session_id, purpose, session_data, expires_at) + VALUES (?, ?, ?, ?, ?, NOW() + make_interval(secs => ?))`, + id, userID, nullableInt64(sessionID), purpose, string(sessionData), ttl.Seconds(), + ) + if err != nil { + return "", fmt.Errorf("failed to store webauthn challenge: %w", err) + } + return id, nil +} + +// ConsumeWebAuthnChallenge atomically deletes and returns an unexpired challenge owned by userID. +// When the challenge was started from a session and sessionID is given, the sessions must match. +func (d *Database) ConsumeWebAuthnChallenge( + id string, + userID int64, + sessionID *int64, + purpose string, +) (*WebAuthnChallenge, error) { + row := d.queryRow(` + DELETE FROM webauthn_challenges + WHERE id = ? AND user_id = ? AND purpose = ? AND expires_at > NOW() + AND (session_id IS NULL OR ?::BIGINT IS NULL OR session_id = ?::BIGINT) + RETURNING id, user_id, session_id, purpose, session_data`, + id, userID, purpose, nullableInt64(sessionID), nullableInt64(sessionID), + ) + return scanWebAuthnChallenge(row) +} + +// ConsumeSessionWebAuthnChallenge atomically deletes and returns the unexpired challenge of the given +// purpose for a user's session (used for registration, where the client holds no challenge ID). +func (d *Database) ConsumeSessionWebAuthnChallenge( + userID int64, + sessionID int64, + purpose string, +) (*WebAuthnChallenge, error) { + row := d.queryRow(` + DELETE FROM webauthn_challenges + WHERE id = ( + SELECT id FROM webauthn_challenges + WHERE user_id = ? AND session_id = ? AND purpose = ? AND expires_at > NOW() + ORDER BY created_at DESC LIMIT 1 + ) + RETURNING id, user_id, session_id, purpose, session_data`, + userID, sessionID, purpose, + ) + return scanWebAuthnChallenge(row) +} + +// UpdateMFAMethodSignCount stores the latest WebAuthn signature counter for a credential. +func (d *Database) UpdateMFAMethodSignCount(methodID int64, signCount uint32) error { + _, err := d.exec("UPDATE mfa_methods SET webauthn_sign_count = ? WHERE id = ?", signCount, methodID) + if err != nil { + return fmt.Errorf("failed to update webauthn sign count: %w", err) + } + return nil +} + +func scanWebAuthnChallenge(row *sql.Row) (*WebAuthnChallenge, error) { + var challenge WebAuthnChallenge + var sessionID sql.NullInt64 + var data string + err := row.Scan(&challenge.ID, &challenge.UserID, &sessionID, &challenge.Purpose, &data) + if errors.Is(err, sql.ErrNoRows) { + return nil, ErrWebAuthnChallengeNotFound + } + if err != nil { + return nil, fmt.Errorf("failed to consume webauthn challenge: %w", err) + } + if sessionID.Valid { + v := sessionID.Int64 + challenge.SessionID = &v + } + challenge.SessionData = []byte(data) + return &challenge, nil +} + +func newChallengeID() (string, error) { + buf := make([]byte, 32) + if _, err := rand.Read(buf); err != nil { + return "", fmt.Errorf("failed to generate webauthn challenge id: %w", err) + } + return "wac_" + base64.RawURLEncoding.EncodeToString(buf), nil +} diff --git a/internal/gateway/handlers/auth_mfa_enrollment.go b/internal/gateway/handlers/auth_mfa_enrollment.go index 0fccddd1..496516fe 100644 --- a/internal/gateway/handlers/auth_mfa_enrollment.go +++ b/internal/gateway/handlers/auth_mfa_enrollment.go @@ -6,7 +6,6 @@ import ( "log" "net/http" "strings" - "time" "github.com/DocSpring/rack-gateway/internal/gateway/audit" @@ -169,7 +168,13 @@ func (h *AuthHandler) StartWebAuthnEnrollment(c *gin.Context) { return } - result, sessionData, err := h.mfaService.StartWebAuthnEnrollment(ctx.userRecord) + sessionID, ok := auth.GetSessionID(c.Request.Context()) + if !ok { + c.JSON(http.StatusUnauthorized, gin.H{"error": "session not found"}) + return + } + + result, err := h.mfaService.StartWebAuthnEnrollment(ctx.userRecord, sessionID) if err != nil { log.Printf( `{"level":"error","event":"webauthn_start_failed","user":%q,"error":%q}`, @@ -180,24 +185,6 @@ func (h *AuthHandler) StartWebAuthnEnrollment(c *gin.Context) { return } - // Store WebAuthn session data in the user's HTTP session metadata - sessionID, ok := auth.GetSessionID(c.Request.Context()) - if !ok { - c.JSON(http.StatusInternalServerError, gin.H{"error": "session not found"}) - return - } - - // Update session metadata with WebAuthn session - metadata := map[string]interface{}{ - "webauthn_enrollment_session": sessionData, - "webauthn_enrollment_expires": time.Now().Add(5 * time.Minute).Unix(), - } - if err := h.database.UpdateSessionMetadata(sessionID, metadata); err != nil { - log.Printf("failed to store webauthn session: %v", err) - c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to store session"}) - return - } - // Debug: log what we're returning optionsJSON, _ := json.Marshal(result.PublicKeyOptions) log.Printf( @@ -240,8 +227,9 @@ func (h *AuthHandler) ConfirmWebAuthnEnrollment(c *gin.Context) { return } - sessionID, sessionDataStr, ok := h.retrieveWebAuthnEnrollmentSession(c) + sessionID, ok := auth.GetSessionID(c.Request.Context()) if !ok { + c.JSON(http.StatusUnauthorized, gin.H{"error": "session not found"}) return } @@ -259,8 +247,8 @@ func (h *AuthHandler) ConfirmWebAuthnEnrollment(c *gin.Context) { methodID, err := h.mfaService.ConfirmWebAuthnEnrollment( ctx.userRecord, + sessionID, req.MethodID, - []byte(sessionDataStr), credentialJSON, label, ) @@ -269,8 +257,6 @@ func (h *AuthHandler) ConfirmWebAuthnEnrollment(c *gin.Context) { return } - h.clearWebAuthnEnrollmentSession(sessionID) - _, ok = h.updateSessionAfterMFA(c, ctx, ctx.authUser.Session.TrustedDeviceID, false) if !ok { return @@ -324,59 +310,3 @@ func defaultMFAMethodLabel(methodType string) string { return "Security Key" } } - -func (h *AuthHandler) retrieveWebAuthnEnrollmentSession(c *gin.Context) (int64, string, bool) { - sessionID, ok := auth.GetSessionID(c.Request.Context()) - if !ok { - c.JSON(http.StatusUnauthorized, gin.H{"error": "session not found"}) - return 0, "", false - } - - session, err := h.database.GetSessionByID(sessionID) - if err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load session"}) - return 0, "", false - } - - var sessionMeta map[string]interface{} - if len(session.Metadata) > 0 { - if err := json.Unmarshal(session.Metadata, &sessionMeta); err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": "invalid session metadata"}) - return 0, "", false - } - } - - sessionDataStr, ok := sessionMeta["webauthn_enrollment_session"].(string) - if !ok || sessionDataStr == "" { - c.JSON(http.StatusBadRequest, gin.H{"error": "webauthn session not found or expired"}) - return 0, "", false - } - - expiresFloat, ok := sessionMeta["webauthn_enrollment_expires"].(float64) - if ok && time.Now().Unix() > int64(expiresFloat) { - c.JSON(http.StatusBadRequest, gin.H{"error": "webauthn session expired"}) - return 0, "", false - } - - return sessionID, sessionDataStr, true -} - -func (h *AuthHandler) clearWebAuthnEnrollmentSession(sessionID int64) { - session, err := h.database.GetSessionByID(sessionID) - if err != nil { - return - } - - var sessionMeta map[string]interface{} - if len(session.Metadata) > 0 { - if err := json.Unmarshal(session.Metadata, &sessionMeta); err != nil { - return - } - } - - delete(sessionMeta, "webauthn_enrollment_session") - delete(sessionMeta, "webauthn_enrollment_expires") - if err := h.database.UpdateSessionMetadata(sessionID, sessionMeta); err != nil { - log.Printf("failed to clear webauthn session: %v", err) - } -} diff --git a/internal/gateway/handlers/auth_mfa_verification.go b/internal/gateway/handlers/auth_mfa_verification.go index b8db4c25..5c8aedac 100644 --- a/internal/gateway/handlers/auth_mfa_verification.go +++ b/internal/gateway/handlers/auth_mfa_verification.go @@ -81,20 +81,20 @@ func (h *AuthHandler) StartWebAuthnAssertion(c *gin.Context) { return } - options, sessionJSON, err := h.mfaService.StartWebAuthnAssertion(ctx.userRecord) - if err != nil { - c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) + if ctx.authUser.Session == nil { + c.JSON(http.StatusUnauthorized, gin.H{"error": "session missing"}) return } - if ctx.authUser.Session == nil { - c.JSON(http.StatusUnauthorized, gin.H{"error": "session missing"}) + options, challengeID, err := h.mfaService.StartWebAuthnAssertion(ctx.userRecord, &ctx.authUser.Session.ID) + if err != nil { + c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()}) return } c.JSON(http.StatusOK, WebAuthnAssertionStartResponse{ Options: options, - SessionData: string(sessionJSON), + SessionData: challengeID, }) } diff --git a/internal/gateway/handlers/dto.go b/internal/gateway/handlers/dto.go index 994b0dfd..baac7a25 100644 --- a/internal/gateway/handlers/dto.go +++ b/internal/gateway/handlers/dto.go @@ -66,15 +66,17 @@ type VerifyMFAResponse struct { TrustedDeviceCookie bool `json:"trusted_device_cookie" validate:"required"` } -// WebAuthnAssertionStartResponse contains WebAuthn assertion options and session data. +// WebAuthnAssertionStartResponse contains WebAuthn assertion options and the challenge ID. type WebAuthnAssertionStartResponse struct { Options interface{} `json:"options" validate:"required"` // protocol.CredentialAssertion - // SessionData is the serialized session to send back with verification + // SessionData is an opaque, single-use challenge ID to send back with the assertion. + // The challenge itself is stored server-side and expires after a few minutes. SessionData string `json:"session_data" validate:"required"` } // VerifyWebAuthnAssertionRequest contains the WebAuthn assertion response to verify. type VerifyWebAuthnAssertionRequest struct { + // SessionData is the challenge ID returned by /auth/mfa/webauthn/assertion/start. SessionData string `json:"session_data" binding:"required"` AssertionResponse string `json:"assertion_response" binding:"required"` TrustDevice bool `json:"trust_device"` diff --git a/internal/gateway/middleware/mfa_stepup_webauthn_test.go b/internal/gateway/middleware/mfa_stepup_webauthn_test.go index 3cdc74ce..e0aa77a1 100644 --- a/internal/gateway/middleware/mfa_stepup_webauthn_test.go +++ b/internal/gateway/middleware/mfa_stepup_webauthn_test.go @@ -46,14 +46,14 @@ func TestEnforceMFARequirements_AllowsInlineWebAuthn(t *testing.T) { mfaService, mfaSettings, sessionManager := setupMFAHelpers(t, database) session := confirmSession(t, sessionManager, user) - _, sessionData, err := mfaService.StartWebAuthnAssertion(user) + options, challengeID, err := mfaService.StartWebAuthnAssertion(user, &session.ID) require.NoError(t, err) - assertionJSON, err := credential.GenerateAssertionForSession(sessionData, "http://localhost") + assertionJSON, err := credential.GenerateAssertion(options, "http://localhost") require.NoError(t, err) inlinePayload := map[string]string{ - "session_data": string(sessionData), + "session_data": challengeID, "assertion_response": assertionJSON, } inlineBytes, err := json.Marshal(inlinePayload) diff --git a/internal/gateway/openapi/generated/swagger.json b/internal/gateway/openapi/generated/swagger.json index f05f76d5..a2a723c2 100644 --- a/internal/gateway/openapi/generated/swagger.json +++ b/internal/gateway/openapi/generated/swagger.json @@ -5394,6 +5394,7 @@ "type": "string" }, "session_data": { + "description": "SessionData is the challenge ID returned by /auth/mfa/webauthn/assertion/start.", "type": "string" }, "trust_device": { @@ -5427,7 +5428,7 @@ "description": "protocol.CredentialAssertion" }, "session_data": { - "description": "SessionData is the serialized session to send back with verification", + "description": "SessionData is an opaque, single-use challenge ID to send back with the assertion.\nThe challenge itself is stored server-side and expires after a few minutes.", "type": "string" } } diff --git a/internal/gateway/testutil/webauthntest/webauthn.go b/internal/gateway/testutil/webauthntest/webauthn.go index d45880d5..385358fe 100644 --- a/internal/gateway/testutil/webauthntest/webauthn.go +++ b/internal/gateway/testutil/webauthntest/webauthn.go @@ -8,14 +8,15 @@ import ( "crypto/sha256" "encoding/asn1" "encoding/base64" + "encoding/binary" "encoding/json" "fmt" "math/big" "strings" "github.com/fxamacker/cbor/v2" + "github.com/go-webauthn/webauthn/protocol" "github.com/go-webauthn/webauthn/protocol/webauthncose" - "github.com/go-webauthn/webauthn/webauthn" ) // MockCredential represents a mock WebAuthn credential for testing @@ -23,6 +24,10 @@ type MockCredential struct { ID []byte PublicKey []byte PrivateKey *ecdsa.PrivateKey + // Counter is the signature counter reported in generated assertions. + Counter uint32 + // WithoutUserVerification clears the UV flag (e.g. a key used without its PIN). + WithoutUserVerification bool } // GenerateMockCredential creates a mock WebAuthn credential with a valid key pair @@ -52,43 +57,38 @@ func GenerateMockCredential() (*MockCredential, error) { }, nil } -// GenerateAssertionForSession creates a valid WebAuthn assertion response for the -// provided session payload returned by StartWebAuthnAssertion. The origin should -// match the configured WebAuthn origin (e.g., "http://localhost"). -func (mc *MockCredential) GenerateAssertionForSession(sessionJSON []byte, origin string) (string, error) { - if len(sessionJSON) == 0 { - return "", fmt.Errorf("session data is required") +// GenerateAssertion creates a valid WebAuthn assertion response for the options returned by +// StartWebAuthnAssertion. The origin should match the configured WebAuthn origin (e.g. "http://localhost"). +// The signature counter in the authenticator data is mc.Counter. +func (mc *MockCredential) GenerateAssertion(options *protocol.CredentialAssertion, origin string) (string, error) { + if options == nil { + return "", fmt.Errorf("assertion options are required") } - - var session webauthn.SessionData - if err := json.Unmarshal(sessionJSON, &session); err != nil { - return "", fmt.Errorf("failed to unmarshal session: %w", err) - } - - if strings.TrimSpace(session.RelyingPartyID) == "" { - return "", fmt.Errorf("session missing relying party id") + request := options.Response + if strings.TrimSpace(request.RelyingPartyID) == "" { + return "", fmt.Errorf("options missing relying party id") } - - if _, err := base64.RawURLEncoding.DecodeString(session.Challenge); err != nil { - return "", fmt.Errorf("failed to decode challenge: %w", err) + if len(request.Challenge) == 0 { + return "", fmt.Errorf("options missing challenge") } credentialID := mc.ID - if len(session.AllowedCredentialIDs) > 0 { - credentialID = session.AllowedCredentialIDs[0] + if len(request.AllowedCredentials) > 0 { + credentialID = request.AllowedCredentials[0].CredentialID } if len(credentialID) == 0 { - return "", fmt.Errorf("session missing credential id") + return "", fmt.Errorf("options missing credential id") } if len(mc.ID) > 0 && !bytes.Equal(mc.ID, credentialID) { - return "", fmt.Errorf("mock credential does not match session credential") + return "", fmt.Errorf("mock credential does not match allowed credential") } if strings.TrimSpace(origin) == "" { - origin = fmt.Sprintf("http://%s", session.RelyingPartyID) + origin = fmt.Sprintf("http://%s", request.RelyingPartyID) } - assertion, err := mc.buildAssertion(session.Challenge, session.RelyingPartyID, credentialID, origin, session.UserID) + challenge := base64.RawURLEncoding.EncodeToString(request.Challenge) + assertion, err := mc.buildAssertion(challenge, request.RelyingPartyID, credentialID, origin, nil) if err != nil { return "", err } @@ -101,6 +101,16 @@ func (mc *MockCredential) GenerateAssertionForSession(sessionJSON []byte, origin return string(assertionBytes), nil } +// GenerateAssertionFromOptionsJSON is GenerateAssertion for the JSON "options" object returned by the +// /auth/mfa/webauthn/assertion/start endpoint. +func (mc *MockCredential) GenerateAssertionFromOptionsJSON(optionsJSON []byte, origin string) (string, error) { + var options protocol.CredentialAssertion + if err := json.Unmarshal(optionsJSON, &options); err != nil { + return "", fmt.Errorf("failed to unmarshal assertion options: %w", err) + } + return mc.GenerateAssertion(&options, origin) +} + func (mc *MockCredential) buildAssertion( challengeEncoded string, rpID string, @@ -112,6 +122,10 @@ func (mc *MockCredential) buildAssertion( authData := make([]byte, 37) copy(authData[0:32], rpIDHash[:]) authData[32] = 0x05 // user present + user verified + if mc.WithoutUserVerification { + authData[32] = 0x01 // user present only + } + binary.BigEndian.PutUint32(authData[33:37], mc.Counter) clientDataJSON := map[string]interface{}{ "type": "webauthn.get", diff --git a/internal/integration/integration_test.go b/internal/integration/integration_test.go index b6d4214c..ab9e1eed 100644 --- a/internal/integration/integration_test.go +++ b/internal/integration/integration_test.go @@ -864,8 +864,8 @@ func testCLIWebAuthnMFA(t *testing.T, s *TestServers) { require.Equal(t, http.StatusOK, startResp.StatusCode) var startResponse struct { - Options map[string]interface{} `json:"options"` - SessionData string `json:"session_data"` + Options json.RawMessage `json:"options"` + SessionData string `json:"session_data"` } err = json.NewDecoder(startResp.Body).Decode(&startResponse) startResp.Body.Close() @@ -874,7 +874,7 @@ func testCLIWebAuthnMFA(t *testing.T, s *TestServers) { // Step 2: Generate valid assertion using mock credential // Use the same origin as the gateway (http://localhost:8448 in dev mode) - assertionJSON, err := credential.GenerateAssertionForSession([]byte(startResponse.SessionData), "http://localhost:"+gatewayPort) + assertionJSON, err := credential.GenerateAssertionFromOptionsJSON(startResponse.Options, "http://localhost:"+gatewayPort) require.NoError(t, err, "failed to generate assertion") // Step 3: Format assertion the way CLI does (base64-encoded JSON with session_data and assertion_response)