diff options
| author | Chia <Chia@93.nz> | 2026-08-05 14:48:00 +1200 |
|---|---|---|
| committer | Chia <Chia@93.nz> | 2026-08-05 14:48:00 +1200 |
| commit | cd0dd91ab93653631904f2ea0e574ccde6d60339 (patch) | |
| tree | c65417b880a3f4a35c504c44edae821bc2122f70 /internal/adminapi | |
| parent | 86b1f42e3c5601ff10621a9779cf0076590797a1 (diff) | |
add passkey, totp.
Diffstat (limited to '')
| -rw-r--r-- | internal/adminapi/api.go | 566 |
1 files changed, 552 insertions, 14 deletions
diff --git a/internal/adminapi/api.go b/internal/adminapi/api.go index adbaea8..4e66d7f 100644 --- a/internal/adminapi/api.go +++ b/internal/adminapi/api.go @@ -1,6 +1,7 @@ package adminapi import ( + "bytes" "context" "crypto/rand" "crypto/sha256" @@ -8,6 +9,7 @@ import ( "encoding/hex" "encoding/json" "errors" + "io" "log/slog" "net" "net/http" @@ -19,6 +21,7 @@ import ( "aigw/internal/billing" "aigw/internal/controlplane" + "github.com/go-webauthn/webauthn/webauthn" "github.com/jackc/pgx/v5/pgconn" ) @@ -32,6 +35,9 @@ type API struct { registrationEnabled bool sessionTTL time.Duration currency string + publicURL string + webauthn *webauthn.WebAuthn + mailEnabled bool } type actorKey struct{} @@ -61,6 +67,9 @@ type Options struct { RegistrationEnabled bool SessionTTL time.Duration Currency string + PublicURL string + WebAuthn *webauthn.WebAuthn + MailEnabled bool } func New(options Options) *API { @@ -79,22 +88,47 @@ func New(options Options) *API { } return &API{store: options.Store, manager: options.Manager, billing: options.Billing, token: []byte(options.Token), logger: options.Logger, prefix: prefix, registrationEnabled: options.RegistrationEnabled, - sessionTTL: options.SessionTTL, currency: options.Currency} + sessionTTL: options.SessionTTL, currency: options.Currency, publicURL: strings.TrimRight(options.PublicURL, "/") + "/", + webauthn: options.WebAuthn, mailEnabled: options.MailEnabled} } func (a *API) Handler() http.Handler { mux := http.NewServeMux() apiPrefix := a.prefix + "/api" mux.HandleFunc("GET "+a.prefix, func(w http.ResponseWriter, r *http.Request) { - http.Redirect(w, r, a.prefix+"/", http.StatusTemporaryRedirect) + target := a.prefix + "/" + if r.URL.RawQuery != "" { + target += "?" + r.URL.RawQuery + } + http.Redirect(w, r, target, http.StatusTemporaryRedirect) }) mux.Handle(a.prefix+"/", http.StripPrefix(a.prefix, adminui.Handler())) mux.HandleFunc("GET "+apiPrefix+"/auth/config", a.public(a.authConfig)) mux.HandleFunc("GET "+apiPrefix+"/auth/session", a.public(a.authSession)) mux.HandleFunc("POST "+apiPrefix+"/auth/register", a.public(a.register)) mux.HandleFunc("POST "+apiPrefix+"/auth/login", a.public(a.login)) + mux.HandleFunc("POST "+apiPrefix+"/auth/email/verify", a.public(a.verifyEmail)) + mux.HandleFunc("POST "+apiPrefix+"/auth/email/resend", a.public(a.resendVerification)) + mux.HandleFunc("POST "+apiPrefix+"/auth/password-reset/request", a.public(a.requestPasswordReset)) + mux.HandleFunc("POST "+apiPrefix+"/auth/password-reset/complete", a.public(a.completePasswordReset)) + mux.HandleFunc("POST "+apiPrefix+"/auth/invite/accept", a.public(a.acceptInvite)) + mux.HandleFunc("POST "+apiPrefix+"/auth/mfa/totp", a.public(a.completeMFA)) + mux.HandleFunc("POST "+apiPrefix+"/auth/passkey/login/options", a.public(a.passkeyLoginOptions)) + mux.HandleFunc("POST "+apiPrefix+"/auth/passkey/login", a.public(a.passkeyLogin)) + mux.HandleFunc("POST "+apiPrefix+"/auth/mfa/passkey/login/options", a.public(a.beginPasskeyMFA)) + mux.HandleFunc("POST "+apiPrefix+"/auth/mfa/passkey/login", a.public(a.finishPasskeyMFA)) mux.HandleFunc("POST "+apiPrefix+"/auth/logout", a.withAuth("overview.read", a.logout)) mux.HandleFunc("POST "+apiPrefix+"/auth/password", a.withAuth("overview.read", a.changePassword)) + mux.HandleFunc("GET "+apiPrefix+"/auth/mfa", a.withAuth("overview.read", a.mfaStatus)) + mux.HandleFunc("POST "+apiPrefix+"/auth/mfa/totp/begin", a.withAuth("overview.read", a.beginTOTP)) + mux.HandleFunc("POST "+apiPrefix+"/auth/mfa/totp/confirm", a.withAuth("overview.read", a.confirmTOTP)) + mux.HandleFunc("POST "+apiPrefix+"/auth/mfa/totp/disable", a.withAuth("overview.read", a.disableTOTP)) + mux.HandleFunc("POST "+apiPrefix+"/auth/mfa/passkey/options", a.withAuth("overview.read", a.beginPasskeyRegistration)) + mux.HandleFunc("POST "+apiPrefix+"/auth/mfa/passkey", a.withAuth("overview.read", a.finishPasskeyRegistration)) + mux.HandleFunc("GET "+apiPrefix+"/auth/sessions", a.withAuth("overview.read", a.listSessions)) + mux.HandleFunc("POST "+apiPrefix+"/auth/sessions/{id}/revoke", a.withAuth("overview.read", a.revokeSession)) + mux.HandleFunc("POST "+apiPrefix+"/auth/sessions/revoke-others", a.withAuth("overview.read", a.revokeOtherSessions)) + mux.HandleFunc("POST "+apiPrefix+"/auth/passkeys/{id}/delete", a.withAuth("overview.read", a.deletePasskey)) mux.HandleFunc("GET "+apiPrefix+"/overview", a.withAuth("overview.read", a.overview)) mux.HandleFunc("GET "+apiPrefix+"/tenants", a.withAuth("tenants.read", a.listTenants)) @@ -114,6 +148,8 @@ func (a *API) Handler() http.Handler { if a.billing != nil { mux.HandleFunc("GET "+apiPrefix+"/billing/accounts", a.withAuth("billing.read", a.listBillingAccounts)) mux.HandleFunc("GET "+apiPrefix+"/billing/ledger", a.withAuth("billing.read", a.listBillingLedger)) + mux.HandleFunc("GET "+apiPrefix+"/billing/orders", a.withAuth("billing.read", a.listTopUpOrders)) + mux.HandleFunc("GET "+apiPrefix+"/billing/orders/{id}", a.withAuth("billing.read", a.getTopUpOrder)) mux.HandleFunc("POST "+apiPrefix+"/billing/adjustments", a.withAuth("billing.adjust", a.adjustBalance)) mux.HandleFunc("POST "+apiPrefix+"/billing/checkout-sessions", a.withAuth("billing.topup", a.createCheckoutSession)) } @@ -149,7 +185,7 @@ func (a *API) securityHeaders(next http.Handler) http.Handler { w.Header().Set("Referrer-Policy", "no-referrer") w.Header().Set("X-Content-Type-Options", "nosniff") w.Header().Set("X-Frame-Options", "DENY") - w.Header().Set("Permissions-Policy", "camera=(), microphone=(), geolocation=(), payment=()") + w.Header().Set("Permissions-Policy", "camera=(), microphone=(), geolocation=(), payment=(), publickey-credentials-create=(self), publickey-credentials-get=(self)") if strings.Contains(r.URL.Path, "/api/") { w.Header().Set("Cache-Control", "no-store") } @@ -239,7 +275,8 @@ func adminRequestID(r *http.Request) string { } func (a *API) authConfig(w http.ResponseWriter, _ *http.Request) { - writeJSON(w, map[string]any{"registration_enabled": a.registrationEnabled}) + writeJSON(w, map[string]any{"registration_enabled": a.registrationEnabled, "email_delivery_enabled": a.mailEnabled, + "passkeys_enabled": a.webauthn != nil}) } func (a *API) authSession(w http.ResponseWriter, r *http.Request) { @@ -270,7 +307,11 @@ func (a *API) register(w http.ResponseWriter, r *http.Request) { if !decodeBody(w, r, &input) { return } - actor, generation, err := a.store.RegisterTenant(r.Context(), input, a.currency) + if !a.consumePublicRateLimit(w, r, "register_ip", remoteIP(r), 5, time.Hour) || + !a.consumePublicRateLimit(w, r, "register_email", input.Email, 3, 24*time.Hour) { + return + } + actor, generation, err := a.store.RegisterTenantPending(r.Context(), input, a.currency, a.publicURL, remoteIP(r)) if err != nil { a.writeAudit(r, controlplane.ConsoleActor{}, "auth.register", http.StatusBadRequest) a.mutationError(w, r, err) @@ -281,14 +322,9 @@ func (a *API) register(w http.ResponseWriter, r *http.Request) { a.logger.Warn("registration_snapshot_reload_failed", "tenant_id", actor.TenantID, "error", err) } } - session, err := a.store.CreateConsoleSession(r.Context(), actor, a.sessionTTL, remoteIP(r), r.UserAgent()) - if err != nil { - a.databaseError(w, r, err) - return - } - a.setSessionCookie(w, r, session) a.writeAudit(r, actor, "auth.register", http.StatusCreated) - writeStatusJSON(w, http.StatusCreated, sessionPayload(session)) + writeStatusJSON(w, http.StatusCreated, map[string]any{"status": "verification_required", "email": actor.Email, + "delivery": map[bool]string{true: "queued", false: "pending_configuration"}[a.mailEnabled]}) } func (a *API) login(w http.ResponseWriter, r *http.Request) { @@ -316,7 +352,18 @@ func (a *API) login(w http.ResponseWriter, r *http.Request) { apierror.Write(w, apierror.Error{Status: status, Type: typeName, Message: message}, requestID(r)) return } - session, err := a.store.CreateConsoleSession(r.Context(), actor, a.sessionTTL, remoteIP(r), r.UserAgent()) + if len(actor.MFAMethods) > 0 { + challenge, err := a.store.BeginMFAChallenge(r.Context(), actor, remoteIP(r), r.UserAgent()) + if err != nil { + a.databaseError(w, r, err) + return + } + a.writeAudit(r, actor, "auth.login.password", http.StatusAccepted) + writeStatusJSON(w, http.StatusAccepted, map[string]any{"mfa_required": true, "challenge_token": challenge.Token, + "methods": challenge.Methods, "expires_at": challenge.ExpiresAt}) + return + } + session, err := a.store.CreateConsoleSessionWithMethod(r.Context(), actor, a.sessionTTL, remoteIP(r), r.UserAgent(), "password", false) if err != nil { a.databaseError(w, r, err) return @@ -326,6 +373,218 @@ func (a *API) login(w http.ResponseWriter, r *http.Request) { writeJSON(w, sessionPayload(session)) } +func (a *API) verifyEmail(w http.ResponseWriter, r *http.Request) { + var input struct { + Token string `json:"token"` + } + if !decodeBody(w, r, &input) { + return + } + actor, err := a.store.VerifyEmailToken(r.Context(), input.Token) + if err != nil { + a.actionError(w, r, err) + return + } + session, err := a.store.CreateConsoleSessionWithMethod(r.Context(), actor, a.sessionTTL, remoteIP(r), r.UserAgent(), "email_link", false) + if err != nil { + a.databaseError(w, r, err) + return + } + a.setSessionCookie(w, r, session) + a.writeAudit(r, actor, "auth.email.verify", http.StatusOK) + writeJSON(w, sessionPayload(session)) +} + +func (a *API) resendVerification(w http.ResponseWriter, r *http.Request) { + var input struct { + Email string `json:"email"` + } + if !decodeBody(w, r, &input) { + return + } + if !a.consumePublicRateLimit(w, r, "verify_resend_ip", remoteIP(r), 10, time.Hour) || + !a.consumePublicRateLimit(w, r, "verify_resend_email", input.Email, 3, time.Hour) { + return + } + if err := a.store.ResendVerification(r.Context(), input.Email, a.publicURL, remoteIP(r)); err != nil { + a.databaseError(w, r, err) + return + } + writeStatusJSON(w, http.StatusAccepted, map[string]any{"status": "queued"}) +} + +func (a *API) requestPasswordReset(w http.ResponseWriter, r *http.Request) { + var input struct { + Email string `json:"email"` + } + if !decodeBody(w, r, &input) { + return + } + if !a.consumePublicRateLimit(w, r, "password_reset_ip", remoteIP(r), 20, time.Hour) || + !a.consumePublicRateLimit(w, r, "password_reset_email", input.Email, 3, time.Hour) { + return + } + if err := a.store.RequestPasswordReset(r.Context(), input.Email, a.publicURL, remoteIP(r)); err != nil { + a.databaseError(w, r, err) + return + } + writeStatusJSON(w, http.StatusAccepted, map[string]any{"status": "queued"}) +} + +func (a *API) completePasswordReset(w http.ResponseWriter, r *http.Request) { + var input controlplane.PasswordResetInput + if !decodeBody(w, r, &input) { + return + } + if err := a.store.ResetPassword(r.Context(), input); err != nil { + a.actionError(w, r, err) + return + } + writeJSON(w, map[string]any{"status": "password_reset", "reauthentication_required": true}) +} + +func (a *API) acceptInvite(w http.ResponseWriter, r *http.Request) { + var input controlplane.InviteAcceptInput + if !decodeBody(w, r, &input) { + return + } + actor, err := a.store.AcceptInvite(r.Context(), input) + if err != nil { + a.actionError(w, r, err) + return + } + session, err := a.store.CreateConsoleSessionWithMethod(r.Context(), actor, a.sessionTTL, remoteIP(r), r.UserAgent(), "invite_link", false) + if err != nil { + a.databaseError(w, r, err) + return + } + a.setSessionCookie(w, r, session) + a.writeAudit(r, actor, "auth.invite.accept", http.StatusOK) + writeJSON(w, sessionPayload(session)) +} + +func (a *API) completeMFA(w http.ResponseWriter, r *http.Request) { + var input controlplane.MFACodeInput + if !decodeBody(w, r, &input) { + return + } + actor, method, err := a.store.CompleteMFAChallenge(r.Context(), input, remoteIP(r)) + if err != nil { + a.actionError(w, r, err) + return + } + session, err := a.store.CreateConsoleSessionWithMethod(r.Context(), actor, a.sessionTTL, remoteIP(r), r.UserAgent(), "password+"+method, true) + if err != nil { + a.databaseError(w, r, err) + return + } + a.setSessionCookie(w, r, session) + a.writeAudit(r, actor, "auth.login."+method, http.StatusOK) + writeJSON(w, sessionPayload(session)) +} + +func (a *API) passkeyLoginOptions(w http.ResponseWriter, r *http.Request) { + if !a.requireWebAuthn(w, r) { + return + } + var input struct { + Email string `json:"email"` + } + if !decodeBody(w, r, &input) { + return + } + if !a.consumePublicRateLimit(w, r, "passkey_login_ip", remoteIP(r), 30, 15*time.Minute) || + !a.consumePublicRateLimit(w, r, "passkey_login_email", input.Email, 10, 15*time.Minute) { + return + } + actor, err := a.store.FindActiveConsoleUser(r.Context(), input.Email) + if err != nil { + a.passkeyAuthError(w, r, err) + return + } + token, options, err := a.store.BeginWebAuthnLogin(r.Context(), actor.ID, "login", "", a.webauthn) + if err != nil { + a.passkeyAuthError(w, r, err) + return + } + writeJSON(w, map[string]any{"challenge_token": token, "options": options}) +} + +func (a *API) passkeyLogin(w http.ResponseWriter, r *http.Request) { + a.finishPasskeyLogin(w, r, "login", "passkey") +} + +func (a *API) beginPasskeyMFA(w http.ResponseWriter, r *http.Request) { + if !a.requireWebAuthn(w, r) { + return + } + var input struct { + ChallengeToken string `json:"challenge_token"` + } + if !decodeBody(w, r, &input) { + return + } + authChallengeID, userID, err := a.store.ResolveMFAChallenge(r.Context(), input.ChallengeToken, remoteIP(r)) + if err != nil { + a.actionError(w, r, err) + return + } + token, options, err := a.store.BeginWebAuthnLogin(r.Context(), userID, "mfa_login", authChallengeID, a.webauthn) + if err != nil { + a.passkeyAuthError(w, r, err) + return + } + writeJSON(w, map[string]any{"challenge_token": token, "options": options}) +} + +func (a *API) finishPasskeyMFA(w http.ResponseWriter, r *http.Request) { + a.finishPasskeyLogin(w, r, "mfa_login", "password+passkey") +} + +type webAuthnFinishInput struct { + ChallengeToken string `json:"challenge_token"` + Name string `json:"name"` + Credential json.RawMessage `json:"credential"` +} + +func (a *API) finishPasskeyLogin(w http.ResponseWriter, r *http.Request, purpose, authMethod string) { + if !a.requireWebAuthn(w, r) { + return + } + var input webAuthnFinishInput + if !decodeBody(w, r, &input) { + return + } + userID, authChallengeID, sessionData, err := a.store.WebAuthnSession(r.Context(), input.ChallengeToken, purpose) + if err != nil { + a.actionError(w, r, err) + return + } + user, err := a.store.WebAuthnUser(r.Context(), userID) + if err != nil { + a.passkeyAuthError(w, r, err) + return + } + credential, err := a.webauthn.FinishLogin(user, sessionData, credentialRequest(r, input.Credential)) + if err != nil { + a.passkeyAuthError(w, r, err) + return + } + actor, err := a.store.FinishWebAuthnLogin(r.Context(), input.ChallengeToken, userID, authChallengeID, credential, purpose) + if err != nil { + a.passkeyAuthError(w, r, err) + return + } + session, err := a.store.CreateConsoleSessionWithMethod(r.Context(), actor, a.sessionTTL, remoteIP(r), r.UserAgent(), authMethod, true) + if err != nil { + a.databaseError(w, r, err) + return + } + a.setSessionCookie(w, r, session) + a.writeAudit(r, actor, "auth.login.passkey", http.StatusOK) + writeJSON(w, sessionPayload(session)) +} + func (a *API) logout(w http.ResponseWriter, r *http.Request) { if cookie, err := r.Cookie("aigw_session"); err == nil { if err := a.store.RevokeConsoleSession(r.Context(), cookie.Value); err != nil { @@ -359,10 +618,201 @@ func (a *API) changePassword(w http.ResponseWriter, r *http.Request) { writeJSON(w, map[string]any{"status": "password_changed", "reauthentication_required": true}) } +func (a *API) mfaStatus(w http.ResponseWriter, r *http.Request) { + actor := a.actor(r) + if actor.ID == "" { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "bootstrap_account", Message: "Bootstrap access has no MFA profile"}, requestID(r)) + return + } + result, err := a.store.MFAStatus(r.Context(), actor.ID) + if err != nil { + a.databaseError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) beginTOTP(w http.ResponseWriter, r *http.Request) { + actor := a.actor(r) + var input struct { + CurrentPassword string `json:"current_password"` + } + if !decodeBody(w, r, &input) { + return + } + if actor.ID == "" || a.store.VerifyConsolePassword(r.Context(), actor.ID, input.CurrentPassword) != nil { + apierror.Write(w, apierror.Error{Status: http.StatusUnauthorized, Type: "invalid_credentials", Message: "Current password is incorrect"}, requestID(r)) + return + } + result, err := a.store.BeginTOTP(r.Context(), actor) + if err != nil { + a.mutationError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) confirmTOTP(w http.ResponseWriter, r *http.Request) { + var input struct { + Code string `json:"code"` + } + if !decodeBody(w, r, &input) { + return + } + codes, err := a.store.ConfirmTOTP(r.Context(), a.actor(r).ID, input.Code) + if err != nil { + a.actionError(w, r, err) + return + } + writeJSON(w, map[string]any{"status": "enabled", "recovery_codes": codes}) +} + +func (a *API) disableTOTP(w http.ResponseWriter, r *http.Request) { + actor := a.actor(r) + var input struct { + CurrentPassword string `json:"current_password"` + Code string `json:"code"` + } + if !decodeBody(w, r, &input) { + return + } + if a.store.VerifyConsolePassword(r.Context(), actor.ID, input.CurrentPassword) != nil || + a.store.VerifyTOTPForUser(r.Context(), actor.ID, input.Code) != nil { + apierror.Write(w, apierror.Error{Status: http.StatusUnauthorized, Type: "invalid_credentials", Message: "Password or authenticator code is incorrect"}, requestID(r)) + return + } + if err := a.store.DisableTOTP(r.Context(), actor.ID); err != nil { + a.databaseError(w, r, err) + return + } + writeJSON(w, map[string]any{"status": "disabled"}) +} + +func (a *API) beginPasskeyRegistration(w http.ResponseWriter, r *http.Request) { + if !a.requireWebAuthn(w, r) { + return + } + actor := a.actor(r) + var input struct { + CurrentPassword string `json:"current_password"` + } + if !decodeBody(w, r, &input) { + return + } + if actor.ID == "" || a.store.VerifyConsolePassword(r.Context(), actor.ID, input.CurrentPassword) != nil { + apierror.Write(w, apierror.Error{Status: http.StatusUnauthorized, Type: "invalid_credentials", Message: "Current password is incorrect"}, requestID(r)) + return + } + token, options, err := a.store.BeginWebAuthnRegistration(r.Context(), actor.ID, a.webauthn) + if err != nil { + a.passkeyAuthError(w, r, err) + return + } + writeJSON(w, map[string]any{"challenge_token": token, "options": options}) +} + +func (a *API) finishPasskeyRegistration(w http.ResponseWriter, r *http.Request) { + if !a.requireWebAuthn(w, r) { + return + } + var input webAuthnFinishInput + if !decodeBody(w, r, &input) { + return + } + actor := a.actor(r) + userID, _, sessionData, err := a.store.WebAuthnSession(r.Context(), input.ChallengeToken, "register") + if err != nil || userID != actor.ID { + a.actionError(w, r, controlplane.ErrActionTokenInvalid) + return + } + user, err := a.store.WebAuthnUser(r.Context(), actor.ID) + if err != nil { + a.passkeyAuthError(w, r, err) + return + } + credential, err := a.webauthn.FinishRegistration(user, sessionData, credentialRequest(r, input.Credential)) + if err != nil { + a.passkeyAuthError(w, r, err) + return + } + if _, err := a.store.FinishWebAuthnRegistration(r.Context(), input.ChallengeToken, actor.ID, input.Name, credential); err != nil { + a.passkeyAuthError(w, r, err) + return + } + writeStatusJSON(w, http.StatusCreated, map[string]any{"status": "registered"}) +} + +func (a *API) deletePasskey(w http.ResponseWriter, r *http.Request) { + actor := a.actor(r) + var input struct { + CurrentPassword string `json:"current_password"` + } + if !decodeBody(w, r, &input) { + return + } + if a.store.VerifyConsolePassword(r.Context(), actor.ID, input.CurrentPassword) != nil { + apierror.Write(w, apierror.Error{Status: http.StatusUnauthorized, Type: "invalid_credentials", Message: "Current password is incorrect"}, requestID(r)) + return + } + if err := a.store.DeletePasskey(r.Context(), actor.ID, r.PathValue("id")); err != nil { + a.mutationError(w, r, err) + return + } + writeJSON(w, map[string]any{"status": "deleted"}) +} + +func (a *API) listSessions(w http.ResponseWriter, r *http.Request) { + actor := a.actor(r) + if actor.ID == "" { + writeJSON(w, []controlplane.DeviceSession{}) + return + } + result, err := a.store.ListDeviceSessions(r.Context(), actor.ID, sessionCookie(r)) + if err != nil { + a.databaseError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) revokeSession(w http.ResponseWriter, r *http.Request) { + actor := a.actor(r) + current := sessionCookie(r) + sessions, err := a.store.ListDeviceSessions(r.Context(), actor.ID, current) + if err != nil { + a.databaseError(w, r, err) + return + } + currentRevoked := false + for _, session := range sessions { + if session.ID == r.PathValue("id") { + currentRevoked = session.Current + break + } + } + if _, err := a.store.RevokeDeviceSession(r.Context(), actor.ID, r.PathValue("id")); err != nil { + a.mutationError(w, r, err) + return + } + if currentRevoked { + a.clearSessionCookie(w, r) + } + writeJSON(w, map[string]any{"status": "revoked", "current": currentRevoked}) +} + +func (a *API) revokeOtherSessions(w http.ResponseWriter, r *http.Request) { + if err := a.store.RevokeOtherDeviceSessions(r.Context(), a.actor(r).ID, sessionCookie(r)); err != nil { + a.databaseError(w, r, err) + return + } + writeJSON(w, map[string]any{"status": "revoked"}) +} + func sessionPayload(session controlplane.ConsoleSession) map[string]any { return map[string]any{ "actor": session.Actor, "permissions": session.Actor.Permissions(), "csrf_token": session.CSRFToken, "expires_at": session.ExpiresAt, + "session_id": session.ID, "auth_method": session.AuthMethod, } } @@ -385,6 +835,62 @@ func requestIsHTTPS(r *http.Request) bool { return r.TLS != nil || strings.EqualFold(strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")), "https") } +func sessionCookie(r *http.Request) string { + if cookie, err := r.Cookie("aigw_session"); err == nil { + return cookie.Value + } + return "" +} + +func credentialRequest(r *http.Request, credential json.RawMessage) *http.Request { + request := r.Clone(r.Context()) + request.Body = io.NopCloser(bytes.NewReader(credential)) + request.ContentLength = int64(len(credential)) + request.Header = r.Header.Clone() + request.Header.Set("Content-Type", "application/json") + return request +} + +func (a *API) requireWebAuthn(w http.ResponseWriter, r *http.Request) bool { + if a.webauthn != nil { + return true + } + apierror.Write(w, apierror.Error{Status: http.StatusNotImplemented, Type: "passkeys_disabled", Message: "Passkeys are not configured for this console"}, requestID(r)) + return false +} + +func (a *API) consumePublicRateLimit(w http.ResponseWriter, r *http.Request, scope, key string, limit int, window time.Duration) bool { + if err := a.store.ConsumeRateLimit(r.Context(), scope, key, limit, window); err != nil { + if errors.Is(err, controlplane.ErrConsoleRateLimited) { + w.Header().Set("Retry-After", "900") + apierror.Write(w, apierror.Error{Status: http.StatusTooManyRequests, Type: "rate_limited", Message: "Too many requests; try again later"}, requestID(r)) + return false + } + a.databaseError(w, r, err) + return false + } + return true +} + +func (a *API) actionError(w http.ResponseWriter, r *http.Request, err error) { + if errors.Is(err, controlplane.ErrActionTokenInvalid) { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "invalid_or_expired_token", Message: "This link or authentication challenge is invalid or expired"}, requestID(r)) + return + } + if errors.Is(err, controlplane.ErrConsoleUnauthorized) { + apierror.Write(w, apierror.Error{Status: http.StatusUnauthorized, Type: "invalid_authenticator", Message: "The authentication code is incorrect"}, requestID(r)) + return + } + a.mutationError(w, r, err) +} + +func (a *API) passkeyAuthError(w http.ResponseWriter, r *http.Request, err error) { + if !errors.Is(err, controlplane.ErrConsoleUnauthorized) && !errors.Is(err, controlplane.ErrNoPasskeys) && !errors.Is(err, controlplane.ErrActionTokenInvalid) { + a.logger.Warn("passkey_authentication_failed", "error", err) + } + apierror.Write(w, apierror.Error{Status: http.StatusUnauthorized, Type: "passkey_failed", Message: "Passkey authentication could not be completed"}, requestID(r)) +} + func isUnsafeMethod(method string) bool { return method != http.MethodGet && method != http.MethodHead && method != http.MethodOptions } @@ -437,6 +943,34 @@ func (a *API) listBillingLedger(w http.ResponseWriter, r *http.Request) { writeJSON(w, result) } +func (a *API) listTopUpOrders(w http.ResponseWriter, r *http.Request) { + actor := a.actor(r) + if actor.TenantID == "" { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "A tenant is required"}, requestID(r)) + return + } + result, err := a.billing.ListTopUpOrders(r.Context(), actor.TenantID, 50) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + +func (a *API) getTopUpOrder(w http.ResponseWriter, r *http.Request) { + actor := a.actor(r) + if actor.TenantID == "" { + apierror.Write(w, apierror.Error{Status: http.StatusBadRequest, Type: "tenant_required", Message: "A tenant is required"}, requestID(r)) + return + } + result, err := a.billing.GetTopUpOrder(r.Context(), actor.TenantID, r.PathValue("id")) + if err != nil { + a.billingError(w, r, err) + return + } + writeJSON(w, result) +} + func (a *API) adjustBalance(w http.ResponseWriter, r *http.Request) { var input billing.AdjustmentInput if !decodeBody(w, r, &input) { @@ -730,7 +1264,7 @@ func (a *API) createUser(w http.ResponseWriter, r *http.Request) { return } } - result, err := a.store.CreateConsoleUser(r.Context(), input) + result, err := a.store.InviteConsoleUser(r.Context(), input, a.publicURL, remoteIP(r)) if err != nil { a.mutationError(w, r, err) return @@ -837,6 +1371,10 @@ func (a *API) billingError(w http.ResponseWriter, r *http.Request, err error) { status = http.StatusServiceUnavailable typeName = "stripe_disabled" message = "Stripe top-ups are disabled" + case errors.Is(err, billing.ErrTopUpOrderNotFound): + status = http.StatusNotFound + typeName = "topup_order_not_found" + message = "Top-up order was not found" default: a.logger.Error("admin_billing_error", "error", err) } |
