feat(all): add resilient server recovery and talent upgrades

This commit is contained in:
2026-10-03 15:09:31 +08:00
parent db1eab764e
commit ce03fe7289
70 changed files with 3664 additions and 419 deletions
+116 -5
View File
@@ -50,6 +50,14 @@ type deviceResult struct {
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")
@@ -731,6 +739,7 @@ func cleanupExpired(tx *sql.Tx, now int64) error {
}{
{`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}},
@@ -747,25 +756,69 @@ func cleanupExpired(tx *sql.Tx, now int64) error {
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) || request.RefreshToken == "" {
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 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=?`, s.store.digest("refresh-token", request.RefreshToken)).Scan(&familyID, &accountID, &provider, &accountStatus, &tokenExpires, &familyExpires, &usedAt, &revokedAt)
if err != nil || revokedAt.Valid || accountStatus != "active" || now.Unix() >= tokenExpires || now.Unix() >= familyExpires {
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)
@@ -775,6 +828,7 @@ func (s *Service) refresh(w http.ResponseWriter, r *http.Request) {
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
}
@@ -788,12 +842,13 @@ func (s *Service) refresh(w http.ResponseWriter, r *http.Request) {
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(), s.store.digest("refresh-token", request.RefreshToken))
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
}
@@ -810,11 +865,67 @@ func (s *Service) refresh(w http.ResponseWriter, r *http.Request) {
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, deviceResult{Provider: provider, AccessToken: newAccess, AccessExpiresIn: int64(s.config.AccessTTL.Seconds()), RefreshToken: newRefresh, RefreshExpiresIn: refreshExpiry - now.Unix()})
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 := r.AccessExpiresAt - now
if accessTTL < 0 {
accessTTL = 0
}
refreshTTL := r.RefreshExpiresAt - now
if refreshTTL < 0 {
refreshTTL = 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) {
+84 -4
View File
@@ -147,7 +147,7 @@ func TestRefreshRotationReplayRevokesFamily(t *testing.T) {
service, store := testService(t)
first := completeAndPoll(t, service, store, "refresh-device", "discord", "https://discord.com", "123456789")
handler := service.Handler()
response := postJSON(handler, "/auth/session/refresh", map[string]string{"refresh_token": first.RefreshToken})
response := postJSON(handler, "/auth/session/refresh", map[string]string{"refresh_token": first.RefreshToken, "attempt_id": "rotation-attempt-0001"})
if response.Code != http.StatusOK {
t.Fatalf("refresh status=%d body=%q", response.Code, response.Body.String())
}
@@ -165,14 +165,14 @@ func TestRefreshRotationReplayRevokesFamily(t *testing.T) {
t.Fatalf("new access token rejected: %v", err)
}
replay := postJSON(handler, "/auth/session/refresh", map[string]string{"refresh_token": first.RefreshToken})
replay := postJSON(handler, "/auth/session/refresh", map[string]string{"refresh_token": first.RefreshToken, "attempt_id": "replay-attempt-000002"})
if replay.Code != http.StatusUnauthorized {
t.Fatalf("replay status=%d body=%q", replay.Code, replay.Body.String())
}
if _, err := service.ValidateAccess(rotated.AccessToken); err == nil {
t.Fatal("refresh replay did not revoke the token family")
}
next := postJSON(handler, "/auth/session/refresh", map[string]string{"refresh_token": rotated.RefreshToken})
next := postJSON(handler, "/auth/session/refresh", map[string]string{"refresh_token": rotated.RefreshToken, "attempt_id": "after-replay-attempt-3"})
if next.Code != http.StatusUnauthorized {
t.Fatalf("family refresh after replay status=%d", next.Code)
}
@@ -191,12 +191,92 @@ func TestRevokeInvalidatesAccessAndRefreshFamily(t *testing.T) {
if _, err := service.ValidateAccess(tokens.AccessToken); err == nil {
t.Fatal("revoked access token remained valid")
}
refresh := postJSON(service.Handler(), "/auth/session/refresh", map[string]string{"refresh_token": tokens.RefreshToken})
refresh := postJSON(service.Handler(), "/auth/session/refresh", map[string]string{"refresh_token": tokens.RefreshToken, "attempt_id": "revoked-attempt-0001"})
if refresh.Code != http.StatusUnauthorized {
t.Fatalf("revoked refresh status=%d", refresh.Code)
}
}
func TestRefreshRetryWithSameAttemptReturnsCommittedRotation(t *testing.T) {
service, store := testService(t)
first := completeAndPoll(t, service, store, "refresh-retry-device", "discord", "https://discord.com", "123456789")
handler := service.Handler()
request := map[string]string{"refresh_token": first.RefreshToken, "attempt_id": "stable-attempt-000001"}
response := postJSON(handler, "/auth/session/refresh", request)
if response.Code != http.StatusOK {
t.Fatalf("first refresh status=%d body=%q", response.Code, response.Body.String())
}
var rotated deviceResult
if err := json.Unmarshal(response.Body.Bytes(), &rotated); err != nil {
t.Fatal(err)
}
retry := postJSON(handler, "/auth/session/refresh", request)
if retry.Code != http.StatusOK {
t.Fatalf("retry status=%d body=%q", retry.Code, retry.Body.String())
}
var replayed deviceResult
if err := json.Unmarshal(retry.Body.Bytes(), &replayed); err != nil {
t.Fatal(err)
}
if replayed.Provider != rotated.Provider || replayed.AccessToken != rotated.AccessToken || replayed.RefreshToken != rotated.RefreshToken {
t.Fatalf("retry returned a different rotation: first=%+v retry=%+v", rotated, replayed)
}
if _, err := service.ValidateAccess(rotated.AccessToken); err != nil {
t.Fatalf("idempotent retry revoked the family: %v", err)
}
var sealed []byte
if err := store.db.QueryRow(`SELECT result_cipher FROM refresh_attempts`).Scan(&sealed); err != nil {
t.Fatal(err)
}
if bytes.Contains(sealed, []byte(rotated.AccessToken)) || bytes.Contains(sealed, []byte(rotated.RefreshToken)) {
t.Fatal("refresh attempt result was stored outside AES-GCM ciphertext")
}
}
func TestRefreshRequiresStableAttemptID(t *testing.T) {
service, store := testService(t)
tokens := completeAndPoll(t, service, store, "refresh-attempt-required", "discord", "https://discord.com", "123456789")
response := postJSON(service.Handler(), "/auth/session/refresh", map[string]string{"refresh_token": tokens.RefreshToken})
if response.Code != http.StatusBadRequest {
t.Fatalf("missing attempt status=%d body=%q", response.Code, response.Body.String())
}
if _, err := service.ValidateAccess(tokens.AccessToken); err != nil {
t.Fatalf("malformed refresh request changed credential family: %v", err)
}
}
func TestRefreshCommittedAttemptReplaysAfterRequestTokenExpiry(t *testing.T) {
service, store := testService(t)
first := completeAndPoll(t, service, store, "refresh-expiry-device", "discord", "https://discord.com", "123456789")
request := map[string]string{"refresh_token": first.RefreshToken, "attempt_id": "expiry-replay-attempt-01"}
response := postJSON(service.Handler(), "/auth/session/refresh", request)
if response.Code != http.StatusOK {
t.Fatalf("first refresh status=%d body=%q", response.Code, response.Body.String())
}
var rotated deviceResult
if err := json.Unmarshal(response.Body.Bytes(), &rotated); err != nil {
t.Fatal(err)
}
if _, err := store.db.Exec(`UPDATE refresh_tokens SET expires_at=? WHERE token_hash=?`, testNowUnix-1, store.digest("refresh-token", first.RefreshToken)); err != nil {
t.Fatal(err)
}
store.now = func() time.Time { return time.Unix(testNowUnix+int64((16*time.Minute).Seconds()), 0) }
replay := postJSON(service.Handler(), "/auth/session/refresh", request)
if replay.Code != http.StatusOK {
t.Fatalf("expired request token replay status=%d body=%q", replay.Code, replay.Body.String())
}
var replayed deviceResult
if err := json.Unmarshal(replay.Body.Bytes(), &replayed); err != nil {
t.Fatal(err)
}
if replayed.AccessToken != rotated.AccessToken || replayed.RefreshToken != rotated.RefreshToken {
t.Fatalf("replay changed committed rotation: first=%+v replay=%+v", rotated, replayed)
}
if replayed.AccessExpiresIn != 0 {
t.Fatalf("expired replay access TTL=%d, want 0", replayed.AccessExpiresIn)
}
}
func TestSensitiveAuthenticationMaterialIsNotStoredInPlaintext(t *testing.T) {
service, store := testService(t)
handler := service.Handler()
+26 -4
View File
@@ -18,7 +18,7 @@ import (
_ "modernc.org/sqlite"
)
const schemaVersion = 1
const schemaVersion = 2
var (
ErrUnauthorized = errors.New("auth: unauthorized")
@@ -84,16 +84,38 @@ func Open(path string, masterKey []byte) (*Store, error) {
var version int
err = tx.QueryRow(`SELECT CAST(value AS INTEGER) FROM metadata WHERE key='schema_version'`).Scan(&version)
if errors.Is(err, sql.ErrNoRows) {
if _, err = tx.Exec(`INSERT INTO metadata(key,value) VALUES('schema_version',?)`, schemaVersion); err != nil {
if _, err = tx.Exec(`INSERT INTO metadata(key,value) VALUES('schema_version',1)`); err != nil {
return nil, err
}
version = schemaVersion
version = 1
} else if err != nil {
return nil, err
}
if version != schemaVersion {
if version < 1 || version > schemaVersion {
return nil, fmt.Errorf("auth: schema version %d, want %d", version, schemaVersion)
}
for version < schemaVersion {
switch version {
case 1:
if _, err = tx.Exec(`CREATE TABLE refresh_attempts (
family_id TEXT NOT NULL REFERENCES families(id),
attempt_hash BLOB NOT NULL,
request_token_hash BLOB NOT NULL,
result_cipher BLOB NOT NULL,
created_at INTEGER NOT NULL,
expires_at INTEGER NOT NULL,
PRIMARY KEY(family_id,attempt_hash)
) WITHOUT ROWID`); err != nil {
return nil, fmt.Errorf("auth: migrate schema 1->2: %w", err)
}
version = 2
default:
return nil, fmt.Errorf("auth: missing adjacent migration %d->%d", version, version+1)
}
if _, err = tx.Exec(`UPDATE metadata SET value=? WHERE key='schema_version'`, version); err != nil {
return nil, fmt.Errorf("auth: record schema version %d: %w", version, err)
}
}
if err = tx.Commit(); err != nil {
return nil, err
}
+39
View File
@@ -2,6 +2,7 @@ package auth
import (
"bytes"
"database/sql"
"path/filepath"
"testing"
)
@@ -23,6 +24,44 @@ func TestStoreRequiresAndClearsExactMasterKey(t *testing.T) {
}
}
func TestStoreMigratesSchemaV1ToV2(t *testing.T) {
path := filepath.Join(t.TempDir(), "auth.db")
store, err := Open(path, bytes.Repeat([]byte{0x61}, 32))
if err != nil {
t.Fatal(err)
}
if err := store.Close(); err != nil {
t.Fatal(err)
}
db, err := sql.Open("sqlite", path)
if err != nil {
t.Fatal(err)
}
if _, err := db.Exec(`DROP TABLE refresh_attempts; UPDATE metadata SET value='1' WHERE key='schema_version'`); err != nil {
db.Close()
t.Fatal(err)
}
if err := db.Close(); err != nil {
t.Fatal(err)
}
reopened, err := Open(path, bytes.Repeat([]byte{0x61}, 32))
if err != nil {
t.Fatal(err)
}
defer reopened.Close()
var version int
if err := reopened.db.QueryRow(`SELECT CAST(value AS INTEGER) FROM metadata WHERE key='schema_version'`).Scan(&version); err != nil {
t.Fatal(err)
}
if version != 2 {
t.Fatalf("schema_version=%d, want 2", version)
}
var table string
if err := reopened.db.QueryRow(`SELECT name FROM sqlite_master WHERE type='table' AND name='refresh_attempts'`).Scan(&table); err != nil {
t.Fatal(err)
}
}
func TestStorePersistsExplicitSchemaVersionAndRejectsUnknownVersion(t *testing.T) {
path := filepath.Join(t.TempDir(), "auth.db")
store, err := Open(path, bytes.Repeat([]byte{0x35}, 32))