Files
bd2/go/internal/server/auth/service.go
T

1027 lines
38 KiB
Go

package auth
import (
"bytes"
"context"
"crypto/sha256"
"crypto/subtle"
"database/sql"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"log/slog"
"net"
"net/http"
"net/url"
"strconv"
"strings"
"sync"
"time"
"bd2server/internal/server/authconfig"
"bd2server/internal/server/wire"
)
type Service struct {
config authconfig.Runtime
store *Store
client *http.Client
limits requestLimiter
}
type limitWindow struct {
started time.Time
count int
}
type requestLimiter struct {
mu sync.Mutex
windows map[string]limitWindow
lastSweep time.Time
}
type deviceResult struct {
Provider string `json:"provider"`
AccessToken string `json:"access_token"`
AccessExpiresIn int64 `json:"access_expires_in"`
RefreshToken string `json:"refresh_token"`
RefreshExpiresIn int64 `json:"refresh_expires_in"`
}
type refreshAttemptResult struct {
Provider string `json:"provider"`
AccessToken string `json:"access_token"`
AccessExpiresAt int64 `json:"access_expires_at"`
RefreshToken string `json:"refresh_token"`
RefreshExpiresAt int64 `json:"refresh_expires_at"`
}
func New(config authconfig.Runtime, store *Store) (*Service, error) {
if config.Mode != "oauth" || store == nil {
return nil, errors.New("auth: OAuth service requires oauth configuration and store")
}
// Store.Open has already derived its purpose-specific keys. Do not retain
// the environment master key in the long-lived HTTP service configuration.
clear(config.MasterKey)
config.MasterKey = nil
return &Service{config: config, store: store, client: &http.Client{Timeout: 15 * time.Second}, limits: requestLimiter{windows: make(map[string]limitWindow)}}, nil
}
func (s *Service) Handler() http.Handler {
mux := http.NewServeMux()
mux.HandleFunc("POST /auth/device", s.createDevice)
mux.HandleFunc("GET /auth/{provider}/start", s.start)
mux.HandleFunc("GET /auth/{provider}/callback", s.callback)
mux.HandleFunc("POST /auth/device/{id}/poll", s.poll)
mux.HandleFunc("POST /auth/session/refresh", s.refresh)
mux.HandleFunc("POST /auth/session/revoke", s.revoke)
return securityHeaders(mux)
}
func securityHeaders(next http.Handler) http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("X-Content-Type-Options", "nosniff")
w.Header().Set("Referrer-Policy", "no-referrer")
w.Header().Set("Content-Security-Policy", "default-src 'none'; frame-ancestors 'none'")
w.Header().Set("X-Frame-Options", "DENY")
next.ServeHTTP(w, r)
})
}
func decodeJSON(w http.ResponseWriter, r *http.Request, target any) bool {
defer func() { _ = r.Body.Close() }()
data, err := io.ReadAll(io.LimitReader(r.Body, 16<<10+1))
if err != nil || len(data) > 16<<10 {
http.Error(w, "request too large", http.StatusRequestEntityTooLarge)
return false
}
decoder := json.NewDecoder(bytes.NewReader(data))
decoder.DisallowUnknownFields()
if err := decoder.Decode(target); err != nil {
http.Error(w, "invalid JSON", http.StatusBadRequest)
return false
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
http.Error(w, "invalid JSON", http.StatusBadRequest)
return false
}
return true
}
func writeJSON(w http.ResponseWriter, status int, value any) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.WriteHeader(status)
_ = json.NewEncoder(w).Encode(value)
}
func (s *Service) createDevice(w http.ResponseWriter, r *http.Request) {
clientIP := remoteIP(r.RemoteAddr)
if !s.limits.allow("create:"+clientIP, s.store.now(), time.Minute, 10) {
w.Header().Set("Retry-After", "60")
http.Error(w, "too many login attempts", http.StatusTooManyRequests)
return
}
var request struct {
Provider string `json:"provider"`
}
if !decodeJSON(w, r, &request) {
return
}
if _, ok := s.config.Providers[request.Provider]; !ok {
http.Error(w, "provider is not enabled", http.StatusBadRequest)
return
}
id, err := randomToken(18)
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
secret, err := randomToken(32)
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
startTicket, err := randomToken(32)
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
now := s.store.now()
clientHash := s.store.digest("client-ip", clientIP)
tx, err := s.store.db.Begin()
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
defer func() { _ = tx.Rollback() }()
if err := cleanupExpired(tx, now.Unix()); err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
var pending int
if err := tx.QueryRow(`SELECT COUNT(*) FROM devices WHERE client_hash=? AND status IN ('created','authorizing') AND expires_at>?`, clientHash, now.Unix()).Scan(&pending); err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
if pending >= 5 {
w.Header().Set("Retry-After", strconv.FormatInt(int64(s.config.DeviceTTL.Seconds()), 10))
http.Error(w, "too many pending login transactions", http.StatusTooManyRequests)
return
}
_, err = tx.Exec(`INSERT INTO devices(id,client_hash,secret_hash,start_hash,provider,status,created_at,expires_at) VALUES(?,?,?,?,?,'created',?,?)`, id, clientHash, s.store.digest("device-secret", secret), s.store.digest("start-ticket", startTicket), request.Provider, now.Unix(), now.Add(s.config.DeviceTTL).Unix())
if err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
if err := tx.Commit(); err != nil {
http.Error(w, "could not create transaction", http.StatusInternalServerError)
return
}
start := *s.config.PublicURLParsed
start.Path = "/auth/" + request.Provider + "/start"
query := start.Query()
query.Set("transaction_id", id)
query.Set("ticket", startTicket)
start.RawQuery = query.Encode()
writeJSON(w, http.StatusCreated, map[string]any{"transaction_id": id, "device_secret": secret, "start_url": start.String(), "expires_in": int64(s.config.DeviceTTL.Seconds()), "poll_interval": 2})
}
func (s *Service) start(w http.ResponseWriter, r *http.Request) {
provider := r.PathValue("provider")
if _, ok := s.config.Providers[provider]; !ok {
http.Error(w, "provider is not enabled", http.StatusNotFound)
return
}
id, ticket := r.URL.Query().Get("transaction_id"), r.URL.Query().Get("ticket")
var storedHash []byte
var storedProvider, status string
var expires int64
err := s.store.db.QueryRow(`SELECT start_hash,provider,status,expires_at FROM devices WHERE id=?`, id).Scan(&storedHash, &storedProvider, &status, &expires)
if err != nil || subtle.ConstantTimeCompare(storedHash, s.store.digest("start-ticket", ticket)) != 1 || storedProvider != provider || status != "created" {
http.Error(w, "invalid login transaction", http.StatusForbidden)
return
}
if s.store.now().Unix() >= expires {
http.Error(w, "login transaction expired", http.StatusGone)
return
}
state, err := randomToken(32)
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
verifier, err := randomToken(32)
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
nonce, err := randomToken(24)
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
verifierCipher, err := s.store.seal(id, "pkce", []byte(verifier))
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
nonceCipher, err := s.store.seal(id, "nonce", []byte(nonce))
if err != nil {
http.Error(w, "could not start authorization", http.StatusInternalServerError)
return
}
result, err := s.store.db.Exec(`UPDATE devices SET state_hash=?,verifier_cipher=?,nonce_cipher=?,start_hash=X'',status='authorizing' WHERE id=? AND status='created'`, s.store.digest("oauth-state", state), verifierCipher, nonceCipher, id)
count, affectedErr := rowsAffected(result)
if err != nil || affectedErr != nil || count != 1 {
http.Error(w, "could not start authorization", http.StatusConflict)
return
}
redirect := s.redirectURL(provider)
challenge := sha256.Sum256([]byte(verifier))
values := url.Values{"client_id": {s.config.Providers[provider].ClientID}, "redirect_uri": {redirect}, "response_type": {"code"}, "scope": {providerScope(provider)}, "state": {state}, "code_challenge": {base64.RawURLEncoding.EncodeToString(challenge[:])}, "code_challenge_method": {"S256"}}
if provider == "google" {
values.Set("nonce", nonce)
}
http.Redirect(w, r, providerAuthorizeURL(provider)+"?"+values.Encode(), http.StatusFound)
}
func (s *Service) callback(w http.ResponseWriter, r *http.Request) {
provider, state, code := r.PathValue("provider"), r.URL.Query().Get("state"), r.URL.Query().Get("code")
if _, ok := s.config.Providers[provider]; !ok {
http.Error(w, "provider is not enabled", http.StatusNotFound)
return
}
if state == "" {
http.Error(w, "authorization was not completed", http.StatusBadRequest)
return
}
if r.URL.Query().Get("error") != "" {
result, err := s.store.db.Exec(`UPDATE devices SET status='failed',error_code='provider_cancelled',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL
WHERE state_hash=? AND provider=? AND status='authorizing' AND expires_at>?`, s.store.digest("oauth-state", state), provider, s.store.now().Unix())
if err != nil {
http.Error(w, "authorization state unavailable", http.StatusInternalServerError)
return
}
if count, err := rowsAffected(result); err != nil || count != 1 {
http.Error(w, "invalid or expired authorization state", http.StatusForbidden)
return
}
http.Error(w, "authorization was cancelled", http.StatusBadRequest)
return
}
if code == "" {
http.Error(w, "authorization was not completed", http.StatusBadRequest)
return
}
var id, storedProvider, status string
var verifierCipher, nonceCipher []byte
var expires int64
err := s.store.db.QueryRow(`SELECT id,provider,status,verifier_cipher,nonce_cipher,expires_at FROM devices WHERE state_hash=?`, s.store.digest("oauth-state", state)).Scan(&id, &storedProvider, &status, &verifierCipher, &nonceCipher, &expires)
if err != nil || provider != storedProvider || status != "authorizing" || s.store.now().Unix() >= expires {
http.Error(w, "invalid or expired authorization state", http.StatusForbidden)
return
}
verifier, err := s.store.open(id, "pkce", verifierCipher)
if err != nil {
http.Error(w, "authorization state unavailable", http.StatusInternalServerError)
return
}
nonce, err := s.store.open(id, "nonce", nonceCipher)
if err != nil {
http.Error(w, "authorization state unavailable", http.StatusInternalServerError)
return
}
identity, err := s.exchangeIdentity(r.Context(), provider, code, string(verifier), string(nonce))
clear(verifier)
clear(nonce)
if err != nil {
_, _ = s.store.db.Exec(`UPDATE devices SET status='failed',error_code='provider_rejected',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL WHERE id=? AND status='authorizing'`, id)
var failure *providerFailure
if errors.As(err, &failure) {
slog.Warn("OAuth provider authorization failed",
"provider", failure.Provider,
"stage", failure.Stage,
"reason", failure.Reason,
"http_status", failure.HTTPStatus,
"oauth_error", failure.OAuthError)
if failure.OAuthError == "invalid_client" {
http.Error(w, "server OAuth configuration is invalid; contact the server administrator", http.StatusBadGateway)
return
}
} else {
slog.Warn("OAuth provider authorization failed", "provider", provider, "reason", "internal_error")
}
http.Error(w, "provider authorization failed", http.StatusBadGateway)
return
}
if err := s.completeDevice(id, provider, identity); err != nil {
code := http.StatusInternalServerError
if errors.Is(err, ErrNotAllowed) {
code = http.StatusForbidden
} else if errors.Is(err, ErrConsumed) {
code = http.StatusConflict
}
http.Error(w, "authorization could not be completed", code)
return
}
w.Header().Set("Content-Security-Policy", "default-src 'none'; style-src 'unsafe-inline'")
w.Header().Set("Content-Type", "text/html; charset=utf-8")
_, _ = io.WriteString(w, `<!doctype html><meta charset="utf-8"><title>BD2 login</title><p>Login complete. You can return to the game.</p>`)
}
type providerIdentity struct{ issuer, subject string }
type providerFailure struct {
Provider string
Stage string
Reason string
HTTPStatus int
OAuthError string
}
func (e *providerFailure) Error() string {
return fmt.Sprintf("provider=%s stage=%s reason=%s status=%d oauth_error=%s", e.Provider, e.Stage, e.Reason, e.HTTPStatus, e.OAuthError)
}
func networkProviderFailure(ctx context.Context, provider, stage string, err error) error {
reason := "network_error"
if errors.Is(ctx.Err(), context.DeadlineExceeded) || errors.Is(err, context.DeadlineExceeded) {
reason = "timeout"
} else if errors.Is(ctx.Err(), context.Canceled) || errors.Is(err, context.Canceled) {
reason = "cancelled"
} else {
var networkError net.Error
if errors.As(err, &networkError) && networkError.Timeout() {
reason = "timeout"
}
}
return &providerFailure{Provider: provider, Stage: stage, Reason: reason}
}
func rejectedProviderFailure(provider, stage string, response *http.Response) error {
failure := &providerFailure{Provider: provider, Stage: stage, Reason: "http_rejected", HTTPStatus: response.StatusCode}
var body struct {
Error string `json:"error"`
}
decoder := json.NewDecoder(io.LimitReader(response.Body, 8<<10))
if decoder.Decode(&body) == nil {
failure.OAuthError = safeOAuthError(body.Error)
}
return failure
}
func invalidProviderResponse(provider, stage string) error {
return &providerFailure{Provider: provider, Stage: stage, Reason: "invalid_response", HTTPStatus: http.StatusOK}
}
func safeOAuthError(value string) string {
switch value {
case "invalid_request", "invalid_client", "invalid_grant", "unauthorized_client",
"unsupported_grant_type", "invalid_scope", "access_denied", "server_error", "temporarily_unavailable":
return value
default:
return "unknown"
}
}
func (s *Service) exchangeIdentity(ctx context.Context, provider, code, verifier, nonce string) (providerIdentity, error) {
values := url.Values{"client_id": {s.config.Providers[provider].ClientID}, "client_secret": {s.config.ProviderSecrets[provider]}, "grant_type": {"authorization_code"}, "code": {code}, "redirect_uri": {s.redirectURL(provider)}, "code_verifier": {verifier}}
request, _ := http.NewRequestWithContext(ctx, http.MethodPost, providerTokenURL(provider), strings.NewReader(values.Encode()))
request.Header.Set("Content-Type", "application/x-www-form-urlencoded")
response, err := s.client.Do(request)
if err != nil {
return providerIdentity{}, networkProviderFailure(ctx, provider, "token_exchange", err)
}
defer func() { _ = response.Body.Close() }()
if response.StatusCode != http.StatusOK {
return providerIdentity{}, rejectedProviderFailure(provider, "token_exchange", response)
}
var token struct {
AccessToken string `json:"access_token"`
IDToken string `json:"id_token"`
}
if err := decodeProviderJSON(response.Body, &token); err != nil || token.AccessToken == "" {
return providerIdentity{}, invalidProviderResponse(provider, "token_exchange")
}
if provider == "google" {
if token.IDToken == "" {
return providerIdentity{}, invalidProviderResponse(provider, "token_exchange")
}
identity, err := s.verifyGoogleIDToken(ctx, token.IDToken, nonce)
token.AccessToken, token.IDToken = "", ""
return identity, err
}
userinfo, _ := http.NewRequestWithContext(ctx, http.MethodGet, providerUserURL(provider), nil)
userinfo.Header.Set("Authorization", "Bearer "+token.AccessToken)
response, err = s.client.Do(userinfo)
token.AccessToken = ""
if err != nil {
return providerIdentity{}, networkProviderFailure(ctx, provider, "userinfo", err)
}
defer func() { _ = response.Body.Close() }()
if response.StatusCode != http.StatusOK {
return providerIdentity{}, rejectedProviderFailure(provider, "userinfo", response)
}
var user struct {
ID string `json:"id"`
Sub string `json:"sub"`
}
if err := decodeProviderJSON(response.Body, &user); err != nil {
return providerIdentity{}, invalidProviderResponse(provider, "userinfo")
}
if provider == "discord" && user.ID != "" {
return providerIdentity{issuer: "https://discord.com", subject: user.ID}, nil
}
return providerIdentity{}, invalidProviderResponse(provider, "userinfo")
}
// verifyGoogleIDToken delegates signature and standard-claim verification to
// Google's HTTPS tokeninfo endpoint, then independently verifies this server's
// audience, nonce and expiry. The raw ID token is never persisted or logged.
func (s *Service) verifyGoogleIDToken(ctx context.Context, idToken, nonce string) (providerIdentity, error) {
endpoint := "https://oauth2.googleapis.com/tokeninfo?id_token=" + url.QueryEscape(idToken)
request, _ := http.NewRequestWithContext(ctx, http.MethodGet, endpoint, nil)
response, err := s.client.Do(request)
if err != nil {
return providerIdentity{}, networkProviderFailure(ctx, "google", "id_token_verify", err)
}
defer func() { _ = response.Body.Close() }()
if response.StatusCode != http.StatusOK {
return providerIdentity{}, rejectedProviderFailure("google", "id_token_verify", response)
}
var claims struct {
Issuer string `json:"iss"`
Audience string `json:"aud"`
Subject string `json:"sub"`
Nonce string `json:"nonce"`
Expires string `json:"exp"`
}
if err := decodeProviderJSON(response.Body, &claims); err != nil {
return providerIdentity{}, invalidProviderResponse("google", "id_token_verify")
}
expires, err := strconv.ParseInt(claims.Expires, 10, 64)
validIssuer := claims.Issuer == "https://accounts.google.com" || claims.Issuer == "accounts.google.com"
if err != nil || !validIssuer || claims.Audience != s.config.Providers["google"].ClientID || claims.Subject == "" || claims.Nonce != nonce || s.store.now().Unix() >= expires {
return providerIdentity{}, &providerFailure{Provider: "google", Stage: "id_token_verify", Reason: "invalid_claims"}
}
return providerIdentity{issuer: "https://accounts.google.com", subject: claims.Subject}, nil
}
func decodeProviderJSON(reader io.Reader, target any) error {
data, err := io.ReadAll(io.LimitReader(reader, 1<<20+1))
if err != nil {
return err
}
if len(data) > 1<<20 {
return errors.New("auth: provider response is too large")
}
decoder := json.NewDecoder(bytes.NewReader(data))
if err := decoder.Decode(target); err != nil {
return err
}
var trailing any
if err := decoder.Decode(&trailing); !errors.Is(err, io.EOF) {
return errors.New("auth: provider response has trailing JSON")
}
return nil
}
func (s *Service) completeDevice(deviceID, provider string, identity providerIdentity) error {
if err := validateProviderIdentity(provider, identity); err != nil {
return err
}
now := s.store.now()
tx, err := s.store.db.Begin()
if err != nil {
return err
}
defer func() { _ = tx.Rollback() }()
var accountID, status string
subjectHash := s.store.identityDigest(identity.issuer, identity.subject)
err = tx.QueryRow(`SELECT i.account_id,a.status FROM identities i JOIN accounts a ON a.id=i.account_id WHERE i.issuer=? AND i.subject_hash=?`, identity.issuer, subjectHash).Scan(&accountID, &status)
if errors.Is(err, sql.ErrNoRows) {
var count int
if err := tx.QueryRow(`SELECT COUNT(*) FROM accounts`).Scan(&count); err != nil {
return err
}
if count != 0 {
result, err := tx.Exec(`UPDATE devices SET status='failed',error_code='not_allowed',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL WHERE id=? AND provider=? AND status='authorizing'`, deviceID, provider)
if err != nil {
return err
}
if count, err := rowsAffected(result); err != nil {
return err
} else if count != 1 {
return ErrConsumed
}
if err := tx.Commit(); err != nil {
return err
}
return ErrNotAllowed
}
accountID, err = randomToken(18)
if err != nil {
return err
}
if _, err = tx.Exec(`INSERT INTO accounts(id,status,created_at,last_login_at) VALUES(?,'active',?,?)`, accountID, now.Unix(), now.Unix()); err != nil {
return err
}
if _, err = tx.Exec(`INSERT INTO identities(provider,issuer,subject_hash,account_id,created_at,last_login_at) VALUES(?,?,?,?,?,?)`, provider, identity.issuer, subjectHash, accountID, now.Unix(), now.Unix()); err != nil {
return err
}
status = "active"
} else if err != nil {
return err
} else {
if _, err = tx.Exec(`UPDATE identities SET last_login_at=? WHERE issuer=? AND subject_hash=?`, now.Unix(), identity.issuer, subjectHash); err != nil {
return err
}
}
if status != "active" {
result, err := tx.Exec(`UPDATE devices SET status='failed',error_code='not_allowed',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL WHERE id=? AND provider=? AND status='authorizing'`, deviceID, provider)
if err != nil {
return err
}
if count, err := rowsAffected(result); err != nil {
return err
} else if count != 1 {
return ErrConsumed
}
if err := tx.Commit(); err != nil {
return err
}
return ErrNotAllowed
}
result, familyID, err := s.issueTokens(tx, accountID, provider, now)
if err != nil {
return err
}
_ = familyID
payload, err := json.Marshal(result)
if err != nil {
return err
}
sealed, err := s.store.seal(deviceID, "result", payload)
clear(payload)
if err != nil {
return err
}
update, err := tx.Exec(`UPDATE devices SET result_cipher=?,status='complete',state_hash=NULL,verifier_cipher=NULL,nonce_cipher=NULL WHERE id=? AND provider=? AND status='authorizing'`, sealed, deviceID, provider)
if err != nil {
return err
}
if count, err := rowsAffected(update); err != nil {
return err
} else if count != 1 {
return ErrConsumed
}
return tx.Commit()
}
func validateProviderIdentity(provider string, identity providerIdentity) error {
switch provider {
case "discord":
if identity.issuer != "https://discord.com" || len(identity.subject) == 0 || len(identity.subject) > 32 {
return errors.New("auth: invalid Discord identity")
}
for _, digit := range identity.subject {
if digit < '0' || digit > '9' {
return errors.New("auth: invalid Discord identity")
}
}
case "google":
if identity.issuer != "https://accounts.google.com" || len(identity.subject) == 0 || len(identity.subject) > 255 {
return errors.New("auth: invalid Google identity")
}
default:
return errors.New("auth: unsupported identity provider")
}
return nil
}
func (s *Service) issueTokens(tx *sql.Tx, accountID, provider string, now time.Time) (deviceResult, string, error) {
familyID, err := randomToken(18)
if err != nil {
return deviceResult{}, "", err
}
access, err := randomToken(32)
if err != nil {
return deviceResult{}, "", err
}
refresh, err := randomToken(32)
if err != nil {
return deviceResult{}, "", err
}
if _, err := tx.Exec(`INSERT INTO families(id,account_id,provider,created_at,expires_at) VALUES(?,?,?,?,?)`, familyID, accountID, provider, now.Unix(), now.Add(s.config.RefreshTTL).Unix()); err != nil {
return deviceResult{}, "", err
}
if _, err := tx.Exec(`INSERT INTO access_tokens(token_hash,family_id,account_id,created_at,expires_at) VALUES(?,?,?,?,?)`, s.store.digest("access-token", access), familyID, accountID, now.Unix(), now.Add(s.config.AccessTTL).Unix()); err != nil {
return deviceResult{}, "", err
}
if _, err := tx.Exec(`INSERT INTO refresh_tokens(token_hash,family_id,created_at,expires_at) VALUES(?,?,?,?)`, s.store.digest("refresh-token", refresh), familyID, now.Unix(), now.Add(s.config.RefreshTTL).Unix()); err != nil {
return deviceResult{}, "", err
}
return deviceResult{Provider: provider, AccessToken: access, AccessExpiresIn: int64(s.config.AccessTTL.Seconds()), RefreshToken: refresh, RefreshExpiresIn: int64(s.config.RefreshTTL.Seconds())}, familyID, nil
}
func (s *Service) poll(w http.ResponseWriter, r *http.Request) {
id := r.PathValue("id")
if !s.limits.allow("poll:"+remoteIP(r.RemoteAddr)+":"+id, s.store.now(), time.Minute, 60) {
w.Header().Set("Retry-After", "2")
http.Error(w, "poll rate exceeded", http.StatusTooManyRequests)
return
}
authorization := r.Header.Get("Authorization")
if !strings.HasPrefix(authorization, "Device ") {
http.Error(w, "invalid device transaction", http.StatusForbidden)
return
}
secret := strings.TrimPrefix(authorization, "Device ")
tx, err := s.store.db.Begin()
if err != nil {
http.Error(w, "login result unavailable", http.StatusInternalServerError)
return
}
defer func() { _ = tx.Rollback() }()
var storedHash, sealed []byte
var status, errorCode string
var expires int64
err = tx.QueryRow(`SELECT secret_hash,status,COALESCE(result_cipher,X''),COALESCE(error_code,''),expires_at FROM devices WHERE id=?`, id).Scan(&storedHash, &status, &sealed, &errorCode, &expires)
if err != nil || subtle.ConstantTimeCompare(storedHash, s.store.digest("device-secret", secret)) != 1 {
http.Error(w, "invalid device transaction", http.StatusForbidden)
return
}
if s.store.now().Unix() >= expires {
http.Error(w, "device transaction expired", http.StatusGone)
return
}
switch status {
case "created", "authorizing":
_ = tx.Rollback()
writeJSON(w, http.StatusAccepted, map[string]any{"status": "pending", "retry_after": 2})
case "failed":
_ = tx.Rollback()
writeJSON(w, http.StatusForbidden, map[string]string{"status": "failed", "error": errorCode})
case "complete":
plain, err := s.store.open(id, "result", sealed)
if err != nil {
http.Error(w, "login result unavailable", http.StatusInternalServerError)
return
}
result, err := tx.Exec(`UPDATE devices SET result_cipher=NULL,status='consumed' WHERE id=? AND status='complete'`, id)
count, affectedErr := rowsAffected(result)
if err != nil || affectedErr != nil || count != 1 {
_ = tx.Rollback()
clear(plain)
http.Error(w, "login result already consumed", http.StatusGone)
return
}
if err := tx.Commit(); err != nil {
clear(plain)
http.Error(w, "login result unavailable", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/json; charset=utf-8")
_, _ = w.Write(plain)
clear(plain)
default:
_ = tx.Rollback()
http.Error(w, "device transaction consumed", http.StatusGone)
}
}
func (l *requestLimiter) allow(key string, now time.Time, duration time.Duration, maximum int) bool {
l.mu.Lock()
defer l.mu.Unlock()
if l.windows == nil {
l.windows = make(map[string]limitWindow)
}
if l.lastSweep.IsZero() || now.Sub(l.lastSweep) >= time.Minute {
for candidate, window := range l.windows {
if now.Sub(window.started) >= duration {
delete(l.windows, candidate)
}
}
l.lastSweep = now
}
window, exists := l.windows[key]
if !exists || now.Sub(window.started) >= duration {
if !exists && len(l.windows) >= 4096 {
return false
}
l.windows[key] = limitWindow{started: now, count: 1}
return true
}
if window.count >= maximum {
return false
}
window.count++
l.windows[key] = window
return true
}
func remoteIP(remoteAddr string) string {
host, _, err := net.SplitHostPort(remoteAddr)
if err == nil && host != "" {
return host
}
return remoteAddr
}
func cleanupExpired(tx *sql.Tx, now int64) error {
statements := []struct {
query string
args []any
}{
{`DELETE FROM devices WHERE expires_at<=?`, []any{now}},
{`DELETE FROM access_tokens WHERE expires_at<=? OR family_id IN (SELECT id FROM families WHERE expires_at<=?)`, []any{now, now}},
{`DELETE FROM refresh_attempts WHERE expires_at<=? OR family_id IN (SELECT id FROM families WHERE expires_at<=?)`, []any{now, now}},
// Used refresh rows remain until their family expires so their reuse can
// still revoke every credential in that family.
{`DELETE FROM refresh_tokens WHERE family_id IN (SELECT id FROM families WHERE expires_at<=?)`, []any{now}},
{`DELETE FROM families WHERE expires_at<=?`, []any{now}},
}
for _, statement := range statements {
if _, err := tx.Exec(statement.query, statement.args...); err != nil {
return err
}
}
return nil
}
func (s *Service) refresh(w http.ResponseWriter, r *http.Request) {
var request struct {
RefreshToken string `json:"refresh_token"`
AttemptID string `json:"attempt_id"`
}
if !decodeJSON(w, r, &request) {
return
}
if request.RefreshToken == "" || !validRefreshAttemptID(request.AttemptID) {
http.Error(w, "refresh_token and valid attempt_id required", http.StatusBadRequest)
return
}
now := s.store.now()
requestTokenHash := s.store.digest("refresh-token", request.RefreshToken)
attemptHash := s.store.digest("refresh-attempt", request.AttemptID)
tx, err := s.store.db.Begin()
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
defer func() { _ = tx.Rollback() }()
if err := cleanupExpired(tx, now.Unix()); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
var familyID, accountID, provider, accountStatus string
var tokenExpires, familyExpires int64
var usedAt, revokedAt sql.NullInt64
err = tx.QueryRow(`SELECT r.family_id,f.account_id,f.provider,a.status,r.expires_at,f.expires_at,r.used_at,COALESCE(r.revoked_at,f.revoked_at) FROM refresh_tokens r JOIN families f ON f.id=r.family_id JOIN accounts a ON a.id=f.account_id WHERE r.token_hash=?`, requestTokenHash).Scan(&familyID, &accountID, &provider, &accountStatus, &tokenExpires, &familyExpires, &usedAt, &revokedAt)
if err != nil || revokedAt.Valid || accountStatus != "active" || now.Unix() >= familyExpires {
w.Header().Set("X-BD2-Refresh-Invalid", "1")
http.Error(w, "refresh token invalid", http.StatusUnauthorized)
return
}
var savedRequestHash, resultCipher []byte
err = tx.QueryRow(`SELECT request_token_hash,result_cipher FROM refresh_attempts WHERE family_id=? AND attempt_hash=?`, familyID, attemptHash).Scan(&savedRequestHash, &resultCipher)
if err == nil {
if subtle.ConstantTimeCompare(savedRequestHash, requestTokenHash) != 1 {
w.Header().Set("X-BD2-Refresh-Invalid", "1")
http.Error(w, "refresh attempt_id already belongs to another request", http.StatusConflict)
return
}
plain, openErr := s.store.open(refreshAttemptSealID(familyID, attemptHash), "result", resultCipher)
if openErr != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
var saved refreshAttemptResult
decodeErr := json.Unmarshal(plain, &saved)
clear(plain)
if decodeErr != nil || saved.Provider == "" || saved.AccessToken == "" || saved.RefreshToken == "" {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
writeJSON(w, http.StatusOK, saved.deviceResult(now.Unix()))
return
}
if !errors.Is(err, sql.ErrNoRows) {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if now.Unix() >= tokenExpires {
w.Header().Set("X-BD2-Refresh-Invalid", "1")
http.Error(w, "refresh token expired", http.StatusUnauthorized)
return
}
if usedAt.Valid {
if _, err := tx.Exec(`UPDATE families SET revoked_at=? WHERE id=? AND revoked_at IS NULL`, now.Unix(), familyID); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if err := tx.Commit(); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
w.Header().Set("X-BD2-Refresh-Invalid", "1")
http.Error(w, "refresh token replayed", http.StatusUnauthorized)
return
}
newAccess, err := randomToken(32)
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
newRefresh, err := randomToken(32)
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
updated, err := tx.Exec(`UPDATE refresh_tokens SET used_at=? WHERE token_hash=? AND used_at IS NULL AND revoked_at IS NULL`, now.Unix(), requestTokenHash)
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if count, err := rowsAffected(updated); err != nil || count != 1 {
w.Header().Set("X-BD2-Refresh-Invalid", "1")
http.Error(w, "refresh token invalid", http.StatusUnauthorized)
return
}
if _, err = tx.Exec(`DELETE FROM access_tokens WHERE family_id=?`, familyID); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if _, err = tx.Exec(`INSERT INTO access_tokens(token_hash,family_id,account_id,created_at,expires_at) VALUES(?,?,?,?,?)`, s.store.digest("access-token", newAccess), familyID, accountID, now.Unix(), now.Add(s.config.AccessTTL).Unix()); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
refreshExpiry := min(familyExpires, now.Add(s.config.RefreshTTL).Unix())
if _, err = tx.Exec(`INSERT INTO refresh_tokens(token_hash,family_id,created_at,expires_at) VALUES(?,?,?,?)`, s.store.digest("refresh-token", newRefresh), familyID, now.Unix(), refreshExpiry); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
result := refreshAttemptResult{
Provider: provider, AccessToken: newAccess, AccessExpiresAt: now.Add(s.config.AccessTTL).Unix(),
RefreshToken: newRefresh, RefreshExpiresAt: refreshExpiry,
}
plain, err := json.Marshal(result)
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
resultCipher, err = s.store.seal(refreshAttemptSealID(familyID, attemptHash), "result", plain)
clear(plain)
if err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if _, err = tx.Exec(`INSERT INTO refresh_attempts(family_id,attempt_hash,request_token_hash,result_cipher,created_at,expires_at) VALUES(?,?,?,?,?,?)`, familyID, attemptHash, requestTokenHash, resultCipher, now.Unix(), refreshExpiry); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
if err = tx.Commit(); err != nil {
http.Error(w, "refresh unavailable", http.StatusInternalServerError)
return
}
writeJSON(w, http.StatusOK, result.deviceResult(now.Unix()))
}
func validRefreshAttemptID(value string) bool {
if len(value) < 16 || len(value) > 128 {
return false
}
for _, item := range value {
if item < 'a' || item > 'z' {
if item < 'A' || item > 'Z' {
if item < '0' || item > '9' {
if item != '-' && item != '_' {
return false
}
}
}
}
}
return true
}
func refreshAttemptSealID(familyID string, attemptHash []byte) string {
return familyID + ":" + base64.RawURLEncoding.EncodeToString(attemptHash)
}
func (r refreshAttemptResult) deviceResult(now int64) deviceResult {
accessTTL := max(r.AccessExpiresAt-now, 0)
refreshTTL := max(r.RefreshExpiresAt-now, 0)
return deviceResult{
Provider: r.Provider, AccessToken: r.AccessToken, AccessExpiresIn: accessTTL,
RefreshToken: r.RefreshToken, RefreshExpiresIn: refreshTTL,
}
}
func (s *Service) revoke(w http.ResponseWriter, r *http.Request) {
authorization := r.Header.Get("Authorization")
if !strings.HasPrefix(authorization, "Bearer ") {
http.Error(w, "access token required", http.StatusUnauthorized)
return
}
token := strings.TrimPrefix(authorization, "Bearer ")
if token == "" {
http.Error(w, "access token required", http.StatusUnauthorized)
return
}
now := s.store.now().Unix()
tx, err := s.store.db.Begin()
if err != nil {
http.Error(w, "revocation unavailable", http.StatusInternalServerError)
return
}
defer func() { _ = tx.Rollback() }()
var familyID, accountStatus string
var expires int64
var revoked sql.NullInt64
err = tx.QueryRow(`SELECT t.family_id,a.status,t.expires_at,COALESCE(t.revoked_at,f.revoked_at)
FROM access_tokens t JOIN families f ON f.id=t.family_id JOIN accounts a ON a.id=t.account_id
WHERE t.token_hash=?`, s.store.digest("access-token", token)).Scan(&familyID, &accountStatus, &expires, &revoked)
if err != nil || accountStatus != "active" || revoked.Valid || now >= expires {
http.Error(w, "access token invalid", http.StatusUnauthorized)
return
}
result, err := tx.Exec(`UPDATE families SET revoked_at=? WHERE id=? AND revoked_at IS NULL`, now, familyID)
if err != nil {
http.Error(w, "revocation unavailable", http.StatusInternalServerError)
return
}
if count, err := rowsAffected(result); err != nil || count != 1 {
http.Error(w, "access token invalid", http.StatusUnauthorized)
return
}
if err := tx.Commit(); err != nil {
http.Error(w, "revocation unavailable", http.StatusInternalServerError)
return
}
w.WriteHeader(http.StatusNoContent)
}
func (s *Service) ValidateAccess(token string) (string, error) {
if token == "" {
return "", ErrUnauthorized
}
var accountID, status string
var expires int64
var revoked sql.NullInt64
err := s.store.db.QueryRow(`SELECT t.account_id,a.status,t.expires_at,COALESCE(t.revoked_at,f.revoked_at) FROM access_tokens t JOIN families f ON f.id=t.family_id JOIN accounts a ON a.id=t.account_id WHERE t.token_hash=?`, s.store.digest("access-token", token)).Scan(&accountID, &status, &expires, &revoked)
if err != nil || status != "active" || revoked.Valid || s.store.now().Unix() >= expires {
return "", ErrUnauthorized
}
return accountID, nil
}
// AuthenticateLogin validates LoginUserRequest.access_token (field 2) before
// the game session is established.
func (s *Service) AuthenticateLogin(request []byte) (string, error) {
token, found, err := wire.Bytes(request, 2)
if err != nil || !found {
return "", ErrUnauthorized
}
return s.ValidateAccess(string(token))
}
func (s *Service) redirectURL(provider string) string {
return s.config.PublicURL + "/auth/" + provider + "/callback"
}
func providerScope(provider string) string {
if provider == "discord" {
return "identify"
}
return "openid"
}
func providerAuthorizeURL(provider string) string {
if provider == "discord" {
return "https://discord.com/oauth2/authorize"
}
return "https://accounts.google.com/o/oauth2/v2/auth"
}
func providerTokenURL(provider string) string {
if provider == "discord" {
return "https://discord.com/api/v10/oauth2/token"
}
return "https://oauth2.googleapis.com/token"
}
func providerUserURL(provider string) string {
return "https://discord.com/api/v10/users/@me"
}
func rowsAffected(result sql.Result) (int64, error) {
if result == nil {
return 0, errors.New("auth: missing SQL result")
}
value, err := result.RowsAffected()
if err != nil {
return 0, fmt.Errorf("auth: count affected rows: %w", err)
}
return value, nil
}