滑动续期、网络错误归类、API Key 激活乐观锁与生效提示

This commit is contained in:
2026-08-11 11:45:03 +08:00
parent 23b2820101
commit e63af49831
21 changed files with 679 additions and 15 deletions
+6 -2
View File
@@ -25,8 +25,12 @@ var ErrInvalidCredentials = errors.New("invalid username or password")
// ErrLoginLocked 表示该 IP+用户名组合因连续失败被锁定;不提示剩余次数与时长细节。
var ErrLoginLocked = errors.New("too many failed attempts, try again later")
// tokenTTL 是登录令牌有效期。
const tokenTTL = 24 * time.Hour
// tokenTTL 是登录令牌有效期;renewThreshold 是滑动续期阈值——剩余有效期
// 低于该值的令牌在鉴权响应中自动换发同会话新令牌(见 MaybeRenew)。
const (
tokenTTL = 24 * time.Hour
renewThreshold = tokenTTL / 2
)
// authClaims 在标准声明外携带令牌版本;版本落后于账号当前值即失效。
// 存量令牌无 ver 字段解析为 0,与存量账号的零值版本兼容(升级不强制登出)。
+3
View File
@@ -190,6 +190,9 @@ func (s *OciConfigService) applyCredentialUpdate(cfg *model.OciConfig, in Update
return fmt.Errorf("encrypt private key: %w", err)
}
cfg.PrivateKeyEnc = enc
// 手工替换私钥同样进入 OCI 公钥传播窗口,记录激活时刻供前端提示
now := time.Now()
cfg.KeyActivatedAt = &now
}
if in.Passphrase == nil {
return nil
+3
View File
@@ -84,6 +84,9 @@ func validateProxyInput(in ProxyInput) error {
if strings.TrimSpace(in.Host) == "" {
return fmt.Errorf("主机不能为空: %w", ErrProxyInvalid)
}
if strings.Contains(in.Host, "@") {
return fmt.Errorf("主机不可包含用户凭据: %w", ErrProxyInvalid)
}
return nil
}
+1
View File
@@ -101,6 +101,7 @@ func TestProxyValidateAndDeleteInUse(t *testing.T) {
{Name: "x", Type: "ss", Host: "h", Port: 1080},
{Name: "x", Type: "http", Host: "h", Port: 0},
{Name: "x", Type: "http", Host: " ", Port: 8080},
{Name: "x", Type: "http", Host: "user:pass@proxy.example.com", Port: 8080},
} {
if _, err := svc.Create(ctx, in); err == nil {
t.Fatalf("Create(%+v) accepted invalid input", in)
+39
View File
@@ -431,3 +431,42 @@ func (s *AuthService) cleanupSessionsOnce(ctx context.Context) {
log.Printf("session cleanup: %v", err)
}
}
// MaybeRenew 滑动续期:对剩余有效期不足 renewThreshold 的令牌换发同会话
// (同 jti / 同版本)新令牌,并把会话行有效期延长到新过期点,不新建会话行。
// 仅供鉴权通过后的请求调用(令牌有效性已由 ParseTokenProof 保证);
// 无需换发或换发失败返回 ok=false,调用方跳过即可。
func (s *AuthService) MaybeRenew(ctx context.Context, tokenString string) (string, time.Time, bool) {
claims := &authClaims{}
if _, err := jwt.ParseWithClaims(tokenString, claims, func(*jwt.Token) (any, error) {
return s.jwtSecret, nil
}); err != nil || claims.ExpiresAt == nil {
return "", time.Time{}, false
}
if time.Until(claims.ExpiresAt.Time) >= renewThreshold {
return "", time.Time{}, false
}
token, expires, _, err := s.signTokenWithJTI(claims.Subject, claims.Ver, claims.ID)
if err != nil {
return "", time.Time{}, false
}
if err := s.extendSessionExpiry(ctx, claims.ID, expires); err != nil {
log.Printf("[WARN] %v", err)
return "", time.Time{}, false
}
return token, expires, true
}
// extendSessionExpiry 把会话行有效期延长到新过期点,保证「活跃会话」展示与
// 清理任务看到真实过期时间;无行(存量令牌)静默跳过。
func (s *AuthService) extendSessionExpiry(ctx context.Context, jti string, expires time.Time) error {
if jti == "" {
return nil
}
err := s.db.WithContext(ctx).Model(&model.UserSession{}).
Where("token_id = ?", jti).UpdateColumn("expires_at", expires).Error
if err != nil {
return fmt.Errorf("extend session expiry: %w", err)
}
return nil
}
+122
View File
@@ -7,6 +7,7 @@ import (
"testing"
"time"
"github.com/golang-jwt/jwt/v5"
"gorm.io/gorm"
"oci-portal/internal/model"
@@ -523,3 +524,124 @@ func TestSessionCleanup(t *testing.T) {
t.Errorf("after cleanup rows = %+v, want only alive", rows)
}
}
// signShortToken 用服务同款密钥手工签指定 TTL 的令牌,构造临近过期态。
func signShortToken(t *testing.T, auth *AuthService, username string, ver uint, jti string, ttl time.Duration) string {
t.Helper()
claims := &authClaims{
RegisteredClaims: jwt.RegisteredClaims{
Subject: username,
ID: jti,
IssuedAt: jwt.NewNumericDate(time.Now()),
ExpiresAt: jwt.NewNumericDate(time.Now().Add(ttl)),
},
Ver: ver,
}
tok, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(auth.jwtSecret)
if err != nil {
t.Fatalf("sign short token: %v", err)
}
return tok
}
func renewalFixture(t *testing.T) (*AuthService, string, TokenProof) {
t.Helper()
auth := newTestAuth(t)
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
t.Fatalf("EnsureAdmin: %v", err)
}
ctx := context.Background()
loginTok, _, err := auth.Login(ctx, "admin", "pass123", "",
SessionMeta{ClientIP: "127.0.0.1", UserAgent: "t"})
if err != nil {
t.Fatalf("Login: %v", err)
}
_, proof, err := auth.ParseTokenProof(ctx, loginTok)
if err != nil {
t.Fatalf("ParseTokenProof: %v", err)
}
return auth, loginTok, proof
}
func setSessionExpiry(t *testing.T, auth *AuthService, jti string, expires time.Time) model.UserSession {
t.Helper()
if err := auth.db.Model(&model.UserSession{}).Where("token_id = ?", jti).
UpdateColumn("expires_at", expires).Error; err != nil {
t.Fatalf("set session expiry: %v", err)
}
var row model.UserSession
if err := auth.db.Where("token_id = ?", jti).First(&row).Error; err != nil {
t.Fatalf("find session row: %v", err)
}
return row
}
func TestMaybeRenewEligibility(t *testing.T) {
auth, loginTok, _ := renewalFixture(t)
tests := []struct {
name string
token string
}{
{name: "剩余时间高于阈值", token: loginTok},
{name: "非法令牌", token: "not.a.token"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if _, _, ok := auth.MaybeRenew(context.Background(), tt.token); ok {
t.Error("MaybeRenew ok = true, want false")
}
})
}
}
func TestMaybeRenewExtendsSession(t *testing.T) {
auth, _, proof := renewalFixture(t)
before := setSessionExpiry(t, auth, proof.Jti, time.Now().Add(time.Hour))
short := signShortToken(t, auth, "admin", proof.Ver, proof.Jti, time.Hour)
renewed, expires, ok := auth.MaybeRenew(context.Background(), short)
if !ok {
t.Fatal("剩余 1h 的令牌应换发")
}
_, newProof, err := auth.ParseTokenProof(context.Background(), renewed)
if err != nil {
t.Fatalf("ParseTokenProof(renewed): %v", err)
}
if newProof.Jti != proof.Jti || newProof.Ver != proof.Ver {
t.Errorf("proof = %+v, want %+v", newProof, proof)
}
var after model.UserSession
if err := auth.db.Where("token_id = ?", proof.Jti).First(&after).Error; err != nil {
t.Fatalf("reload session row: %v", err)
}
if delta := after.ExpiresAt.Sub(before.ExpiresAt); delta < 22*time.Hour {
t.Errorf("session expiry delta = %v, want >= 22h", delta)
}
if gap := after.ExpiresAt.Sub(expires); gap < -time.Second || gap > time.Second {
t.Errorf("session expiry = %v, token expiry = %v", after.ExpiresAt, expires)
}
}
func TestMaybeRenewedTokenFollowsSessionRevocation(t *testing.T) {
auth, loginTok, proof := renewalFixture(t)
short := signShortToken(t, auth, "admin", proof.Ver, proof.Jti, time.Hour)
renewed, _, ok := auth.MaybeRenew(context.Background(), short)
if !ok {
t.Fatal("剩余 1h 的令牌应换发")
}
auth.Logout(context.Background(), loginTok)
if _, err := auth.ParseToken(context.Background(), renewed); err == nil {
t.Error("会话撤销后换发令牌仍有效, want 失效")
}
}
func TestMaybeRenewUpdateFailureDoesNotIssueToken(t *testing.T) {
auth, _, proof := renewalFixture(t)
short := signShortToken(t, auth, "admin", proof.Ver, proof.Jti, time.Hour)
if err := auth.db.Migrator().DropTable(&model.UserSession{}); err != nil {
t.Fatalf("drop sessions table: %v", err)
}
token, expires, ok := auth.MaybeRenew(context.Background(), short)
if ok || token != "" || !expires.IsZero() {
t.Errorf("MaybeRenew = (%q, %v, %v), want empty result", token, expires, ok)
}
}
+13 -4
View File
@@ -10,6 +10,8 @@ import (
"fmt"
"time"
"gorm.io/gorm"
"oci-portal/internal/model"
"oci-portal/internal/oci"
)
@@ -106,7 +108,7 @@ func (s *OciConfigService) ActivateApiKey(ctx context.Context, id uint, userID,
if err := s.waitApiKeyUsable(ctx, newCred); err != nil {
return err
}
return s.persistSigningKey(cfg, newCred)
return s.persistSigningKey(ctx, cfg, newCred)
}
// waitApiKeyUsable 用新凭据测活,等待上传的公钥在 OCI 侧生效。
@@ -126,7 +128,7 @@ func (s *OciConfigService) waitApiKeyUsable(ctx context.Context, cred oci.Creden
}
// persistSigningKey 加密新私钥,更新配置签名用户与指纹并清空口令密文(面板生成的 key 无口令)。
func (s *OciConfigService) persistSigningKey(cfg *model.OciConfig, newCred oci.Credentials) error {
func (s *OciConfigService) persistSigningKey(ctx context.Context, cfg *model.OciConfig, newCred oci.Credentials) error {
enc, err := s.cipher.EncryptString(newCred.PrivateKey)
if err != nil {
return fmt.Errorf("encrypt private key: %w", err)
@@ -134,9 +136,16 @@ func (s *OciConfigService) persistSigningKey(cfg *model.OciConfig, newCred oci.C
updates := map[string]any{
"user_oc_id": newCred.UserOCID, "fingerprint": newCred.Fingerprint,
"private_key_enc": enc, "passphrase_enc": "",
// 记录激活时刻:OCI 公钥全球传播为分钟级,前端据此做窗口期提示
"key_activated_at": time.Now(),
}
if err := s.db.Model(cfg).Updates(updates).Error; err != nil {
return fmt.Errorf("persist rotated key: %w", err)
res := s.db.WithContext(ctx).Model(&model.OciConfig{}).
Where("id = ? AND updated_at = ?", cfg.ID, cfg.UpdatedAt).Updates(updates)
if res.Error != nil {
return fmt.Errorf("persist rotated key: %w", res.Error)
}
if res.RowsAffected != 1 {
return fmt.Errorf("persist rotated key: %w", gorm.ErrRecordNotFound)
}
return nil
}
+90
View File
@@ -9,6 +9,8 @@ import (
"testing"
"time"
"gorm.io/gorm"
"oci-portal/internal/model"
"oci-portal/internal/oci"
)
@@ -194,6 +196,9 @@ func TestActivateApiKey(t *testing.T) {
if err != nil || plain != newKey {
t.Fatalf("persisted key mismatch (err=%v)", err)
}
if got.KeyActivatedAt == nil || time.Since(*got.KeyActivatedAt) > time.Minute {
t.Fatalf("keyActivatedAt = %v, want 刚写入的时间", got.KeyActivatedAt)
}
if len(fc.validated) == 0 || fc.validated[0] != "11:22" {
t.Fatalf("validated = %v", fc.validated)
}
@@ -203,3 +208,88 @@ func TestActivateApiKey(t *testing.T) {
})
}
}
func TestPersistSigningKeyRejectsStaleSnapshot(t *testing.T) {
tests := []struct {
name string
deleted bool
}{
{name: "租户已删除", deleted: true},
{name: "凭据已被并发更新"},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
testStaleSigningKeyPersistence(t, tt.deleted)
})
}
}
func TestPersistSigningKeyAdvancesSnapshotVersion(t *testing.T) {
s := newTestService(t, &apiKeyClient{})
cfg := seedApiKeyConfig(t, s)
stale := *cfg
first := oci.Credentials{UserOCID: cfg.UserOCID, Fingerprint: "11:22", PrivateKey: "first-key"}
if err := s.persistSigningKey(context.Background(), cfg, first); err != nil {
t.Fatalf("first persist: %v", err)
}
second := oci.Credentials{UserOCID: cfg.UserOCID, Fingerprint: "33:44", PrivateKey: "second-key"}
if err := s.persistSigningKey(context.Background(), &stale, second); !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("stale persist err = %v, want ErrRecordNotFound", err)
}
var got model.OciConfig
if err := s.db.First(&got, cfg.ID).Error; err != nil || got.Fingerprint != "11:22" {
t.Fatalf("persisted config = %+v, err %v", got, err)
}
}
func testStaleSigningKeyPersistence(t *testing.T, deleted bool) {
t.Helper()
s := newTestService(t, &apiKeyClient{})
cfg := seedApiKeyConfig(t, s)
invalidateSigningSnapshot(t, s, cfg, deleted)
cred := oci.Credentials{UserOCID: cfg.UserOCID, Fingerprint: "11:22", PrivateKey: "new-key"}
err := s.persistSigningKey(context.Background(), cfg, cred)
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("err = %v, want ErrRecordNotFound", err)
}
assertSigningSnapshotPreserved(t, s, cfg.ID, deleted)
}
func invalidateSigningSnapshot(t *testing.T, s *OciConfigService, cfg *model.OciConfig, deleted bool) {
t.Helper()
if deleted {
if err := s.db.Delete(&model.OciConfig{}, cfg.ID).Error; err != nil {
t.Fatalf("delete config: %v", err)
}
return
}
enc, err := s.cipher.EncryptString("concurrent-key")
if err != nil {
t.Fatalf("encrypt concurrent key: %v", err)
}
res := s.db.Model(&model.OciConfig{}).Where("id = ?", cfg.ID).UpdateColumns(map[string]any{
"fingerprint": "cc:dd", "private_key_enc": enc, "updated_at": cfg.UpdatedAt.Add(time.Second),
})
if res.Error != nil || res.RowsAffected != 1 {
t.Fatalf("mutate config = rows %d, err %v", res.RowsAffected, res.Error)
}
}
func assertSigningSnapshotPreserved(t *testing.T, s *OciConfigService, id uint, deleted bool) {
t.Helper()
var got model.OciConfig
err := s.db.First(&got, id).Error
if deleted {
if !errors.Is(err, gorm.ErrRecordNotFound) {
t.Fatalf("deleted config reload err = %v", err)
}
return
}
if err != nil || got.Fingerprint != "cc:dd" {
t.Fatalf("concurrent config = %+v, err %v", got, err)
}
plain, err := s.cipher.DecryptString(got.PrivateKeyEnc)
if err != nil || plain != "concurrent-key" {
t.Fatalf("concurrent key = %q, err %v", plain, err)
}
}