滑动续期、网络错误归类、API Key 激活乐观锁与生效提示
This commit is contained in:
@@ -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,与存量账号的零值版本兼容(升级不强制登出)。
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user