@@ -12,6 +12,7 @@ import (
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
"oci-portal/internal/service"
|
||||
)
|
||||
|
||||
// seedGatewayModel 直插启用渠道与模型缓存,绕过云同步。
|
||||
@@ -29,7 +30,7 @@ func seedGatewayModel(t *testing.T, db *gorm.DB, name string) {
|
||||
|
||||
func TestAiGatewayKeyModelRestrict(t *testing.T) {
|
||||
r, auth, _, db := newTestRouterDB(t)
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "", service.SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
@@ -125,7 +126,7 @@ func TestAiGatewayKeyModelRestrict(t *testing.T) {
|
||||
|
||||
func TestAiGatewayModelsFilteredByKey(t *testing.T) {
|
||||
r, auth, _, db := newTestRouterDB(t)
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "", service.SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
@@ -172,7 +173,7 @@ func TestAiGatewayModelsFilteredByKey(t *testing.T) {
|
||||
|
||||
func TestAiKeyModelsUpdateRoundTrip(t *testing.T) {
|
||||
r, auth, _, _ := newTestRouterDB(t)
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "", service.SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
|
||||
@@ -3,7 +3,6 @@ package api
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -34,7 +33,9 @@ type loginRequest struct {
|
||||
// @Param body body loginRequest true "登录凭据"
|
||||
// @Success 200 {object} tokenResponse "token 与 expiresAt"
|
||||
// @Failure 401 {object} errorResponse "凭据错误"
|
||||
// @Failure 403 {object} errorResponse "密码登录已禁用"
|
||||
// @Failure 428 {object} totpRequiredResponse "需要两步验证码(totpRequired=true)"
|
||||
// @Failure 429 {object} errorResponse "登录守卫锁定(无 code);全局 IP 限流时 code=RateLimited"
|
||||
// @Router /api/v1/auth/login [post]
|
||||
func (h *authHandler) login(c *gin.Context) {
|
||||
var req loginRequest
|
||||
@@ -43,7 +44,7 @@ func (h *authHandler) login(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
start := time.Now()
|
||||
token, expires, err := h.svc.Login(c.Request.Context(), req.Username, req.Password, requestIP(c), req.TotpCode)
|
||||
token, expires, err := h.svc.Login(c.Request.Context(), req.Username, req.Password, req.TotpCode, sessionMetaOf(c))
|
||||
if errors.Is(err, service.ErrTotpRequired) {
|
||||
// 密码已通过,引导前端弹出二次验证输入;不算失败不留痕
|
||||
c.JSON(http.StatusPreconditionRequired, gin.H{"error": err.Error(), "totpRequired": true})
|
||||
@@ -80,9 +81,8 @@ func (h *authHandler) login(c *gin.Context) {
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/logout [post]
|
||||
func (h *authHandler) logout(c *gin.Context) {
|
||||
token, ok := strings.CutPrefix(c.GetHeader("Authorization"), "Bearer ")
|
||||
if ok && token != "" {
|
||||
h.svc.Logout(token)
|
||||
if token := bearerToken(c); token != "" {
|
||||
h.svc.Logout(c.Request.Context(), token)
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
|
||||
+51
-27
@@ -16,9 +16,11 @@ import (
|
||||
|
||||
// authxHandler 处理两步验证与外部身份(OAuth)接口。
|
||||
type authxHandler struct {
|
||||
auth *service.AuthService
|
||||
oauth *service.OAuthService
|
||||
logs *service.SystemLogService
|
||||
auth *service.AuthService
|
||||
oauth *service.OAuthService
|
||||
passkeys *service.PasskeyService
|
||||
wallets *service.WalletService
|
||||
logs *service.SystemLogService
|
||||
}
|
||||
|
||||
// ---- TOTP(JWT 组内) ----
|
||||
@@ -66,7 +68,7 @@ func (h *authxHandler) totpSetup(c *gin.Context) {
|
||||
// @Tags 认证
|
||||
// @Param body body object true "{code: 6 位验证码}"
|
||||
// @Success 200 {object} tokenResponse "已启用,返回换发的新 token"
|
||||
// @Success 204 "已启用但新 token 签发失败,需重新登录"
|
||||
// @Failure 401 {object} errorResponse "请求处理期间会话已撤销"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/totp/activate [post]
|
||||
func (h *authxHandler) totpActivate(c *gin.Context) {
|
||||
@@ -77,7 +79,9 @@ func (h *authxHandler) totpActivate(c *gin.Context) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
err := h.auth.ActivateTotp(c.Request.Context(), c.GetString(usernameKey), req.Code)
|
||||
token, expires, err := h.auth.ActivateTotp(
|
||||
c.Request.Context(), c.GetString(usernameKey), req.Code,
|
||||
bearerToken(c), sessionMetaOf(c), tokenProofOf(c))
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrTotpInvalid) || errors.Is(err, service.ErrTotpNotSetup) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
@@ -86,7 +90,7 @@ func (h *authxHandler) totpActivate(c *gin.Context) {
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
h.respondFreshToken(c)
|
||||
c.JSON(http.StatusOK, gin.H{"token": token, "expiresAt": expires})
|
||||
}
|
||||
|
||||
// totpDisable 停用两步验证;需当前验证码或登录密码任一确认。
|
||||
@@ -95,7 +99,7 @@ func (h *authxHandler) totpActivate(c *gin.Context) {
|
||||
// @Tags 认证
|
||||
// @Param body body object true "{password 或 code 任一确认}"
|
||||
// @Success 200 {object} tokenResponse "已停用,返回换发的新 token"
|
||||
// @Success 204 "已停用但新 token 签发失败,需重新登录"
|
||||
// @Failure 401 {object} errorResponse "请求处理期间会话已撤销"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/totp/disable [post]
|
||||
func (h *authxHandler) totpDisable(c *gin.Context) {
|
||||
@@ -107,7 +111,9 @@ func (h *authxHandler) totpDisable(c *gin.Context) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
err := h.auth.DisableTotp(c.Request.Context(), c.GetString(usernameKey), req.Password, req.Code)
|
||||
token, expires, err := h.auth.DisableTotp(
|
||||
c.Request.Context(), c.GetString(usernameKey), req.Password, req.Code,
|
||||
bearerToken(c), sessionMetaOf(c), tokenProofOf(c))
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrTotpConfirm) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
@@ -116,13 +122,14 @@ func (h *authxHandler) totpDisable(c *gin.Context) {
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
h.respondFreshToken(c)
|
||||
c.JSON(http.StatusOK, gin.H{"token": token, "expiresAt": expires})
|
||||
}
|
||||
|
||||
// respondFreshToken 敏感变更后为操作者签发新令牌返回(版本已递增,
|
||||
// 旧令牌全部失效);签发失败降级 204,前端按会话失效走重新登录。
|
||||
func (h *authxHandler) respondFreshToken(c *gin.Context) {
|
||||
token, expires, err := h.auth.IssueToken(c.Request.Context(), c.GetString(usernameKey))
|
||||
// 旧令牌全部失效);旧令牌的会话行接续到新令牌,列表中当前设备保持连续。
|
||||
// 签发失败降级 204,前端按会话失效走重新登录。
|
||||
func respondFreshToken(c *gin.Context, auth *service.AuthService) {
|
||||
token, expires, err := auth.RenewToken(c.Request.Context(), c.GetString(usernameKey), bearerToken(c), sessionMetaOf(c))
|
||||
if err != nil {
|
||||
c.Status(http.StatusNoContent)
|
||||
return
|
||||
@@ -138,7 +145,7 @@ func (h *authxHandler) respondFreshToken(c *gin.Context) {
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/revoke-sessions [post]
|
||||
func (h *authxHandler) revokeSessions(c *gin.Context) {
|
||||
token, expires, err := h.auth.RevokeSessions(c.Request.Context(), c.GetString(usernameKey))
|
||||
token, expires, err := h.auth.RevokeSessions(c.Request.Context(), c.GetString(usernameKey), bearerToken(c), sessionMetaOf(c), tokenProofOf(c))
|
||||
if err != nil {
|
||||
respondError(c, err)
|
||||
return
|
||||
@@ -183,7 +190,7 @@ func (h *authxHandler) updateCredentials(c *gin.Context) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
finalName, err := h.auth.UpdateCredentials(c.Request.Context(), c.GetString(usernameKey), req)
|
||||
finalName, err := h.auth.UpdateCredentials(c.Request.Context(), c.GetString(usernameKey), req, tokenProofOf(c))
|
||||
if errors.Is(err, service.ErrCredentialConfirm) {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -197,7 +204,7 @@ func (h *authxHandler) updateCredentials(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
c.Set(usernameKey, finalName)
|
||||
h.respondFreshToken(c)
|
||||
respondFreshToken(c, h.auth)
|
||||
}
|
||||
|
||||
// updatePasswordLogin 保存密码登录禁用开关;开启需至少绑定一个外部身份。
|
||||
@@ -218,7 +225,7 @@ func (h *authxHandler) updatePasswordLogin(c *gin.Context) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
err := h.auth.SetPasswordLoginDisabled(c.Request.Context(), c.GetString(usernameKey), *req.Disabled)
|
||||
err := h.auth.SetPasswordLoginDisabled(c.Request.Context(), c.GetString(usernameKey), *req.Disabled, tokenProofOf(c))
|
||||
if errors.Is(err, service.ErrNeedIdentity) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -227,26 +234,31 @@ func (h *authxHandler) updatePasswordLogin(c *gin.Context) {
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
h.respondFreshToken(c)
|
||||
respondFreshToken(c, h.auth)
|
||||
}
|
||||
|
||||
// ---- OAuth ----
|
||||
|
||||
// oauthProviders 返回已配置的 provider 列表(公开,登录页展示按钮),
|
||||
// 并附带密码登录禁用开关:开启且存在可用外部身份时,登录页隐藏密码表单;
|
||||
// provider 清空时不下发禁用,保证界面始终留有登录入口(后端 Login 仍会拒绝)。
|
||||
// 并附带密码登录禁用开关与通行密钥登录可用性:开关开启且存在可用外部身份时,
|
||||
// 登录页隐藏密码表单;provider 清空时不下发禁用,保证界面始终留有登录入口
|
||||
// (后端 Login 仍会拒绝)。
|
||||
//
|
||||
// @Summary 外部登录 provider 列表
|
||||
// @Tags 认证
|
||||
// @Success 200 {object} oauthProvidersResponse "providers 与 passwordLoginDisabled"
|
||||
// @Success 200 {object} oauthProvidersResponse "providers、passwordLoginDisabled 与 passkeyLogin"
|
||||
// @Router /api/v1/auth/oauth/providers [get]
|
||||
func (h *authxHandler) oauthProviders(c *gin.Context) {
|
||||
providers := h.oauth.Providers(c.Request.Context())
|
||||
passkey := h.passkeys != nil && h.passkeys.HasAny(c.Request.Context())
|
||||
wallet := h.wallets != nil && h.wallets.HasAny(c.Request.Context())
|
||||
disabled := false
|
||||
if off, err := h.auth.PasswordLoginDisabled(c.Request.Context()); err == nil && off && len(providers) > 0 {
|
||||
// 有任一免密入口(OAuth/通行密钥/钱包)才如实下发禁用状态;全部缺席时隐藏,
|
||||
// 保证登录页始终留有入口(后端 Login 仍会拒绝)
|
||||
if off, err := h.auth.PasswordLoginDisabled(c.Request.Context()); err == nil && off && (len(providers) > 0 || passkey || wallet) {
|
||||
disabled = true
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"providers": providers, "passwordLoginDisabled": disabled})
|
||||
c.JSON(http.StatusOK, gin.H{"providers": providers, "passwordLoginDisabled": disabled, "passkeyLogin": passkey, "walletLogin": wallet})
|
||||
}
|
||||
|
||||
// oauthAuthorize 返回授权跳转 URL;mode=bind 需有效 JWT(绑定到当前账号),login 公开。
|
||||
@@ -256,6 +268,9 @@ func (h *authxHandler) oauthProviders(c *gin.Context) {
|
||||
// @Param provider path string true "oidc / github"
|
||||
// @Param mode query string false "bind=绑定当前账号(需 Bearer),缺省登录"
|
||||
// @Success 200 {object} urlResponse "url"
|
||||
// @Failure 401 {object} errorResponse "bind 模式未携带有效 Bearer"
|
||||
// @Failure 404 {object} errorResponse "未知 provider"
|
||||
// @Failure 409 {object} errorResponse "provider 未配置 / 面板地址缺失 / 已禁用"
|
||||
// @Router /api/v1/auth/oauth/{provider}/authorize [get]
|
||||
func (h *authxHandler) oauthAuthorize(c *gin.Context) {
|
||||
provider := c.Param("provider")
|
||||
@@ -272,7 +287,7 @@ func (h *authxHandler) oauthAuthorize(c *gin.Context) {
|
||||
}
|
||||
mode, username = "bind", name
|
||||
}
|
||||
authURL, err := h.oauth.AuthorizeURL(c.Request.Context(), provider, mode, username)
|
||||
authURL, err := h.oauth.AuthorizeURL(c.Request.Context(), provider, mode, username, bearerToken(c))
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrOAuthNotConfigured) || errors.Is(err, service.ErrOAuthNoAppURL) || errors.Is(err, service.ErrOAuthDisabled) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
@@ -310,7 +325,7 @@ func (h *authxHandler) oauthCallback(c *gin.Context) {
|
||||
start := time.Now()
|
||||
provider := c.Param("provider")
|
||||
token, username, mode, err := h.oauth.HandleCallback(
|
||||
c.Request.Context(), provider, c.Query("state"), c.Query("code"))
|
||||
c.Request.Context(), provider, c.Query("state"), c.Query("code"), sessionMetaOf(c))
|
||||
if err != nil {
|
||||
h.recordOauth(c, username, http.StatusUnauthorized, oauthErrText(err), start)
|
||||
target := "/login"
|
||||
@@ -390,7 +405,7 @@ func (h *authxHandler) unbindIdentity(c *gin.Context) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid id"})
|
||||
return
|
||||
}
|
||||
if err := h.oauth.Unbind(c.Request.Context(), c.GetString(usernameKey), uint(id)); err != nil {
|
||||
if err := h.oauth.Unbind(c.Request.Context(), c.GetString(usernameKey), uint(id), tokenProofOf(c)); err != nil {
|
||||
if errors.Is(err, service.ErrLastIdentity) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
@@ -398,7 +413,7 @@ func (h *authxHandler) unbindIdentity(c *gin.Context) {
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
h.respondFreshToken(c)
|
||||
respondFreshToken(c, h.auth)
|
||||
}
|
||||
|
||||
// ---- OAuth provider 设置(JWT 组内) ----
|
||||
@@ -426,6 +441,9 @@ func (h *authxHandler) getOAuthSettings(c *gin.Context, settings *service.Settin
|
||||
// @Tags 设置
|
||||
// @Param body body service.UpdateOAuthInput true "出现的字段才会被更新"
|
||||
// @Success 200 {object} service.OAuthProvidersView
|
||||
// @Failure 400 {object} errorResponse "请求体非法"
|
||||
// @Failure 401 {object} errorResponse "请求处理期间会话已撤销"
|
||||
// @Failure 409 {object} errorResponse "密码登录禁用期间,该变更将移除最后可用登录方式"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/settings/oauth [patch]
|
||||
func (h *authxHandler) updateOAuthSettings(c *gin.Context, settings *service.SettingService) {
|
||||
@@ -434,7 +452,13 @@ func (h *authxHandler) updateOAuthSettings(c *gin.Context, settings *service.Set
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := settings.UpdateOAuth(c.Request.Context(), req); err != nil {
|
||||
err := settings.UpdateOAuthAuthenticated(
|
||||
c.Request.Context(), req, h.auth, c.GetString(usernameKey), tokenProofOf(c))
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrProviderLastLogin) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -26,7 +26,7 @@ func newIconTestRouter(t *testing.T) (*gin.Engine, string, uint, *service.System
|
||||
t.Helper()
|
||||
router, auth, logs, db := newTestRouterDB(t)
|
||||
id := seedIconConfig(t, db)
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "", service.SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
|
||||
@@ -2,6 +2,7 @@ package api
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"sync"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
@@ -202,6 +203,7 @@ func (h *ociConfigHandler) updateInstance(c *gin.Context) {
|
||||
// @Param id path int true "配置 ID"
|
||||
// @Param instanceId path string true "instanceId"
|
||||
// @Success 204 "无内容"
|
||||
// @Failure 409 {object} errorResponse "该实例的生命周期操作正在处理中"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/oci-configs/{id}/instances/{instanceId} [delete]
|
||||
func (h *ociConfigHandler) terminateInstance(c *gin.Context) {
|
||||
@@ -209,20 +211,43 @@ func (h *ociConfigHandler) terminateInstance(c *gin.Context) {
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
instanceID := c.Param("instanceId")
|
||||
if !beginInstanceLifecycleOperation(instanceID) {
|
||||
respondInstanceLifecycleBusy(c)
|
||||
return
|
||||
}
|
||||
defer endInstanceLifecycleOperation(instanceID)
|
||||
preserve := c.Query("preserveBootVolume") == "true"
|
||||
if err := h.svc.TerminateInstance(c.Request.Context(), id, c.Query("region"), c.Param("instanceId"), preserve); err != nil {
|
||||
if err := h.svc.TerminateInstance(c.Request.Context(), id, c.Query("region"), instanceID, preserve); err != nil {
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
|
||||
// instanceLifecycleGuard 按实例 OCID 互斥 HTTP 请求飞行期内的终止与电源操作。
|
||||
var instanceLifecycleGuard sync.Map
|
||||
|
||||
func beginInstanceLifecycleOperation(instanceID string) bool {
|
||||
_, busy := instanceLifecycleGuard.LoadOrStore(instanceID, struct{}{})
|
||||
return !busy
|
||||
}
|
||||
|
||||
func endInstanceLifecycleOperation(instanceID string) {
|
||||
instanceLifecycleGuard.Delete(instanceID)
|
||||
}
|
||||
|
||||
func respondInstanceLifecycleBusy(c *gin.Context) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": "该实例的生命周期操作正在处理中"})
|
||||
}
|
||||
|
||||
// @Summary 实例电源操作(启停/重启)
|
||||
// @Tags 计算
|
||||
// @Param id path int true "配置 ID"
|
||||
// @Param instanceId path string true "instanceId"
|
||||
// @Param body body instanceActionRequest true "请求体"
|
||||
// @Success 200 {object} oci.Instance
|
||||
// @Failure 409 {object} errorResponse "该实例的生命周期操作正在处理中"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/oci-configs/{id}/instances/{instanceId}/action [post]
|
||||
func (h *ociConfigHandler) instanceAction(c *gin.Context) {
|
||||
@@ -235,7 +260,13 @@ func (h *ociConfigHandler) instanceAction(c *gin.Context) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
instance, err := h.svc.InstanceAction(c.Request.Context(), id, req.Region, c.Param("instanceId"), req.Action)
|
||||
instanceID := c.Param("instanceId")
|
||||
if !beginInstanceLifecycleOperation(instanceID) {
|
||||
respondInstanceLifecycleBusy(c)
|
||||
return
|
||||
}
|
||||
defer endInstanceLifecycleOperation(instanceID)
|
||||
instance, err := h.svc.InstanceAction(c.Request.Context(), id, req.Region, instanceID, req.Action)
|
||||
if err != nil {
|
||||
respondError(c, err)
|
||||
return
|
||||
|
||||
@@ -0,0 +1,78 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
)
|
||||
|
||||
func countLifecycleWinners(instanceID string, workers int) int {
|
||||
start := make(chan struct{})
|
||||
results := make(chan bool, workers)
|
||||
var wg sync.WaitGroup
|
||||
for range workers {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
<-start
|
||||
results <- beginInstanceLifecycleOperation(instanceID)
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
wg.Wait()
|
||||
close(results)
|
||||
winners := 0
|
||||
for won := range results {
|
||||
if won {
|
||||
winners++
|
||||
}
|
||||
}
|
||||
return winners
|
||||
}
|
||||
|
||||
func TestBeginInstanceLifecycleOperationSingleFlight(t *testing.T) {
|
||||
instanceID := "ocid1.instance.oc1..single-flight"
|
||||
instanceLifecycleGuard.Delete(instanceID)
|
||||
t.Cleanup(func() { instanceLifecycleGuard.Delete(instanceID) })
|
||||
if got, want := countLifecycleWinners(instanceID, 16), 1; got != want {
|
||||
t.Errorf("winners = %d, want %d", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func instanceOperationRouter() *gin.Engine {
|
||||
h := &ociConfigHandler{}
|
||||
r := gin.New()
|
||||
r.DELETE("/oci-configs/:id/instances/:instanceId", h.terminateInstance)
|
||||
r.POST("/oci-configs/:id/instances/:instanceId/action", h.instanceAction)
|
||||
return r
|
||||
}
|
||||
|
||||
func TestInstanceLifecycleOperationRejectsLaterRequest(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
instanceID := "ocid1.instance.oc1..same-physical-instance"
|
||||
instanceLifecycleGuard.Store(instanceID, struct{}{})
|
||||
t.Cleanup(func() { instanceLifecycleGuard.Delete(instanceID) })
|
||||
r := instanceOperationRouter()
|
||||
tests := []struct {
|
||||
name, method, path, body string
|
||||
}{
|
||||
{"action先占位时跨配置terminate被拒", http.MethodDelete, "/oci-configs/2/instances/" + instanceID, ""},
|
||||
{"terminate先占位时action被拒", http.MethodPost, "/oci-configs/1/instances/" + instanceID + "/action", `{"action":"STOP"}`},
|
||||
{"action先占位时第二个action被拒", http.MethodPost, "/oci-configs/3/instances/" + instanceID + "/action", `{"action":"RESET"}`},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
req := httptest.NewRequest(tt.method, tt.path, strings.NewReader(tt.body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
w := httptest.NewRecorder()
|
||||
r.ServeHTTP(w, req)
|
||||
if got, want := w.Code, http.StatusConflict; got != want {
|
||||
t.Errorf("status = %d, want %d, body %s", got, want, w.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
@@ -12,11 +13,17 @@ import (
|
||||
"oci-portal/internal/service"
|
||||
)
|
||||
|
||||
// usernameKey 是鉴权通过后写入 gin.Context 的用户名键。
|
||||
const usernameKey = "username"
|
||||
// usernameKey 是鉴权通过后写入 gin.Context 的用户名键;
|
||||
// tokenVerKey/tokenJtiKey 是鉴权时观察到的令牌快照,敏感事务提交前复核。
|
||||
const (
|
||||
usernameKey = "username"
|
||||
tokenVerKey = "tokenVer"
|
||||
tokenJtiKey = "tokenJti"
|
||||
)
|
||||
|
||||
// RequireAuth 校验 Authorization: Bearer 令牌(签名、有效期与令牌版本),
|
||||
// 通过后把用户名放进上下文。
|
||||
// 通过后把用户名与令牌快照(版本/jti)放进上下文——快照供敏感事务在
|
||||
// 行锁下复核,防「请求挂起期间令牌被撤销,恢复后仍完成变更」的在途绕过。
|
||||
func RequireAuth(auth *service.AuthService) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
token, ok := strings.CutPrefix(c.GetHeader("Authorization"), "Bearer ")
|
||||
@@ -24,16 +31,25 @@ func RequireAuth(auth *service.AuthService) gin.HandlerFunc {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "missing bearer token"})
|
||||
return
|
||||
}
|
||||
username, err := auth.ParseToken(c.Request.Context(), token)
|
||||
username, proof, err := auth.ParseTokenProof(c.Request.Context(), token)
|
||||
if err != nil {
|
||||
c.AbortWithStatusJSON(http.StatusUnauthorized, gin.H{"error": "invalid or expired token"})
|
||||
return
|
||||
}
|
||||
c.Set(usernameKey, username)
|
||||
c.Set(tokenVerKey, proof.Ver)
|
||||
c.Set(tokenJtiKey, proof.Jti)
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// tokenProofOf 取出鉴权时的令牌快照,交给敏感 service 事务复核。
|
||||
func tokenProofOf(c *gin.Context) service.TokenProof {
|
||||
v, _ := c.Get(tokenVerKey)
|
||||
ver, _ := v.(uint)
|
||||
return service.TokenProof{Ver: ver, Jti: c.GetString(tokenJtiKey)}
|
||||
}
|
||||
|
||||
// bodyLimit 给请求体套上限(S-03):超限读取由 MaxBytesReader 截断报错,
|
||||
// 绑定失败返回 400;webhook 入口另有独立 256KiB 自限,不经此中间件。
|
||||
func bodyLimit(n int64) gin.HandlerFunc {
|
||||
@@ -45,7 +61,7 @@ func bodyLimit(n int64) gin.HandlerFunc {
|
||||
}
|
||||
}
|
||||
|
||||
// systemLogMiddleware 把写请求(POST/PUT/DELETE)异步记入系统日志,只读请求跳过。
|
||||
// systemLogMiddleware 把写请求(POST/PUT/PATCH/DELETE)异步记入系统日志,只读请求跳过。
|
||||
// 只记方法、路径等元数据,绝不读请求体——请求体可能含私钥 / 口令等敏感数据。
|
||||
func systemLogMiddleware(logs *service.SystemLogService) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
@@ -101,17 +117,23 @@ func errMsgOf(status int, head []byte) string {
|
||||
return truncateLogField(body.Error, 256)
|
||||
}
|
||||
|
||||
// truncateLogField 按字节截断留痕字段,防超长 UA / 错误串撑爆列宽。
|
||||
// truncateLogField 按字节截断留痕字段,防超长 UA / 错误串撑爆列宽;
|
||||
// 截断点回退到 rune 边界,避免产生非法 UTF-8(严格 MySQL/PG 会拒写整行)。
|
||||
func truncateLogField(s string, max int) string {
|
||||
if len(s) <= max {
|
||||
return s
|
||||
}
|
||||
for max > 0 && !utf8.RuneStart(s[max]) {
|
||||
max--
|
||||
}
|
||||
return s[:max]
|
||||
}
|
||||
|
||||
// isWriteMethod 判断是否需要留痕的写方法。
|
||||
// isWriteMethod 判断是否需要留痕的写方法;PATCH 承载 OAuth / 安全设置等
|
||||
// 敏感配置变更,必须与 POST/PUT/DELETE 一样进入审计。
|
||||
func isWriteMethod(m string) bool {
|
||||
return m == http.MethodPost || m == http.MethodPut || m == http.MethodDelete
|
||||
return m == http.MethodPost || m == http.MethodPut ||
|
||||
m == http.MethodPatch || m == http.MethodDelete
|
||||
}
|
||||
|
||||
// requestPath 优先取注册的路由模板(含 :id 占位符,避免把敏感 query 记入日志),
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
// PATCH 承载 OAuth / 安全设置变更,必须纳入审计;只读方法不留痕。
|
||||
func TestIsWriteMethod(t *testing.T) {
|
||||
tests := []struct {
|
||||
method string
|
||||
want bool
|
||||
}{
|
||||
{http.MethodPost, true},
|
||||
{http.MethodPut, true},
|
||||
{http.MethodPatch, true},
|
||||
{http.MethodDelete, true},
|
||||
{http.MethodGet, false},
|
||||
{http.MethodHead, false},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
if got := isWriteMethod(tt.method); got != tt.want {
|
||||
t.Errorf("isWriteMethod(%s) = %v, want %v", tt.method, got, tt.want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestTruncateLogField(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in string
|
||||
max int
|
||||
want string
|
||||
}{
|
||||
{name: "短串原样", in: "abc", max: 10, want: "abc"},
|
||||
{name: "ASCII 截断", in: "abcdef", max: 4, want: "abcd"},
|
||||
{name: "中文边界回退", in: "中文日志", max: 4, want: "中"},
|
||||
{name: "恰好落在边界", in: "中文", max: 3, want: "中"},
|
||||
{name: "emoji 拦腰", in: "a😀b", max: 3, want: "a"},
|
||||
{name: "超长 UA 不产非法序列", in: strings.Repeat("界", 200), max: 255, want: strings.Repeat("界", 85)},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := truncateLogField(tt.in, tt.max)
|
||||
if got != tt.want {
|
||||
t.Errorf("truncateLogField(%q, %d) = %q, want %q", tt.in, tt.max, got, tt.want)
|
||||
}
|
||||
if !utf8.ValidString(got) {
|
||||
t.Errorf("result %q is not valid UTF-8", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -277,6 +277,12 @@ func respondError(c *gin.Context, err error) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "资源不存在"})
|
||||
return
|
||||
}
|
||||
// 请求处理期间令牌已失效(撤销/注销/并发敏感操作后到):按未认证处理,
|
||||
// 前端「发送时令牌==当前令牌」快照决定是否登出
|
||||
if errors.Is(err, service.ErrTokenStale) {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
id := newRequestID()
|
||||
log.Printf("[ERR %s] %s %s: %v", id, c.Request.Method, requestPath(c), err)
|
||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "服务器内部错误", "requestId": id})
|
||||
|
||||
@@ -0,0 +1,196 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
"oci-portal/internal/service"
|
||||
)
|
||||
|
||||
// passkeyHandler 处理通行密钥(WebAuthn)注册、登录与凭据管理接口。
|
||||
type passkeyHandler struct {
|
||||
svc *service.PasskeyService
|
||||
auth *service.AuthService
|
||||
logs *service.SystemLogService
|
||||
}
|
||||
|
||||
// registerBegin 生成注册 options;sessionId 需原样带回 finish。
|
||||
//
|
||||
// @Summary 发起通行密钥注册
|
||||
// @Tags 认证
|
||||
// @Success 200 {object} passkeyOptionsResponse "sessionId 与 WebAuthn options"
|
||||
// @Failure 409 {object} errorResponse "面板地址未设置或数量达上限"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/passkey/register/begin [post]
|
||||
func (h *passkeyHandler) registerBegin(c *gin.Context) {
|
||||
sid, opts, err := h.svc.BeginRegister(c.Request.Context(), c.GetString(usernameKey))
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrPasskeyNoAppURL) || errors.Is(err, service.ErrPasskeyLimit) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"sessionId": sid, "options": opts})
|
||||
}
|
||||
|
||||
// passkeyFinishRequest 是注册/登录 finish 的请求体;credential 为浏览器
|
||||
// navigator.credentials 返回的 WebAuthn JSON,原样透传后端解析。
|
||||
type passkeyFinishRequest struct {
|
||||
SessionId string `json:"sessionId" binding:"required,max=64"`
|
||||
Name string `json:"name" binding:"max=64"`
|
||||
Credential json.RawMessage `json:"credential" binding:"required"`
|
||||
}
|
||||
|
||||
// registerFinish 校验注册响应并保存凭据;属敏感变更,换发新令牌。
|
||||
//
|
||||
// @Summary 完成通行密钥注册
|
||||
// @Tags 认证
|
||||
// @Param body body passkeyFinishRequest true "sessionId、名称与凭据响应"
|
||||
// @Success 200 {object} tokenResponse "已注册,返回换发的新 token"
|
||||
// @Success 204 "已注册但新 token 签发失败,需重新登录"
|
||||
// @Failure 400 {object} errorResponse "请求体非法、会话过期或凭据校验失败"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/passkey/register/finish [post]
|
||||
func (h *passkeyHandler) registerFinish(c *gin.Context) {
|
||||
var req passkeyFinishRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
err := h.svc.FinishRegister(c.Request.Context(), c.GetString(usernameKey),
|
||||
req.SessionId, req.Name, bytes.NewReader(req.Credential), tokenProofOf(c))
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrPasskeySession) || errors.Is(err, service.ErrPasskeyVerify) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
respondFreshToken(c, h.auth)
|
||||
}
|
||||
|
||||
// list 列出当前账号的通行密钥。
|
||||
//
|
||||
// @Summary 通行密钥列表
|
||||
// @Tags 认证
|
||||
// @Success 200 {object} itemsResponse[model.UserPasskey] "items"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/passkeys [get]
|
||||
func (h *passkeyHandler) list(c *gin.Context) {
|
||||
items, err := h.svc.List(c.Request.Context(), c.GetString(usernameKey))
|
||||
if err != nil {
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||
}
|
||||
|
||||
// remove 删除通行密钥;属敏感变更,换发新令牌。
|
||||
// 密码登录禁用期间,最后一个免密登录方式不可删除(防自锁)。
|
||||
//
|
||||
// @Summary 删除通行密钥
|
||||
// @Tags 认证
|
||||
// @Param id path int true "凭据 ID"
|
||||
// @Success 200 {object} tokenResponse "已删除,返回换发的新 token"
|
||||
// @Success 204 "已删除但新 token 签发失败,需重新登录"
|
||||
// @Failure 409 {object} errorResponse "密码登录已禁用且这是最后一个免密登录方式"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/passkeys/{id} [delete]
|
||||
func (h *passkeyHandler) remove(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid id"})
|
||||
return
|
||||
}
|
||||
if err := h.svc.Remove(c.Request.Context(), c.GetString(usernameKey), uint(id), tokenProofOf(c)); err != nil {
|
||||
if errors.Is(err, service.ErrLastIdentity) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
respondFreshToken(c, h.auth)
|
||||
}
|
||||
|
||||
// loginBegin 生成登录断言 options(公开;无凭据也正常下发,不泄露账号状态)。
|
||||
//
|
||||
// @Summary 发起通行密钥登录
|
||||
// @Tags 认证
|
||||
// @Success 200 {object} passkeyOptionsResponse "sessionId 与 WebAuthn options"
|
||||
// @Failure 409 {object} errorResponse "面板地址未设置"
|
||||
// @Router /api/v1/auth/passkey/login/begin [post]
|
||||
func (h *passkeyHandler) loginBegin(c *gin.Context) {
|
||||
sid, opts, err := h.svc.BeginLogin(c.Request.Context())
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrPasskeyNoAppURL) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"sessionId": sid, "options": opts})
|
||||
}
|
||||
|
||||
// loginFinish 校验断言并签发 JWT;凭 UV 豁免 TOTP。
|
||||
//
|
||||
// @Summary 完成通行密钥登录
|
||||
// @Tags 认证
|
||||
// @Param body body passkeyFinishRequest true "sessionId 与断言响应(name 忽略)"
|
||||
// @Success 200 {object} tokenResponse "token 与 expiresAt"
|
||||
// @Failure 400 {object} errorResponse "请求体非法"
|
||||
// @Failure 401 {object} errorResponse "校验失败"
|
||||
// @Failure 429 {object} errorResponse "连续失败已锁定"
|
||||
// @Router /api/v1/auth/passkey/login/finish [post]
|
||||
func (h *passkeyHandler) loginFinish(c *gin.Context) {
|
||||
var req passkeyFinishRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
start := time.Now()
|
||||
token, expires, username, err := h.svc.FinishLogin(
|
||||
c.Request.Context(), req.SessionId, sessionMetaOf(c), bytes.NewReader(req.Credential))
|
||||
if errors.Is(err, service.ErrLoginLocked) {
|
||||
h.recordLogin(c, username, http.StatusTooManyRequests, "连续失败已锁定", start)
|
||||
c.JSON(http.StatusTooManyRequests, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
// 统一 401 文案,不区分会话过期/校验失败,防探测
|
||||
h.recordLogin(c, username, http.StatusUnauthorized, service.ErrPasskeyVerify.Error(), start)
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": service.ErrPasskeyVerify.Error()})
|
||||
return
|
||||
}
|
||||
h.recordLogin(c, username, http.StatusOK, "", start)
|
||||
c.JSON(http.StatusOK, gin.H{"token": token, "expiresAt": expires})
|
||||
}
|
||||
|
||||
// recordLogin 记录通行密钥登录成败——公开端点不经系统日志中间件,
|
||||
// 登录失败是安全关键事件,必须留痕;失败时账号未知,username 为空。
|
||||
func (h *passkeyHandler) recordLogin(c *gin.Context, username string, status int, errMsg string, start time.Time) {
|
||||
if h.logs == nil {
|
||||
return
|
||||
}
|
||||
h.logs.Record(model.SystemLog{
|
||||
Username: username,
|
||||
Method: http.MethodPost,
|
||||
Path: requestPath(c),
|
||||
Status: status,
|
||||
DurationMs: time.Since(start).Milliseconds(),
|
||||
ClientIP: requestIP(c),
|
||||
UserAgent: truncateLogField(c.Request.UserAgent(), 256),
|
||||
ErrMsg: errMsg,
|
||||
})
|
||||
}
|
||||
@@ -77,7 +77,9 @@ func IPRateMiddleware(settings *service.SettingService) gin.HandlerFunc {
|
||||
return func(c *gin.Context) {
|
||||
sec := settings.SecurityCached()
|
||||
if !lm.get(requestIP(c), sec.IPRateRPS, sec.IPRateBurst).Allow() {
|
||||
c.AbortWithStatusJSON(http.StatusTooManyRequests, gin.H{"error": "rate limit exceeded"})
|
||||
// code 供前端区分全局 IP 限流与登录守卫锁定(后者不带 code,留在表单内提示)
|
||||
c.AbortWithStatusJSON(http.StatusTooManyRequests,
|
||||
gin.H{"error": "rate limit exceeded", "code": "RateLimited"})
|
||||
return
|
||||
}
|
||||
c.Next()
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
|
||||
// NewRouter 组装 HTTP 路由骨架:中间件链、公开/鉴权分组与各域注册;
|
||||
// 具体路由按域拆在 routes_*.go。/auth/login 公开,业务路由要求 JWT。
|
||||
func NewRouter(auth *service.AuthService, oauth *service.OAuthService, ociConfigs *service.OciConfigService, tasks *service.TaskService, console *service.ConsoleService, settings *service.SettingService, notifier *service.Notifier, systemLogs *service.SystemLogService, logEvents *service.LogEventService, proxies *service.ProxyService, aiGateway *service.AiGatewayService) *gin.Engine {
|
||||
func NewRouter(auth *service.AuthService, oauth *service.OAuthService, passkeys *service.PasskeyService, wallets *service.WalletService, ociConfigs *service.OciConfigService, tasks *service.TaskService, console *service.ConsoleService, settings *service.SettingService, notifier *service.Notifier, systemLogs *service.SystemLogService, logEvents *service.LogEventService, proxies *service.ProxyService, aiGateway *service.AiGatewayService) *gin.Engine {
|
||||
r := gin.New()
|
||||
// 禁用 gin 内置代理信任:真实 IP 统一由 RealIPMiddleware 按安全设置解析
|
||||
_ = r.SetTrustedProxies(nil)
|
||||
@@ -23,15 +23,15 @@ func NewRouter(auth *service.AuthService, oauth *service.OAuthService, ociConfig
|
||||
|
||||
// 面板 API 请求体统一 1MB 上限(S-03);AI 网关按对话体量单独 10MB
|
||||
v1 := r.Group("/api/v1", bodyLimit(1<<20))
|
||||
registerAuthPublic(v1, auth, oauth, systemLogs)
|
||||
registerAuthPublic(v1, auth, oauth, passkeys, wallets, systemLogs)
|
||||
// AI 网关对外端点:独立密钥鉴权,挂在全局 Use 之后,仍受 IP 限速与 Recovery 保护
|
||||
registerAiGateway(r, aiGateway)
|
||||
|
||||
// 系统日志中间件挂在鉴权之后,写请求自动留痕(webconsole ws 为 GET,天然跳过)
|
||||
secured := v1.Group("", RequireAuth(auth), systemLogMiddleware(systemLogs))
|
||||
registerAuthSecured(secured, auth, oauth, settings, systemLogs)
|
||||
registerAuthSecured(secured, auth, oauth, passkeys, settings, systemLogs)
|
||||
registerConsole(v1, secured, console, auth)
|
||||
registerSettings(secured, settings, notifier, systemLogs, proxies)
|
||||
registerSettings(secured, auth, settings, notifier, systemLogs, proxies)
|
||||
registerTasksAndLogs(secured, tasks, logEvents)
|
||||
registerOci(secured, ociConfigs)
|
||||
// 上传类接口独立成组:文件本体上限 1MB,multipart 边界与字段头另占空间,
|
||||
|
||||
@@ -86,7 +86,7 @@ func newTestRouterDB(t *testing.T) (*gin.Engine, *service.AuthService, *service.
|
||||
t.Fatalf("db handle: %v", err)
|
||||
}
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
if err := db.AutoMigrate(&model.User{}, &model.OciConfig{}, &model.Task{}, &model.TaskLog{}, &model.Setting{}, &model.SystemLog{}, &model.LogEvent{}, &model.Proxy{}, &model.AiKey{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}); err != nil {
|
||||
if err := db.AutoMigrate(&model.User{}, &model.UserPasskey{}, &model.UserSession{}, &model.OciConfig{}, &model.Task{}, &model.TaskLog{}, &model.Setting{}, &model.SystemLog{}, &model.LogEvent{}, &model.Proxy{}, &model.AiKey{}, &model.AiChannel{}, &model.AiModelCache{}, &model.AiCallLog{}); err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
cipher, err := crypto.NewCipher("test-key")
|
||||
@@ -104,7 +104,9 @@ func newTestRouterDB(t *testing.T) (*gin.Engine, *service.AuthService, *service.
|
||||
systemLogs := service.NewSystemLogService(db)
|
||||
logEvents := service.NewLogEventService(db)
|
||||
oauth := service.NewOAuthService(db, settings, auth)
|
||||
r := NewRouter(auth, oauth, ociConfigs, tasks, service.NewConsoleService(ociConfigs), settings, notifier, systemLogs, logEvents, service.NewProxyService(db, cipher), service.NewAiGatewayService(db, ociConfigs, nullClient{}))
|
||||
passkeys := service.NewPasskeyService(db, settings, auth)
|
||||
wallets := service.NewWalletService(db, settings, auth)
|
||||
r := NewRouter(auth, oauth, passkeys, wallets, ociConfigs, tasks, service.NewConsoleService(ociConfigs), settings, notifier, systemLogs, logEvents, service.NewProxyService(db, cipher), service.NewAiGatewayService(db, ociConfigs, nullClient{}))
|
||||
return r, auth, systemLogs, db
|
||||
}
|
||||
|
||||
@@ -146,7 +148,7 @@ func TestLoginEndpoint(t *testing.T) {
|
||||
|
||||
func TestSecuredRoutesRequireToken(t *testing.T) {
|
||||
r, auth, _ := newTestRouter(t)
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "", service.SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
@@ -171,7 +173,7 @@ func TestSecuredRoutesRequireToken(t *testing.T) {
|
||||
|
||||
func TestSystemLogsEndpoint(t *testing.T) {
|
||||
r, auth, logs := newTestRouter(t)
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "", service.SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
|
||||
@@ -6,23 +6,44 @@ import (
|
||||
"oci-portal/internal/service"
|
||||
)
|
||||
|
||||
// registerAuthPublic 公开登录/OAuth 路由(不经 JWT)。
|
||||
func registerAuthPublic(v1 *gin.RouterGroup, auth *service.AuthService, oauth *service.OAuthService, systemLogs *service.SystemLogService) {
|
||||
// registerAuthPublic 公开登录/OAuth/通行密钥/钱包路由(不经 JWT)。
|
||||
func registerAuthPublic(v1 *gin.RouterGroup, auth *service.AuthService, oauth *service.OAuthService, passkeys *service.PasskeyService, wallets *service.WalletService, systemLogs *service.SystemLogService) {
|
||||
ah := &authHandler{svc: auth, logs: systemLogs}
|
||||
v1.POST("/auth/login", ah.login)
|
||||
|
||||
// 外部身份登录:provider 列表 / 授权跳转(bind 模式 handler 内校验 JWT)/ 回调
|
||||
ax := &authxHandler{auth: auth, oauth: oauth, logs: systemLogs}
|
||||
ax := &authxHandler{auth: auth, oauth: oauth, passkeys: passkeys, wallets: wallets, logs: systemLogs}
|
||||
v1.GET("/auth/oauth/providers", ax.oauthProviders)
|
||||
v1.GET("/auth/oauth/:provider/authorize", ax.oauthAuthorize)
|
||||
v1.GET("/auth/oauth/:provider/callback", ax.oauthCallback)
|
||||
|
||||
// 通行密钥登录:断言 options 公开下发,finish 内部接入登录守卫
|
||||
pk := &passkeyHandler{svc: passkeys, auth: auth, logs: systemLogs}
|
||||
v1.POST("/auth/passkey/login/begin", pk.loginBegin)
|
||||
v1.POST("/auth/passkey/login/finish", pk.loginFinish)
|
||||
|
||||
// 钱包登录/绑定:挑战-签名两段式(bind 模式 handler 内校验 JWT)
|
||||
wl := &walletHandler{svc: wallets, auth: auth, logs: systemLogs}
|
||||
v1.POST("/auth/wallet/challenge", wl.challenge)
|
||||
v1.POST("/auth/wallet/verify", wl.verify)
|
||||
}
|
||||
|
||||
// registerAuthSecured 登出、两步验证、外部身份管理(JWT 组内)。
|
||||
func registerAuthSecured(secured *gin.RouterGroup, auth *service.AuthService, oauth *service.OAuthService, settings *service.SettingService, systemLogs *service.SystemLogService) {
|
||||
// registerAuthSecured 登出、两步验证、外部身份与通行密钥管理(JWT 组内)。
|
||||
func registerAuthSecured(secured *gin.RouterGroup, auth *service.AuthService, oauth *service.OAuthService, passkeys *service.PasskeyService, settings *service.SettingService, systemLogs *service.SystemLogService) {
|
||||
ah := &authHandler{svc: auth, logs: systemLogs}
|
||||
secured.POST("/auth/logout", ah.logout)
|
||||
|
||||
pk := &passkeyHandler{svc: passkeys, auth: auth}
|
||||
secured.POST("/auth/passkey/register/begin", pk.registerBegin)
|
||||
secured.POST("/auth/passkey/register/finish", pk.registerFinish)
|
||||
secured.GET("/auth/passkeys", pk.list)
|
||||
secured.DELETE("/auth/passkeys/:id", pk.remove)
|
||||
|
||||
// 活跃会话:查看与定点撤销
|
||||
sh := &sessionHandler{auth: auth}
|
||||
secured.GET("/auth/sessions", sh.list)
|
||||
secured.DELETE("/auth/sessions/:id", sh.revoke)
|
||||
|
||||
ax := &authxHandler{auth: auth, oauth: oauth}
|
||||
secured.GET("/auth/totp", ax.totpStatus)
|
||||
secured.POST("/auth/totp/setup", ax.totpSetup)
|
||||
|
||||
@@ -7,11 +7,11 @@ import (
|
||||
)
|
||||
|
||||
// registerSettings 挂载 /settings/* 与 /about,以及系统日志与代理管理。
|
||||
func registerSettings(secured *gin.RouterGroup, settings *service.SettingService, notifier *service.Notifier, systemLogs *service.SystemLogService, proxies *service.ProxyService) {
|
||||
func registerSettings(secured *gin.RouterGroup, auth *service.AuthService, settings *service.SettingService, notifier *service.Notifier, systemLogs *service.SystemLogService, proxies *service.ProxyService) {
|
||||
secured.GET("/about", about)
|
||||
secured.GET("/system-logs", (&systemLogHandler{svc: systemLogs}).list)
|
||||
|
||||
st := &settingsHandler{svc: settings, notifier: notifier}
|
||||
st := &settingsHandler{svc: settings, auth: auth, notifier: notifier}
|
||||
secured.GET("/settings/telegram", st.getTelegram)
|
||||
secured.PUT("/settings/telegram", st.updateTelegram)
|
||||
secured.POST("/settings/telegram/test", st.testTelegram)
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/service"
|
||||
)
|
||||
|
||||
// sessionHandler 处理活跃会话查看与定点撤销接口。
|
||||
type sessionHandler struct {
|
||||
auth *service.AuthService
|
||||
}
|
||||
|
||||
// bearerToken 取 Authorization 头中的原始令牌串;不存在返回空串。
|
||||
func bearerToken(c *gin.Context) string {
|
||||
token, ok := strings.CutPrefix(c.GetHeader("Authorization"), "Bearer ")
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
return token
|
||||
}
|
||||
|
||||
// sessionMetaOf 采集会话登记所需的客户端上下文;登录方式由 service 层填写。
|
||||
func sessionMetaOf(c *gin.Context) service.SessionMeta {
|
||||
return service.SessionMeta{
|
||||
ClientIP: requestIP(c),
|
||||
UserAgent: truncateLogField(c.Request.UserAgent(), 256),
|
||||
}
|
||||
}
|
||||
|
||||
// list 列出当前账号的活跃会话,最近活跃在前,当前会话带 current 标记。
|
||||
//
|
||||
// @Summary 活跃会话列表
|
||||
// @Tags 认证
|
||||
// @Success 200 {object} itemsResponse[service.SessionInfo] "items"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/sessions [get]
|
||||
func (h *sessionHandler) list(c *gin.Context) {
|
||||
items, err := h.auth.ListSessions(c.Request.Context(), c.GetString(usernameKey), bearerToken(c))
|
||||
if err != nil {
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"items": items})
|
||||
}
|
||||
|
||||
// revoke 定点撤销一个其他会话;该会话令牌随即失效,不影响其余会话。
|
||||
//
|
||||
// @Summary 撤销单个会话
|
||||
// @Tags 认证
|
||||
// @Param id path int true "会话 ID"
|
||||
// @Success 204 "已撤销"
|
||||
// @Failure 400 {object} errorResponse "id 非法"
|
||||
// @Failure 404 {object} errorResponse "会话不存在"
|
||||
// @Failure 409 {object} errorResponse "不能撤销当前会话"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/auth/sessions/{id} [delete]
|
||||
func (h *sessionHandler) revoke(c *gin.Context) {
|
||||
id, err := strconv.ParseUint(c.Param("id"), 10, 64)
|
||||
if err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid id"})
|
||||
return
|
||||
}
|
||||
err = h.auth.RevokeSession(c.Request.Context(), c.GetString(usernameKey), bearerToken(c), uint(id))
|
||||
if errors.Is(err, service.ErrSessionCurrent) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
c.JSON(http.StatusNotFound, gin.H{"error": "会话不存在或已失效"})
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
c.Status(http.StatusNoContent)
|
||||
}
|
||||
@@ -12,6 +12,7 @@ import (
|
||||
// settingsHandler 处理系统设置(Telegram 通知)相关请求。
|
||||
type settingsHandler struct {
|
||||
svc *service.SettingService
|
||||
auth *service.AuthService
|
||||
notifier *service.Notifier
|
||||
}
|
||||
|
||||
@@ -323,6 +324,9 @@ func (h *settingsHandler) getSecurity(c *gin.Context) {
|
||||
// @Tags 设置
|
||||
// @Param body body service.SecurityPatch true "出现的字段才会被更新"
|
||||
// @Success 200 {object} service.SecuritySettings
|
||||
// @Failure 400 {object} errorResponse "字段越界或非法"
|
||||
// @Failure 401 {object} errorResponse "请求处理期间会话已撤销"
|
||||
// @Failure 409 {object} errorResponse "密码登录禁用期间,面板地址变更将移除最后可用登录方式"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/settings/security [patch]
|
||||
func (h *settingsHandler) updateSecurity(c *gin.Context) {
|
||||
@@ -331,11 +335,17 @@ func (h *settingsHandler) updateSecurity(c *gin.Context) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if err := h.svc.UpdateSecurity(c.Request.Context(), req); err != nil {
|
||||
err := h.svc.UpdateSecurityAuthenticated(
|
||||
c.Request.Context(), req, h.auth, c.GetString(usernameKey), tokenProofOf(c))
|
||||
if err != nil {
|
||||
if errors.Is(err, service.ErrInvalidSecurity) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if errors.Is(err, service.ErrProviderLastLogin) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -15,6 +15,9 @@ import (
|
||||
// errorResponse 是统一错误响应外壳。
|
||||
type errorResponse struct {
|
||||
Error string `json:"error"`
|
||||
// Code 是机器可读错误码,多数错误不携带;当前仅全局 IP 限流的 429
|
||||
// 返回 "RateLimited",前端据此与登录守卫的 429(无 code)区分
|
||||
Code string `json:"code,omitempty"`
|
||||
}
|
||||
|
||||
// itemsResponse 是 {"items": [...]} 列表外壳。
|
||||
@@ -71,6 +74,20 @@ type credentialsResponse struct {
|
||||
type oauthProvidersResponse struct {
|
||||
Providers []service.ProviderInfo `json:"providers"`
|
||||
PasswordLoginDisabled bool `json:"passwordLoginDisabled"`
|
||||
PasskeyLogin bool `json:"passkeyLogin"`
|
||||
WalletLogin bool `json:"walletLogin"`
|
||||
}
|
||||
|
||||
// walletChallengeResponse 是钱包签名挑战响应;message 需原样 personal_sign。
|
||||
type walletChallengeResponse struct {
|
||||
Nonce string `json:"nonce"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// passkeyOptionsResponse 是 WebAuthn 仪式 options 响应;sessionId 需原样带回 finish。
|
||||
type passkeyOptionsResponse struct {
|
||||
SessionId string `json:"sessionId"`
|
||||
Options json.RawMessage `json:"options"`
|
||||
}
|
||||
|
||||
// urlResponse 是跳转地址响应。
|
||||
|
||||
@@ -439,10 +439,10 @@ func (h *ociConfigHandler) deleteTenantUserApiKey(c *gin.Context) {
|
||||
}
|
||||
|
||||
// @Summary 启用 API Key 为面板签名凭据
|
||||
// @Description 将刚创建的 key 设为本配置签名凭据:验证可用后落库,不删除旧 key;私钥为创建时下发的那份回传。
|
||||
// @Description 将刚创建的 key 设为本配置签名凭据:验证可用后落库,不删除旧 key;私钥为创建时下发的那份回传。userId 给定且异于当前签名用户时,一并把签名用户切换为该用户。
|
||||
// @Tags 租户 IAM
|
||||
// @Param id path int true "配置 ID"
|
||||
// @Param body body object true "请求体:fingerprint 与 privateKey"
|
||||
// @Param body body object true "请求体:fingerprint 与 privateKey,可选 userId(切换签名用户)"
|
||||
// @Success 204 "无内容"
|
||||
// @Security BearerAuth
|
||||
// @Router /api/v1/oci-configs/{id}/activate-api-key [post]
|
||||
@@ -452,6 +452,7 @@ func (h *ociConfigHandler) activateApiKey(c *gin.Context) {
|
||||
return
|
||||
}
|
||||
var req struct {
|
||||
UserID string `json:"userId"`
|
||||
Fingerprint string `json:"fingerprint"`
|
||||
PrivateKey string `json:"privateKey"`
|
||||
}
|
||||
@@ -459,7 +460,7 @@ func (h *ociConfigHandler) activateApiKey(c *gin.Context) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": "fingerprint 与 privateKey 必填"})
|
||||
return
|
||||
}
|
||||
if err := h.svc.ActivateApiKey(c.Request.Context(), id, req.Fingerprint, req.PrivateKey); err != nil {
|
||||
if err := h.svc.ActivateApiKey(c.Request.Context(), id, req.UserID, req.Fingerprint, req.PrivateKey); err != nil {
|
||||
respondError(c, err)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -0,0 +1,161 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
"oci-portal/internal/service"
|
||||
)
|
||||
|
||||
// walletHandler 处理钱包(SIWE)签名挑战与校验接口。
|
||||
type walletHandler struct {
|
||||
svc *service.WalletService
|
||||
auth *service.AuthService
|
||||
logs *service.SystemLogService
|
||||
}
|
||||
|
||||
// walletChallengeRequest 是签名挑战请求;mode=bind 需有效 Bearer(绑定到当前账号)。
|
||||
type walletChallengeRequest struct {
|
||||
Address string `json:"address" binding:"required,max=64"`
|
||||
Mode string `json:"mode" binding:"omitempty,oneof=login bind"`
|
||||
}
|
||||
|
||||
// challenge 下发 EIP-4361 消息与一次性 nonce;前端对消息原样 personal_sign。
|
||||
//
|
||||
// @Summary 发起钱包签名挑战
|
||||
// @Tags 认证
|
||||
// @Param body body walletChallengeRequest true "钱包地址与模式(bind 需 Bearer)"
|
||||
// @Success 200 {object} walletChallengeResponse "nonce 与待签名消息"
|
||||
// @Failure 400 {object} errorResponse "地址格式不正确"
|
||||
// @Failure 401 {object} errorResponse "bind 模式未登录"
|
||||
// @Failure 409 {object} errorResponse "面板地址未设置"
|
||||
// @Router /api/v1/auth/wallet/challenge [post]
|
||||
func (h *walletHandler) challenge(c *gin.Context) {
|
||||
var req walletChallengeRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
mode, username := "login", ""
|
||||
if req.Mode == "bind" {
|
||||
name, ok := (&authxHandler{auth: h.auth}).bearerUser(c)
|
||||
if !ok {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": "绑定需要先登录"})
|
||||
return
|
||||
}
|
||||
mode, username = "bind", name
|
||||
}
|
||||
nonce, message, err := h.svc.Challenge(c.Request.Context(), req.Address, mode, username, bearerToken(c))
|
||||
if err != nil {
|
||||
h.respondChallengeErr(c, err)
|
||||
return
|
||||
}
|
||||
c.JSON(http.StatusOK, gin.H{"nonce": nonce, "message": message})
|
||||
}
|
||||
|
||||
// respondChallengeErr 映射挑战阶段错误。
|
||||
func (h *walletHandler) respondChallengeErr(c *gin.Context, err error) {
|
||||
if errors.Is(err, service.ErrWalletAddress) {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
if errors.Is(err, service.ErrWalletNoAppURL) {
|
||||
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
respondError(c, err)
|
||||
}
|
||||
|
||||
// walletVerifyRequest 是签名校验请求;signature 为 65 字节 r||s||v 的 hex。
|
||||
type walletVerifyRequest struct {
|
||||
Nonce string `json:"nonce" binding:"required,max=64"`
|
||||
Signature string `json:"signature" binding:"required,max=200"`
|
||||
}
|
||||
|
||||
// verify 校验签名:login 模式签发 JWT,bind 模式绑定后换发新令牌。
|
||||
//
|
||||
// @Summary 校验钱包签名
|
||||
// @Tags 认证
|
||||
// @Param body body walletVerifyRequest true "nonce 与签名"
|
||||
// @Success 200 {object} tokenResponse "token 与 expiresAt(bind 模式为换发的新 token)"
|
||||
// @Failure 400 {object} errorResponse "挑战无效已过期,或 bind 模式签名校验失败"
|
||||
// @Failure 401 {object} errorResponse "签名校验失败(login 模式)"
|
||||
// @Failure 403 {object} errorResponse "地址未绑定账号"
|
||||
// @Failure 409 {object} errorResponse "地址已绑定过"
|
||||
// @Failure 429 {object} errorResponse "连续失败已锁定"
|
||||
// @Router /api/v1/auth/wallet/verify [post]
|
||||
func (h *walletHandler) verify(c *gin.Context) {
|
||||
var req walletVerifyRequest
|
||||
if err := c.ShouldBindJSON(&req); err != nil {
|
||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
start := time.Now()
|
||||
res, err := h.svc.Verify(c.Request.Context(), req.Nonce, req.Signature, sessionMetaOf(c))
|
||||
if err != nil {
|
||||
status := walletErrStatus(err)
|
||||
// bind 是已登录流程:验签失败不可回 401,否则前端全局拦截会误登出有效会话
|
||||
if res.Mode == "bind" && status == http.StatusUnauthorized {
|
||||
status = http.StatusBadRequest
|
||||
}
|
||||
h.recordWallet(c, res.Username, status, err, start)
|
||||
c.JSON(status, gin.H{"error": walletErrText(err)})
|
||||
return
|
||||
}
|
||||
h.recordWallet(c, res.Username, http.StatusOK, nil, start)
|
||||
// login 与 bind 同构响应;bind 的 token 为版本递增后换发的新会话
|
||||
c.JSON(http.StatusOK, gin.H{"token": res.Token, "expiresAt": res.ExpiresAt})
|
||||
}
|
||||
|
||||
// walletErrStatus 映射校验阶段错误码。
|
||||
func walletErrStatus(err error) int {
|
||||
switch {
|
||||
case errors.Is(err, service.ErrLoginLocked):
|
||||
return http.StatusTooManyRequests
|
||||
case errors.Is(err, service.ErrWalletChallenge):
|
||||
return http.StatusBadRequest
|
||||
case errors.Is(err, service.ErrWalletNotBound):
|
||||
return http.StatusForbidden
|
||||
case errors.Is(err, service.ErrWalletBound):
|
||||
return http.StatusConflict
|
||||
case errors.Is(err, service.ErrWalletSig):
|
||||
return http.StatusUnauthorized
|
||||
default:
|
||||
return http.StatusInternalServerError
|
||||
}
|
||||
}
|
||||
|
||||
// walletErrText 把流程错误转为用户可读文案;内部错误不透出细节。
|
||||
func walletErrText(err error) string {
|
||||
for _, known := range []error{service.ErrLoginLocked, service.ErrWalletChallenge, service.ErrWalletNotBound, service.ErrWalletBound, service.ErrWalletSig} {
|
||||
if errors.Is(err, known) {
|
||||
return known.Error()
|
||||
}
|
||||
}
|
||||
return "校验失败,请重试"
|
||||
}
|
||||
|
||||
// recordWallet 记录钱包登录/绑定结果——公开端点不经系统日志中间件,须显式留痕。
|
||||
func (h *walletHandler) recordWallet(c *gin.Context, username string, status int, err error, start time.Time) {
|
||||
if h.logs == nil {
|
||||
return
|
||||
}
|
||||
errMsg := ""
|
||||
if err != nil {
|
||||
errMsg = walletErrText(err)
|
||||
}
|
||||
h.logs.Record(model.SystemLog{
|
||||
Username: username,
|
||||
Method: http.MethodPost,
|
||||
Path: requestPath(c),
|
||||
Status: status,
|
||||
DurationMs: time.Since(start).Milliseconds(),
|
||||
ClientIP: requestIP(c),
|
||||
UserAgent: truncateLogField(c.Request.UserAgent(), 256),
|
||||
ErrMsg: errMsg,
|
||||
})
|
||||
}
|
||||
@@ -64,7 +64,7 @@ func buildDialector(driver, dsn, path string) (gorm.Dialector, error) {
|
||||
|
||||
func autoMigrate(db *gorm.DB) error {
|
||||
return db.AutoMigrate(
|
||||
&model.User{}, &model.UserIdentity{}, &model.OciConfig{}, &model.Task{}, &model.TaskLog{},
|
||||
&model.User{}, &model.UserIdentity{}, &model.UserPasskey{}, &model.UserSession{}, &model.OciConfig{}, &model.Task{}, &model.TaskLog{},
|
||||
&model.CheckSnapshot{}, &model.CostSnapshot{},
|
||||
&model.RegionCache{}, &model.CompartmentCache{}, &model.Setting{},
|
||||
&model.SystemLog{}, &model.LogEvent{}, &model.Proxy{},
|
||||
|
||||
@@ -86,6 +86,45 @@ type UserIdentity struct {
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
}
|
||||
|
||||
// UserPasskey 是账号绑定的 WebAuthn 通行密钥;公钥非敏感,凭据整段 JSON 明文落库。
|
||||
type UserPasskey struct {
|
||||
ID uint `gorm:"primaryKey" json:"id"`
|
||||
UserID uint `gorm:"index" json:"-"`
|
||||
Name string `gorm:"size:64" json:"name"`
|
||||
// 凭据 ID(base64url)原文:规范允许原始 ID 达 1023 字节、编码后逾千字符,
|
||||
// 存 text 不设索引;唯一约束与登录反查走下方定长哈希列
|
||||
CredentialID string `gorm:"type:text" json:"-"`
|
||||
// 凭据 ID 的 SHA-256 hex(定长 64):唯一索引防同一验证器重复注册
|
||||
CredentialIDHash string `gorm:"size:64;uniqueIndex" json:"-"`
|
||||
// 注册时的 WebAuthn origin(scheme://host):面板地址变更后按它判定
|
||||
// 该凭据是否仍可登录(RP ID / origin 绑定注册时域名)
|
||||
Origin string `gorm:"size:255" json:"-"`
|
||||
// webauthn.Credential 完整 JSON(公钥/signCount/flags);attestation 可上 KB,
|
||||
// 显式 text 防 MySQL DefaultStringSize 截断;校验后 signCount 有更新,整段重写
|
||||
Credential string `gorm:"type:text" json:"-"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
LastUsedAt *time.Time `json:"lastUsedAt"`
|
||||
}
|
||||
|
||||
// UserSession 是一次已签发登录会话(JWT)的落地记录;
|
||||
// 「活跃」判定 = 未撤销 && 未过期 && token_ver == 账号当前版本。
|
||||
// 升级前签发的存量令牌无对应行,校验时放行(兼容),仅不可见/不可定点撤销。
|
||||
type UserSession struct {
|
||||
ID uint `gorm:"primaryKey" json:"id"`
|
||||
UserID uint `gorm:"index" json:"-"`
|
||||
// JWT jti;敏感变更换发时同行更新(会话接续),不产生新条目
|
||||
TokenID string `gorm:"size:64;uniqueIndex" json:"-"`
|
||||
TokenVer uint `json:"-"`
|
||||
// 登录方式:password / oidc / github / passkey / wallet;接续换发保留原值
|
||||
Method string `gorm:"size:16" json:"method"`
|
||||
ClientIP string `gorm:"size:64" json:"clientIp"`
|
||||
UserAgent string `gorm:"size:256" json:"userAgent"`
|
||||
CreatedAt time.Time `json:"createdAt"`
|
||||
LastSeenAt time.Time `json:"lastSeenAt"`
|
||||
ExpiresAt time.Time `json:"expiresAt"`
|
||||
RevokedAt *time.Time `json:"-"`
|
||||
}
|
||||
|
||||
// 任务类型取值。
|
||||
const (
|
||||
TaskTypeHealthCheck = "health_check" // 定时测活
|
||||
|
||||
+97
-32
@@ -39,13 +39,19 @@ type authClaims struct {
|
||||
// 用户不存在分支比对它以对齐耗时,防用户枚举与时序侧信道。
|
||||
const dummyBcryptHash = "$2a$10$N9qo8uLOickgx2ZMRZoMye3xW1Wq8p1zEIfQpXCXbXyE3xY5C6P6W"
|
||||
|
||||
// AuthService 负责账号初始化、登录校验(密码 + 可选 TOTP)和 JWT 签发与验证。
|
||||
// AuthService 负责账号初始化、登录校验(密码 + 可选 TOTP)、JWT 签发与验证,
|
||||
// 以及会话落地与管理(session.go)。
|
||||
type AuthService struct {
|
||||
db *gorm.DB
|
||||
jwtSecret []byte
|
||||
guard *loginGuard
|
||||
// revoked 是登出令牌黑名单(哈希→占位),TTL 对齐令牌剩余有效期,过期自动清出
|
||||
revoked *cache.Cache
|
||||
// seen 是会话最后活跃回写的节流缓存(jti→占位);撤销时清键保证即时生效
|
||||
seen *cache.Cache
|
||||
// revokedJti 是定点撤销的负缓存:无容量上限的专用结构,
|
||||
// 既不与高频 seen 回写争容量,也不存在满载淘汰导致的令牌复活
|
||||
revokedJti *jtiTombstones
|
||||
|
||||
notifier *Notifier
|
||||
settings *SettingService
|
||||
@@ -53,6 +59,9 @@ type AuthService struct {
|
||||
|
||||
totpMu sync.Mutex
|
||||
totpPending map[string]pendingTotp // username → setup 暂存密钥
|
||||
|
||||
// cleanupWG 追踪会话清理 goroutine;关停时 Wait 保证清理查询已结束
|
||||
cleanupWG sync.WaitGroup
|
||||
}
|
||||
|
||||
// NewAuthService 组装依赖。
|
||||
@@ -61,6 +70,8 @@ func NewAuthService(db *gorm.DB, jwtSecret string) *AuthService {
|
||||
db: db, jwtSecret: []byte(jwtSecret),
|
||||
guard: newLoginGuard(),
|
||||
revoked: cache.New(4096),
|
||||
seen: cache.New(4096),
|
||||
revokedJti: newJtiTombstones(),
|
||||
totpPending: map[string]pendingTotp{},
|
||||
}
|
||||
}
|
||||
@@ -109,11 +120,12 @@ func (s *AuthService) createUser(username, password string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Login 校验用户名密码与可选 TOTP,成功后签发 JWT;按「IP+用户名」滑动窗口防爆破,
|
||||
// 锁定期内一律 ErrLoginLocked(正确密码同样拒绝);阈值与时长取安全设置。
|
||||
// 已启用两步验证时:密码通过但缺验证码返回 ErrTotpRequired(不计失败),验证码错误计入守卫。
|
||||
func (s *AuthService) Login(ctx context.Context, username, password, clientIP, totpCode string) (string, time.Time, error) {
|
||||
key := guardKey(clientIP, username)
|
||||
// Login 校验用户名密码与可选 TOTP,成功后签发 JWT 并落地会话;
|
||||
// 按「IP+用户名」滑动窗口防爆破,锁定期内一律 ErrLoginLocked(正确密码同样拒绝);
|
||||
// 阈值与时长取安全设置。已启用两步验证时:密码通过但缺验证码返回
|
||||
// ErrTotpRequired(不计失败),验证码错误计入守卫。
|
||||
func (s *AuthService) Login(ctx context.Context, username, password, totpCode string, meta SessionMeta) (string, time.Time, error) {
|
||||
key := guardKey(meta.ClientIP, username)
|
||||
now := time.Now()
|
||||
sec := securityOf(s.settings)
|
||||
lockFor := time.Duration(sec.LoginLockMinutes) * time.Minute
|
||||
@@ -129,23 +141,24 @@ func (s *AuthService) Login(ctx context.Context, username, password, clientIP, t
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
// 与密码错误分支对齐耗时,防用户枚举
|
||||
_ = bcrypt.CompareHashAndPassword([]byte(dummyBcryptHash), []byte(password))
|
||||
return "", time.Time{}, s.failLogin(key, now, username, clientIP, sec)
|
||||
return "", time.Time{}, s.failLogin(key, now, username, meta.ClientIP, sec)
|
||||
}
|
||||
return "", time.Time{}, fmt.Errorf("find user: %w", err)
|
||||
}
|
||||
if bcrypt.CompareHashAndPassword([]byte(user.PasswordHash), []byte(password)) != nil {
|
||||
return "", time.Time{}, s.failLogin(key, now, username, clientIP, sec)
|
||||
return "", time.Time{}, s.failLogin(key, now, username, meta.ClientIP, sec)
|
||||
}
|
||||
if user.TotpSecretEnc != "" {
|
||||
if totpCode == "" {
|
||||
return "", time.Time{}, ErrTotpRequired
|
||||
}
|
||||
if !s.verifyTotp(&user, totpCode) {
|
||||
return "", time.Time{}, s.failLogin(key, now, username, clientIP, sec)
|
||||
return "", time.Time{}, s.failLogin(key, now, username, meta.ClientIP, sec)
|
||||
}
|
||||
}
|
||||
s.guard.success(key)
|
||||
return s.signToken(user.Username, user.TokenVersion)
|
||||
meta.Method = "password"
|
||||
return s.signSessionToken(ctx, &user, meta)
|
||||
}
|
||||
|
||||
// failLogin 记失败;达到阈值转锁定并推送告警(开关 login_lock,缺省开)。
|
||||
@@ -176,7 +189,13 @@ func (s *AuthService) notifyLock(username, clientIP string, sec SecuritySettings
|
||||
})
|
||||
}
|
||||
|
||||
func (s *AuthService) signToken(username string, ver uint) (string, time.Time, error) {
|
||||
// signToken 签发 JWT,返回令牌、过期时间与 jti(会话落地用)。
|
||||
func (s *AuthService) signToken(username string, ver uint) (string, time.Time, string, error) {
|
||||
return s.signTokenWithJTI(username, ver, newTokenID())
|
||||
}
|
||||
|
||||
// signTokenWithJTI 为敏感换发沿用会话 jti,使旧/新 JWT 始终指向同一会话。
|
||||
func (s *AuthService) signTokenWithJTI(username string, ver uint, jti string) (string, time.Time, string, error) {
|
||||
now := time.Now()
|
||||
expires := now.Add(tokenTTL)
|
||||
claims := authClaims{
|
||||
@@ -184,31 +203,37 @@ func (s *AuthService) signToken(username string, ver uint) (string, time.Time, e
|
||||
Subject: username,
|
||||
IssuedAt: jwt.NewNumericDate(now),
|
||||
ExpiresAt: jwt.NewNumericDate(expires),
|
||||
// jti:同一秒签发的令牌若无唯一 ID 字节全同,登出一个会连坐全部
|
||||
ID: newTokenID(),
|
||||
ID: jti,
|
||||
},
|
||||
Ver: ver,
|
||||
}
|
||||
token, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(s.jwtSecret)
|
||||
if err != nil {
|
||||
return "", time.Time{}, fmt.Errorf("sign token: %w", err)
|
||||
return "", time.Time{}, "", fmt.Errorf("sign token: %w", err)
|
||||
}
|
||||
return token, expires, nil
|
||||
return token, expires, jti, nil
|
||||
}
|
||||
|
||||
// IssueToken 按账号当前令牌版本签发新 JWT;敏感操作递增版本后用它为
|
||||
// 操作者重签,避免操作者自身会话中断。
|
||||
// IssueToken 按账号当前令牌版本签发新 JWT(不落会话行);
|
||||
// 需要会话接续的换发场景用 RenewToken。
|
||||
func (s *AuthService) IssueToken(ctx context.Context, username string) (string, time.Time, error) {
|
||||
user, err := s.findUser(ctx, username)
|
||||
if err != nil {
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
return s.signToken(user.Username, user.TokenVersion)
|
||||
token, expires, _, err := s.signToken(user.Username, user.TokenVersion)
|
||||
return token, expires, err
|
||||
}
|
||||
|
||||
// bumpTokenVersion 原子递增账号令牌版本,使所有已签发令牌立即失效。
|
||||
func (s *AuthService) bumpTokenVersion(ctx context.Context, username string) error {
|
||||
err := s.db.WithContext(ctx).Model(&model.User{}).Where("username = ?", username).
|
||||
return bumpTokenVersionTx(s.db.WithContext(ctx), username)
|
||||
}
|
||||
|
||||
// bumpTokenVersionTx 事务内递增令牌版本:与认证因子写入同事务提交,
|
||||
// 避免「因子已生效而旧令牌仍有效」的半程状态。
|
||||
func bumpTokenVersionTx(tx *gorm.DB, username string) error {
|
||||
err := tx.Model(&model.User{}).Where("username = ?", username).
|
||||
UpdateColumn("token_version", gorm.Expr("token_version + 1")).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("bump token version: %w", err)
|
||||
@@ -216,12 +241,28 @@ func (s *AuthService) bumpTokenVersion(ctx context.Context, username string) err
|
||||
return nil
|
||||
}
|
||||
|
||||
// RevokeSessions 撤销账号全部会话(版本递增),并为操作者重签新令牌。
|
||||
func (s *AuthService) RevokeSessions(ctx context.Context, username string) (string, time.Time, error) {
|
||||
if err := s.bumpTokenVersion(ctx, username); err != nil {
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
return s.IssueToken(ctx, username)
|
||||
// RevokeSessions 撤销账号全部会话(版本递增),并为操作者重签新令牌;
|
||||
// oldToken 非空时其会话行接续到新令牌(当前设备在列表中保持连续)。
|
||||
// 行锁下复核 proof:请求挂起期间令牌已失效则拒绝,并发敏感操作后到者拒。
|
||||
func (s *AuthService) RevokeSessions(ctx context.Context, username, oldToken string, meta SessionMeta, proof TokenProof) (string, time.Time, error) {
|
||||
var token string
|
||||
var expires time.Time
|
||||
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.ensureTokenCurrentTx(tx, user, proof); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := bumpTokenVersionTx(tx, username); err != nil {
|
||||
return err
|
||||
}
|
||||
user.TokenVersion++
|
||||
token, expires, err = s.renewSessionTx(tx, user, oldToken, meta)
|
||||
return err
|
||||
})
|
||||
return token, expires, err
|
||||
}
|
||||
|
||||
// newTokenID 生成 128 位随机令牌 ID(crypto/rand 自 Go 1.24 起不会失败)。
|
||||
@@ -231,9 +272,28 @@ func newTokenID() string {
|
||||
return hex.EncodeToString(b)
|
||||
}
|
||||
|
||||
// TokenProof 是鉴权时观察到的令牌快照(版本 + jti):敏感事务在用户行锁下
|
||||
// 复核它仍是账号当前状态,防「请求挂起期间用户撤销 / 注销,恢复后仍完成
|
||||
// 敏感变更并换出新令牌」的在途绕过。
|
||||
type TokenProof struct {
|
||||
Ver uint
|
||||
Jti string
|
||||
}
|
||||
|
||||
// ErrTokenStale 表示请求令牌在处理期间已失效(版本已递增或会话已撤销);
|
||||
// api 层映射 401,前端按「发送时令牌 == 当前令牌」决定是否登出。
|
||||
var ErrTokenStale = errors.New("会话已失效,请重新登录")
|
||||
|
||||
// ParseToken 验证 JWT 签名、有效期与令牌版本,返回其中的用户名。
|
||||
// 版本落后于账号当前值(凭据等已变更)按无效处理,不区分具体原因。
|
||||
func (s *AuthService) ParseToken(ctx context.Context, tokenString string) (string, error) {
|
||||
username, _, err := s.ParseTokenProof(ctx, tokenString)
|
||||
return username, err
|
||||
}
|
||||
|
||||
// ParseTokenProof 验证令牌并返回用户名与版本 / jti 快照(RequireAuth 与
|
||||
// 绑定发起点用;快照随后交由敏感事务复核)。
|
||||
func (s *AuthService) ParseTokenProof(ctx context.Context, tokenString string) (string, TokenProof, error) {
|
||||
claims := &authClaims{}
|
||||
_, err := jwt.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (any, error) {
|
||||
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
@@ -242,23 +302,27 @@ func (s *AuthService) ParseToken(ctx context.Context, tokenString string) (strin
|
||||
return s.jwtSecret, nil
|
||||
})
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("parse token: %w", err)
|
||||
return "", TokenProof{}, fmt.Errorf("parse token: %w", err)
|
||||
}
|
||||
if _, hit := s.revoked.Get(tokenHash(tokenString)); hit {
|
||||
return "", errors.New("token revoked")
|
||||
return "", TokenProof{}, errors.New("token revoked")
|
||||
}
|
||||
user, err := s.findUser(ctx, claims.Subject)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("token subject: %w", err)
|
||||
return "", TokenProof{}, fmt.Errorf("token subject: %w", err)
|
||||
}
|
||||
if claims.Ver != user.TokenVersion {
|
||||
return "", errors.New("token version outdated")
|
||||
return "", TokenProof{}, errors.New("token version outdated")
|
||||
}
|
||||
return claims.Subject, nil
|
||||
if err := s.checkSession(ctx, claims.ID); err != nil {
|
||||
return "", TokenProof{}, err
|
||||
}
|
||||
return claims.Subject, TokenProof{Ver: claims.Ver, Jti: claims.ID}, nil
|
||||
}
|
||||
|
||||
// Logout 把令牌拉进黑名单直至其自然过期;无效/已过期令牌直接视为成功(幂等)。
|
||||
func (s *AuthService) Logout(tokenString string) {
|
||||
// Logout 把令牌拉进黑名单直至其自然过期,并标记对应会话行已撤销;
|
||||
// 无效/已过期令牌直接视为成功(幂等)。
|
||||
func (s *AuthService) Logout(ctx context.Context, tokenString string) {
|
||||
claims := &jwt.RegisteredClaims{}
|
||||
_, err := jwt.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (any, error) {
|
||||
return s.jwtSecret, nil
|
||||
@@ -271,6 +335,7 @@ func (s *AuthService) Logout(tokenString string) {
|
||||
return
|
||||
}
|
||||
s.revoked.Set(tokenHash(tokenString), struct{}{}, ttl)
|
||||
s.revokeSessionByJTI(ctx, claims.Subject, claims.ID, ttl)
|
||||
}
|
||||
|
||||
// tokenHash 取令牌 SHA-256 摘要作黑名单键,不在内存长期保留原令牌串。
|
||||
|
||||
@@ -21,7 +21,7 @@ func newTestAuth(t *testing.T) *AuthService {
|
||||
if err != nil {
|
||||
t.Fatalf("open in-memory sqlite: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.User{}); err != nil {
|
||||
if err := db.AutoMigrate(&model.User{}, &model.UserSession{}); err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
return NewAuthService(db, "test-jwt-secret")
|
||||
@@ -32,7 +32,7 @@ func TestEnsureAdminCreatesUser(t *testing.T) {
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
if _, _, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", ""); err != nil {
|
||||
if _, _, err := auth.Login(context.Background(), "admin", "pass123", "", SessionMeta{ClientIP: "127.0.0.1"}); err != nil {
|
||||
t.Errorf("Login after EnsureAdmin: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -45,10 +45,10 @@ func TestEnsureAdminDoesNotResetPassword(t *testing.T) {
|
||||
if err := auth.EnsureAdmin("admin", "second"); err != nil {
|
||||
t.Fatalf("EnsureAdmin twice: %v", err)
|
||||
}
|
||||
if _, _, err := auth.Login(context.Background(), "admin", "first", "127.0.0.1", ""); err != nil {
|
||||
if _, _, err := auth.Login(context.Background(), "admin", "first", "", SessionMeta{ClientIP: "127.0.0.1"}); err != nil {
|
||||
t.Errorf("Login with original password: %v", err)
|
||||
}
|
||||
if _, _, err := auth.Login(context.Background(), "admin", "second", "127.0.0.1", ""); !errors.Is(err, ErrInvalidCredentials) {
|
||||
if _, _, err := auth.Login(context.Background(), "admin", "second", "", SessionMeta{ClientIP: "127.0.0.1"}); !errors.Is(err, ErrInvalidCredentials) {
|
||||
t.Errorf("Login with new password: got %v, want ErrInvalidCredentials", err)
|
||||
}
|
||||
}
|
||||
@@ -89,7 +89,7 @@ func TestLoginFailures(t *testing.T) {
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, _, err := auth.Login(context.Background(), tt.username, tt.password, "127.0.0.1", "")
|
||||
_, _, err := auth.Login(context.Background(), tt.username, tt.password, "", SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if !errors.Is(err, ErrInvalidCredentials) {
|
||||
t.Errorf("Login(%q, %q) error = %v, want ErrInvalidCredentials", tt.username, tt.password, err)
|
||||
}
|
||||
@@ -102,7 +102,7 @@ func TestTokenRoundTrip(t *testing.T) {
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
token, expires, err := auth.Login(context.Background(), "admin", "pass123", "127.0.0.1", "")
|
||||
token, expires, err := auth.Login(context.Background(), "admin", "pass123", "", SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("Login: %v", err)
|
||||
}
|
||||
@@ -125,7 +125,7 @@ func TestParseTokenRejectsForged(t *testing.T) {
|
||||
if err := other.EnsureAdmin("admin", "pass"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
forged, _, err := other.Login(context.Background(), "admin", "pass", "127.0.0.1", "")
|
||||
forged, _, err := other.Login(context.Background(), "admin", "pass", "", SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("Login: %v", err)
|
||||
}
|
||||
@@ -143,21 +143,21 @@ func TestLogoutRevokesToken(t *testing.T) {
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
token, _, err := auth.signToken("admin", 0)
|
||||
token, _, _, err := auth.signToken("admin", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("signToken: %v", err)
|
||||
}
|
||||
if _, err := auth.ParseToken(context.Background(), token); err != nil {
|
||||
t.Fatalf("ParseToken before logout: %v", err)
|
||||
}
|
||||
auth.Logout(token)
|
||||
auth.Logout(context.Background(), token)
|
||||
if _, err := auth.ParseToken(context.Background(), token); err == nil {
|
||||
t.Error("ParseToken after logout: got nil error, want revoked")
|
||||
}
|
||||
// 幂等:重复登出与无效令牌登出都不应 panic,也不影响其他令牌
|
||||
auth.Logout(token)
|
||||
auth.Logout("not.a.token")
|
||||
fresh, _, err := auth.signToken("admin", 0)
|
||||
auth.Logout(context.Background(), token)
|
||||
auth.Logout(context.Background(), "not.a.token")
|
||||
fresh, _, _, err := auth.signToken("admin", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("signToken fresh: %v", err)
|
||||
}
|
||||
@@ -174,13 +174,13 @@ func TestTokenVersionInvalidatesOldToken(t *testing.T) {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
old, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", "")
|
||||
old, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("Login: %v", err)
|
||||
}
|
||||
finalName, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{
|
||||
NewPassword: "changed-456", CurrentPassword: "pass123",
|
||||
})
|
||||
}, proofOf(t, auth.db, "admin"))
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateCredentials: %v", err)
|
||||
}
|
||||
@@ -203,11 +203,11 @@ func TestRevokeSessions(t *testing.T) {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
old, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", "")
|
||||
old, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("Login: %v", err)
|
||||
}
|
||||
fresh, _, err := auth.RevokeSessions(ctx, "admin")
|
||||
fresh, _, err := auth.RevokeSessions(ctx, "admin", "", SessionMeta{}, proofOf(t, auth.db, "admin"))
|
||||
if err != nil {
|
||||
t.Fatalf("RevokeSessions: %v", err)
|
||||
}
|
||||
|
||||
+234
-24
@@ -20,13 +20,16 @@ var ErrCredentialConfirm = errors.New("当前密码不正确")
|
||||
var ErrCredentialInvalid = errors.New("凭据输入非法")
|
||||
|
||||
// ErrPasswordLoginDisabled 表示密码登录已被禁用;api 层映射 403。
|
||||
var ErrPasswordLoginDisabled = errors.New("密码登录已禁用,请使用外部身份登录")
|
||||
var ErrPasswordLoginDisabled = errors.New("密码登录已禁用,请使用免密方式登录")
|
||||
|
||||
// ErrNeedIdentity 表示未绑定外部身份不可禁用密码登录;api 层映射 409。
|
||||
var ErrNeedIdentity = errors.New("至少绑定一个外部身份后才能禁用密码登录")
|
||||
// ErrNeedIdentity 表示无任何可用免密登录方式时不可禁用密码登录;api 层映射 409。
|
||||
var ErrNeedIdentity = errors.New("需先有可用的免密登录方式(通行密钥、钱包或已启用的外部登录),才能禁用密码登录")
|
||||
|
||||
// ErrLastIdentity 表示密码登录禁用期间不可解绑最后一个身份(防自锁);api 层映射 409。
|
||||
var ErrLastIdentity = errors.New("密码登录已禁用,不能解绑最后一个外部身份;请先允许密码登录")
|
||||
// ErrProviderLastLogin 表示密码登录禁用期间不可禁用/清空最后可用的登录方式;api 层映射 409。
|
||||
var ErrProviderLastLogin = errors.New("密码登录已禁用,该操作将移除最后可用的登录方式;请先允许密码登录")
|
||||
|
||||
// ErrLastIdentity 表示密码登录禁用期间不可移除最后一个免密登录方式(防自锁);api 层映射 409。
|
||||
var ErrLastIdentity = errors.New("密码登录已禁用,不能移除最后一个免密登录方式;请先允许密码登录")
|
||||
|
||||
// UpdateCredentialsInput 是修改登录凭据的请求体;两项至少改一项,当前密码必填。
|
||||
type UpdateCredentialsInput struct {
|
||||
@@ -35,9 +38,23 @@ type UpdateCredentialsInput struct {
|
||||
CurrentPassword string `json:"currentPassword" binding:"required"`
|
||||
}
|
||||
|
||||
type authenticatedMutation struct {
|
||||
auth *AuthService
|
||||
username string
|
||||
proof TokenProof
|
||||
}
|
||||
|
||||
func (m *authenticatedMutation) lockAndCheck(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, m.username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return m.auth.ensureTokenCurrentTx(tx, user, m.proof)
|
||||
}
|
||||
|
||||
// UpdateCredentials 修改用户名 / 密码:当前密码必验;成功后令牌版本递增
|
||||
// (全部旧 JWT 立即失效),返回最终用户名供调用方为操作者重签新令牌。
|
||||
func (s *AuthService) UpdateCredentials(ctx context.Context, username string, in UpdateCredentialsInput) (string, error) {
|
||||
func (s *AuthService) UpdateCredentials(ctx context.Context, username string, in UpdateCredentialsInput, proof TokenProof) (string, error) {
|
||||
user, err := s.findUser(ctx, username)
|
||||
if err != nil {
|
||||
return "", err
|
||||
@@ -53,14 +70,40 @@ func (s *AuthService) UpdateCredentials(ctx context.Context, username string, in
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
// 同一条 UPDATE 里递增令牌版本,与凭据变更保持原子
|
||||
updates["token_version"] = gorm.Expr("token_version + 1")
|
||||
if err := s.db.WithContext(ctx).Model(&model.User{}).Where("id = ?", user.ID).Updates(updates).Error; err != nil {
|
||||
return "", fmt.Errorf("update credentials: %w", err)
|
||||
if err := s.applyCredentialUpdates(ctx, username, proof, updates); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return finalName, nil
|
||||
}
|
||||
|
||||
func (s *AuthService) applyCredentialUpdates(
|
||||
ctx context.Context, username string, proof TokenProof, updates map[string]any,
|
||||
) error {
|
||||
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
locked, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.ensureTokenCurrentTx(tx, locked, proof); err != nil {
|
||||
return err
|
||||
}
|
||||
return updateCredentialsTx(tx, locked.ID, proof.Ver, updates)
|
||||
})
|
||||
}
|
||||
|
||||
func updateCredentialsTx(tx *gorm.DB, userID uint, version uint, updates map[string]any) error {
|
||||
res := tx.Model(&model.User{}).
|
||||
Where("id = ? AND token_version = ?", userID, version).Updates(updates)
|
||||
if res.Error != nil {
|
||||
return fmt.Errorf("update credentials: %w", res.Error)
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
return ErrTokenStale
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// credentialUpdates 组装凭据变更字段并返回最终用户名。
|
||||
func (s *AuthService) credentialUpdates(ctx context.Context, user *model.User, newName, newPassword string) (map[string]any, string, error) {
|
||||
updates := map[string]any{}
|
||||
@@ -119,9 +162,10 @@ func (s *AuthService) PasswordLoginDisabled(ctx context.Context) (bool, error) {
|
||||
return s.settings.PasswordLoginDisabled(ctx)
|
||||
}
|
||||
|
||||
// SetPasswordLoginDisabled 保存开关;开启前须至少绑定一个外部身份,防止自锁。
|
||||
// 检查与写入在同一事务内并锁定用户行,防与解绑身份并发绕过「至少一种登录方式」。
|
||||
func (s *AuthService) SetPasswordLoginDisabled(ctx context.Context, username string, disabled bool) error {
|
||||
// SetPasswordLoginDisabled 保存开关;开启前须至少有一种免密登录方式
|
||||
// (通行密钥或外部身份),防止自锁。检查与写入在同一事务内并锁定用户行,
|
||||
// 防与解绑身份/删除通行密钥并发绕过「至少一种登录方式」。
|
||||
func (s *AuthService) SetPasswordLoginDisabled(ctx context.Context, username string, disabled bool, proof TokenProof) error {
|
||||
if s.settings == nil {
|
||||
return errors.New("settings unavailable")
|
||||
}
|
||||
@@ -129,27 +173,41 @@ func (s *AuthService) SetPasswordLoginDisabled(ctx context.Context, username str
|
||||
if disabled {
|
||||
value = "1"
|
||||
}
|
||||
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if disabled {
|
||||
n, err := identityCountTx(tx, user.ID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n == 0 {
|
||||
return ErrNeedIdentity
|
||||
}
|
||||
if err := s.ensureTokenCurrentTx(tx, user, proof); err != nil {
|
||||
return err
|
||||
}
|
||||
return saveSettingTx(tx, settingSecPasswordLoginOff, value)
|
||||
if err := s.ensurePasswordlessForToggle(tx, user.ID, disabled); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := saveSettingTx(tx, settingSecPasswordLoginOff, value); err != nil {
|
||||
return err
|
||||
}
|
||||
// 登录策略属敏感变更:版本递增与开关写入同事务提交,不留半程状态
|
||||
return bumpTokenVersionTx(tx, username)
|
||||
})
|
||||
}
|
||||
|
||||
func (s *AuthService) ensurePasswordlessForToggle(tx *gorm.DB, userID uint, disabled bool) error {
|
||||
if !disabled {
|
||||
return nil
|
||||
}
|
||||
origin, err := effectiveOriginTx(tx, s.settings)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// 登录策略属敏感变更:递增令牌版本,已签发会话全部失效
|
||||
return s.bumpTokenVersion(ctx, username)
|
||||
ok, err := usablePasswordlessTx(tx, userID, 0, 0, origin)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
return ErrNeedIdentity
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// lockUserForAuthChange 事务内锁定用户行(SQLite 单写天然串行,MySQL/PG 靠行锁),
|
||||
@@ -164,6 +222,64 @@ func lockUserForAuthChange(tx *gorm.DB, username string) (*model.User, error) {
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
func lockUserByIDForAuthChange(tx *gorm.DB, userID uint) (*model.User, error) {
|
||||
var user model.User
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).First(&user, userID).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("find user: %w", err)
|
||||
}
|
||||
return &user, nil
|
||||
}
|
||||
|
||||
// ensureTokenCurrentTx 行锁下复核请求令牌仍是账号当前状态:版本一致
|
||||
// (改密 / 撤销全部会递增),且对应会话行未被注销或定点撤销
|
||||
// (Logout 联动标记 revoked_at,查行即可覆盖两者);存量令牌无行时仅校验版本。
|
||||
func (s *AuthService) ensureTokenCurrentTx(tx *gorm.DB, user *model.User, proof TokenProof) error {
|
||||
if user.TokenVersion != proof.Ver {
|
||||
return ErrTokenStale
|
||||
}
|
||||
if proof.Jti == "" {
|
||||
return nil
|
||||
}
|
||||
if s.revokedJti.has(proof.Jti) {
|
||||
return ErrTokenStale
|
||||
}
|
||||
var row model.UserSession
|
||||
err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).
|
||||
Select("id", "revoked_at").Where("token_id = ?", proof.Jti).First(&row).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("find session: %w", err)
|
||||
}
|
||||
if row.RevokedAt != nil {
|
||||
return ErrTokenStale
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// lockUsersForAuthChange 锁定全部用户行(单管理员面板即一行):与
|
||||
// lockUserForAuthChange 竞争同一把行锁,防「改 provider / 面板地址配置」
|
||||
// 与「禁用密码 / 删除最后因子」并发交错绕过登录方式不变量。
|
||||
func lockUsersForAuthChange(tx *gorm.DB) error {
|
||||
var users []model.User
|
||||
if err := tx.Clauses(clause.Locking{Strength: "UPDATE"}).Find(&users).Error; err != nil {
|
||||
return fmt.Errorf("lock users: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// identityProviderCountTx 统计某 provider 的身份总数(单管理员面板,不区分账号)。
|
||||
func identityProviderCountTx(tx *gorm.DB, provider string) (int64, error) {
|
||||
var n int64
|
||||
err := tx.Model(&model.UserIdentity{}).Where("provider = ?", provider).Count(&n).Error
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("count identities: %w", err)
|
||||
}
|
||||
return n, nil
|
||||
}
|
||||
|
||||
func identityCountTx(tx *gorm.DB, userID uint) (int64, error) {
|
||||
var count int64
|
||||
err := tx.Model(&model.UserIdentity{}).
|
||||
@@ -173,3 +289,97 @@ func identityCountTx(tx *gorm.DB, userID uint) (int64, error) {
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// passkeyCountTx 统计账号的通行密钥数(事务内,防自锁检查用)。
|
||||
func passkeyCountTx(tx *gorm.DB, userID uint) (int64, error) {
|
||||
var count int64
|
||||
err := tx.Model(&model.UserPasskey{}).
|
||||
Where("user_id = ?", userID).Count(&count).Error
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("count passkeys: %w", err)
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// usablePasswordlessTx 事务内判断(排除给定身份/通行密钥行后)是否仍存在
|
||||
// 可实际登录的免密方式:通行密钥仅统计注册 origin 与当前面板地址一致的
|
||||
// (地址变更后旧域名凭据不可登录,不得计入);钱包按地址绑定,恒可用;
|
||||
// GitHub/OIDC 身份须对应 provider 已配置且未禁用才计入(端到端防自锁)。
|
||||
// origin 为空时所有方式均不可用:Passkey RP、钱包 SIWE 与 OAuth 回调都依赖面板地址。
|
||||
func usablePasswordlessTx(tx *gorm.DB, userID uint, excludeIdentity, excludePasskey uint, origin string) (bool, error) {
|
||||
if origin == "" {
|
||||
return false, nil
|
||||
}
|
||||
pk, err := passkeyCountExcludingTx(tx, userID, excludePasskey, origin)
|
||||
if err != nil || pk > 0 {
|
||||
return pk > 0, err
|
||||
}
|
||||
if n, err := identityCountByProviderTx(tx, userID, "wallet", excludeIdentity); err != nil || n > 0 {
|
||||
return n > 0, err
|
||||
}
|
||||
for _, p := range []string{"github", "oidc"} {
|
||||
n, err := identityCountByProviderTx(tx, userID, p, excludeIdentity)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if n == 0 {
|
||||
continue
|
||||
}
|
||||
if ok, err := oauthProviderUsableTx(tx, p); err != nil || ok {
|
||||
return ok, err
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// oauthProviderUsableTx 事务内判断 provider 当前可实际登录:
|
||||
// clientID 与 secret 均已配置(oidc 还需 issuer)且未禁用。
|
||||
func oauthProviderUsableTx(tx *gorm.DB, provider string) (bool, error) {
|
||||
need := []string{settingOauthGithubClientID, settingOauthGithubClientSecret}
|
||||
offKey := settingOauthGithubDisabled
|
||||
if provider == "oidc" {
|
||||
need = []string{settingOauthOidcClientID, settingOauthOidcClientSecret, settingOauthOidcIssuer}
|
||||
offKey = settingOauthOidcDisabled
|
||||
}
|
||||
for _, k := range need {
|
||||
v, err := settingValueTx(tx, k)
|
||||
if err != nil || v == "" {
|
||||
return false, err
|
||||
}
|
||||
}
|
||||
off, err := settingValueTx(tx, offKey)
|
||||
return off != "1", err
|
||||
}
|
||||
|
||||
// identityCountByProviderTx 统计账号某 provider 的身份数,可排除一行(解绑前判定用)。
|
||||
func identityCountByProviderTx(tx *gorm.DB, userID uint, provider string, excludeID uint) (int64, error) {
|
||||
q := tx.Model(&model.UserIdentity{}).Where("user_id = ? AND provider = ?", userID, provider)
|
||||
if excludeID != 0 {
|
||||
q = q.Where("id <> ?", excludeID)
|
||||
}
|
||||
var count int64
|
||||
if err := q.Count(&count).Error; err != nil {
|
||||
return 0, fmt.Errorf("count identities: %w", err)
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// passkeyCountExcludingTx 统计「当前地址下可用」的通行密钥数,可排除一行;
|
||||
// userID 为 0 表示全表(单管理员面板);origin 非空时仅计注册来源一致的凭据。
|
||||
func passkeyCountExcludingTx(tx *gorm.DB, userID, excludeID uint, origin string) (int64, error) {
|
||||
q := tx.Model(&model.UserPasskey{})
|
||||
if userID != 0 {
|
||||
q = q.Where("user_id = ?", userID)
|
||||
}
|
||||
if excludeID != 0 {
|
||||
q = q.Where("id <> ?", excludeID)
|
||||
}
|
||||
if origin != "" {
|
||||
q = q.Where("origin = ?", origin)
|
||||
}
|
||||
var count int64
|
||||
if err := q.Count(&count).Error; err != nil {
|
||||
return 0, fmt.Errorf("count passkeys: %w", err)
|
||||
}
|
||||
return count, nil
|
||||
}
|
||||
|
||||
@@ -5,6 +5,8 @@ import (
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/crypto"
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
@@ -17,7 +19,11 @@ func newCredEnv(t *testing.T) *AuthService {
|
||||
if err != nil {
|
||||
t.Fatalf("new cipher: %v", err)
|
||||
}
|
||||
auth.SetNotifier(nil, NewSettingService(db, cipher))
|
||||
settings := NewSettingService(db, cipher)
|
||||
if err := settings.UpdateSecurity(context.Background(), SecurityPatch{AppURL: strPtr("https://app.example.com")}); err != nil {
|
||||
t.Fatalf("seed app url: %v", err)
|
||||
}
|
||||
auth.SetNotifier(nil, settings)
|
||||
return auth
|
||||
}
|
||||
|
||||
@@ -26,16 +32,16 @@ func TestUpdateCredentials(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// 当前密码错误
|
||||
_, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "newpass-123", CurrentPassword: "wrong"})
|
||||
_, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "newpass-123", CurrentPassword: "wrong"}, proofOf(t, auth.db, "admin"))
|
||||
if !errors.Is(err, ErrCredentialConfirm) {
|
||||
t.Fatalf("wrong current password: err = %v, want ErrCredentialConfirm", err)
|
||||
}
|
||||
// 短密码 / 无变更均拒绝
|
||||
_, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "short", CurrentPassword: "pass123"})
|
||||
_, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewPassword: "short", CurrentPassword: "pass123"}, proofOf(t, auth.db, "admin"))
|
||||
if !errors.Is(err, ErrCredentialInvalid) {
|
||||
t.Fatalf("short password: err = %v, want ErrCredentialInvalid", err)
|
||||
}
|
||||
_, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewUsername: "admin", CurrentPassword: "pass123"})
|
||||
_, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{NewUsername: "admin", CurrentPassword: "pass123"}, proofOf(t, auth.db, "admin"))
|
||||
if !errors.Is(err, ErrCredentialInvalid) {
|
||||
t.Fatalf("no-op change: err = %v, want ErrCredentialInvalid", err)
|
||||
}
|
||||
@@ -43,24 +49,60 @@ func TestUpdateCredentials(t *testing.T) {
|
||||
// 同时改名改密
|
||||
_, err = auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{
|
||||
NewUsername: "root", NewPassword: "newpass-123", CurrentPassword: "pass123",
|
||||
})
|
||||
}, proofOf(t, auth.db, "admin"))
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateCredentials: %v", err)
|
||||
}
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", ""); !errors.Is(err, ErrInvalidCredentials) {
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "127.0.0.1"}); !errors.Is(err, ErrInvalidCredentials) {
|
||||
t.Errorf("old username login: err = %v, want ErrInvalidCredentials", err)
|
||||
}
|
||||
if _, _, err := auth.Login(ctx, "root", "newpass-123", "127.0.0.2", ""); err != nil {
|
||||
if _, _, err := auth.Login(ctx, "root", "newpass-123", "", SessionMeta{ClientIP: "127.0.0.2"}); err != nil {
|
||||
t.Errorf("new credentials login: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUpdateCredentialsRejectsRevokedJTI(t *testing.T) {
|
||||
auth := newCredEnv(t)
|
||||
ctx := context.Background()
|
||||
stale := loginSession(t, auth, "10.10.0.1", "stale")
|
||||
current := loginSession(t, auth, "10.10.0.2", "current")
|
||||
proof := mustProof(t, auth, stale)
|
||||
for _, item := range mustSessions(t, auth, current) {
|
||||
if !item.Current {
|
||||
if err := auth.RevokeSession(ctx, "admin", current, item.ID); err != nil {
|
||||
t.Fatalf("RevokeSession: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
_, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{
|
||||
NewPassword: "newpass-123", CurrentPassword: "pass123",
|
||||
}, proof)
|
||||
if !errors.Is(err, ErrTokenStale) {
|
||||
t.Fatalf("UpdateCredentials err = %v, want ErrTokenStale", err)
|
||||
}
|
||||
if _, _, err := auth.Login(
|
||||
ctx, "admin", "pass123", "", SessionMeta{ClientIP: "10.10.0.3"},
|
||||
); err != nil {
|
||||
t.Fatalf("old password should remain valid: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// proofOf 取账号当前令牌版本组装快照(jti 空=不校验会话行),敏感调用测试用。
|
||||
func proofOf(t *testing.T, db *gorm.DB, username string) TokenProof {
|
||||
t.Helper()
|
||||
var user model.User
|
||||
if err := db.Where("username = ?", username).First(&user).Error; err != nil {
|
||||
t.Fatalf("find user: %v", err)
|
||||
}
|
||||
return TokenProof{Ver: user.TokenVersion}
|
||||
}
|
||||
|
||||
func TestPasswordLoginToggle(t *testing.T) {
|
||||
auth := newCredEnv(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// 未绑定外部身份不可禁用
|
||||
if err := auth.SetPasswordLoginDisabled(ctx, "admin", true); !errors.Is(err, ErrNeedIdentity) {
|
||||
if err := auth.SetPasswordLoginDisabled(ctx, "admin", true, proofOf(t, auth.db, "admin")); !errors.Is(err, ErrNeedIdentity) {
|
||||
t.Fatalf("disable without identity: err = %v, want ErrNeedIdentity", err)
|
||||
}
|
||||
|
||||
@@ -68,36 +110,73 @@ func TestPasswordLoginToggle(t *testing.T) {
|
||||
if err := auth.db.Where("username = ?", "admin").First(&user).Error; err != nil {
|
||||
t.Fatalf("find admin: %v", err)
|
||||
}
|
||||
ident := model.UserIdentity{UserID: user.ID, Provider: "github", Subject: "1", Display: "tester"}
|
||||
// 用钱包身份:不依赖 provider 配置即可登录(github/oidc 身份须 provider 启用才计入)
|
||||
ident := model.UserIdentity{UserID: user.ID, Provider: "wallet", Subject: "0xabc", Display: "tester"}
|
||||
if err := auth.db.Create(&ident).Error; err != nil {
|
||||
t.Fatalf("seed identity: %v", err)
|
||||
}
|
||||
|
||||
if err := auth.SetPasswordLoginDisabled(ctx, "admin", true); err != nil {
|
||||
if err := auth.SetPasswordLoginDisabled(ctx, "admin", true, proofOf(t, auth.db, "admin")); err != nil {
|
||||
t.Fatalf("disable with identity: %v", err)
|
||||
}
|
||||
if off, err := auth.PasswordLoginDisabled(ctx); err != nil || !off {
|
||||
t.Fatalf("PasswordLoginDisabled = %v, %v; want true", off, err)
|
||||
}
|
||||
// 密码登录被拒且不计失败(正确密码亦拒)
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", ""); !errors.Is(err, ErrPasswordLoginDisabled) {
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "127.0.0.1"}); !errors.Is(err, ErrPasswordLoginDisabled) {
|
||||
t.Fatalf("login while disabled: err = %v, want ErrPasswordLoginDisabled", err)
|
||||
}
|
||||
|
||||
// 开关开着不能解绑最后一个身份
|
||||
oauth := NewOAuthService(auth.db, nil, auth)
|
||||
if err := oauth.Unbind(ctx, "admin", ident.ID); !errors.Is(err, ErrLastIdentity) {
|
||||
if err := oauth.Unbind(ctx, "admin", ident.ID, proofOf(t, auth.db, "admin")); !errors.Is(err, ErrLastIdentity) {
|
||||
t.Fatalf("unbind last identity: err = %v, want ErrLastIdentity", err)
|
||||
}
|
||||
|
||||
// 恢复密码登录后可解绑、可登录
|
||||
if err := auth.SetPasswordLoginDisabled(ctx, "admin", false); err != nil {
|
||||
if err := auth.SetPasswordLoginDisabled(ctx, "admin", false, proofOf(t, auth.db, "admin")); err != nil {
|
||||
t.Fatalf("enable password login: %v", err)
|
||||
}
|
||||
if err := oauth.Unbind(ctx, "admin", ident.ID); err != nil {
|
||||
if err := oauth.Unbind(ctx, "admin", ident.ID, proofOf(t, auth.db, "admin")); err != nil {
|
||||
t.Fatalf("unbind after enable: %v", err)
|
||||
}
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.3", ""); err != nil {
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "127.0.0.3"}); err != nil {
|
||||
t.Errorf("login after enable: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPasswordLoginToggleWithPasskey 验证通行密钥计入禁用密码登录的门槛:
|
||||
// 仅有通行密钥即可开启;禁用期间最后一把钥匙受防自锁保护,身份可兜底解绑。
|
||||
func TestPasswordLoginToggleWithPasskey(t *testing.T) {
|
||||
auth := newCredEnv(t)
|
||||
ctx := context.Background()
|
||||
pk := model.UserPasskey{
|
||||
UserID: 1, Name: "k", CredentialID: "cred-1",
|
||||
CredentialIDHash: "hash-1", Credential: "{}", Origin: "https://app.example.com",
|
||||
}
|
||||
if err := auth.db.Create(&pk).Error; err != nil {
|
||||
t.Fatalf("seed passkey: %v", err)
|
||||
}
|
||||
|
||||
// 无外部身份、仅通行密钥:可禁用密码登录
|
||||
if err := auth.SetPasswordLoginDisabled(ctx, "admin", true, proofOf(t, auth.db, "admin")); err != nil {
|
||||
t.Fatalf("disable with passkey only: %v", err)
|
||||
}
|
||||
// 禁用期间删除最后一把钥匙被拒(防自锁)
|
||||
passkeys := NewPasskeyService(auth.db, auth.settings, auth)
|
||||
if err := passkeys.Remove(ctx, "admin", pk.ID, proofOf(t, auth.db, "admin")); !errors.Is(err, ErrLastIdentity) {
|
||||
t.Fatalf("remove last passkey: err = %v, want ErrLastIdentity", err)
|
||||
}
|
||||
// 绑定身份后钥匙可删;身份成为最后方式后解绑又被拒
|
||||
ident := model.UserIdentity{UserID: 1, Provider: "wallet", Subject: "0xdef", Display: "t"}
|
||||
if err := auth.db.Create(&ident).Error; err != nil {
|
||||
t.Fatalf("seed identity: %v", err)
|
||||
}
|
||||
if err := passkeys.Remove(ctx, "admin", pk.ID, proofOf(t, auth.db, "admin")); err != nil {
|
||||
t.Fatalf("remove passkey with identity fallback: %v", err)
|
||||
}
|
||||
oauth := NewOAuthService(auth.db, nil, auth)
|
||||
if err := oauth.Unbind(ctx, "admin", ident.ID, proofOf(t, auth.db, "admin")); !errors.Is(err, ErrLastIdentity) {
|
||||
t.Fatalf("unbind last identity: err = %v, want ErrLastIdentity", err)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,89 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
|
||||
"golang.org/x/crypto/sha3"
|
||||
)
|
||||
|
||||
// 以太坊 personal_sign(EIP-191)验签工具:仅覆盖 EOA 签名恢复,
|
||||
// 不做 EIP-1271 合约钱包(需链上 RPC,超出自托管零依赖边界)。
|
||||
|
||||
// errEthSig 是签名格式/恢复失败的内部标记;对外统一映射 ErrWalletSig。
|
||||
var errEthSig = errors.New("invalid ethereum signature")
|
||||
|
||||
// keccak256 计算 Keccak-256 摘要(以太坊场景的非 NIST 填充变体)。
|
||||
func keccak256(data ...[]byte) []byte {
|
||||
h := sha3.NewLegacyKeccak256()
|
||||
for _, d := range data {
|
||||
h.Write(d)
|
||||
}
|
||||
return h.Sum(nil)
|
||||
}
|
||||
|
||||
// personalSignDigest 计算 EIP-191 personal_sign 摘要:
|
||||
// keccak256("\x19Ethereum Signed Message:\n" + 消息字节长度十进制 + 消息)。
|
||||
func personalSignDigest(msg []byte) []byte {
|
||||
prefix := []byte("\x19Ethereum Signed Message:\n" + strconv.Itoa(len(msg)))
|
||||
return keccak256(prefix, msg)
|
||||
}
|
||||
|
||||
// recoverEthAddress 从 65 字节签名(r||s||v)恢复签名者地址(EIP-55 格式)。
|
||||
// v 兼容钱包生态的 27/28 与原始 0/1 两种取值;dcrec 的 compact 格式要求
|
||||
// 恢复位打头(27+v),与以太坊的尾置布局相反,需重排。
|
||||
func recoverEthAddress(digest, sig []byte) (string, error) {
|
||||
if len(sig) != 65 {
|
||||
return "", fmt.Errorf("%w: length %d", errEthSig, len(sig))
|
||||
}
|
||||
v := sig[64]
|
||||
if v >= 27 {
|
||||
v -= 27
|
||||
}
|
||||
if v > 1 {
|
||||
return "", fmt.Errorf("%w: recovery id %d", errEthSig, sig[64])
|
||||
}
|
||||
compact := make([]byte, 65)
|
||||
compact[0] = 27 + v
|
||||
copy(compact[1:], sig[:64])
|
||||
pub, _, err := ecdsa.RecoverCompact(compact, digest)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("%w: %v", errEthSig, err)
|
||||
}
|
||||
// 地址 = keccak256(未压缩公钥去掉 0x04 前缀)的后 20 字节
|
||||
return eip55(keccak256(pub.SerializeUncompressed()[1:])[12:]), nil
|
||||
}
|
||||
|
||||
// eip55 把 20 字节地址编码为 EIP-55 混合大小写校验和格式:
|
||||
// 对小写 hex 串取 keccak256,对应半字节 ≥ 8 的字母转大写。
|
||||
func eip55(addr []byte) string {
|
||||
lower := hex.EncodeToString(addr)
|
||||
sum := keccak256([]byte(lower))
|
||||
out := []byte(lower)
|
||||
for i, ch := range out {
|
||||
nibble := sum[i/2] >> 4
|
||||
if i%2 == 1 {
|
||||
nibble = sum[i/2] & 0x0f
|
||||
}
|
||||
if ch >= 'a' && nibble >= 8 {
|
||||
out[i] = ch - 'a' + 'A'
|
||||
}
|
||||
}
|
||||
return "0x" + string(out)
|
||||
}
|
||||
|
||||
// normalizeEthAddress 校验 0x+40hex 地址并返回 EIP-55 规范形;大小写不敏感。
|
||||
func normalizeEthAddress(addr string) (string, error) {
|
||||
if len(addr) != 42 || !strings.HasPrefix(addr, "0x") {
|
||||
return "", ErrWalletAddress
|
||||
}
|
||||
raw, err := hex.DecodeString(strings.ToLower(addr[2:]))
|
||||
if err != nil {
|
||||
return "", ErrWalletAddress
|
||||
}
|
||||
return eip55(raw), nil
|
||||
}
|
||||
@@ -0,0 +1,136 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/decred/dcrd/dcrec/secp256k1/v4"
|
||||
"github.com/decred/dcrd/dcrec/secp256k1/v4/ecdsa"
|
||||
)
|
||||
|
||||
// ethAddrOf 从私钥派生 EIP-55 地址(keccak256(未压缩公钥[1:]) 后 20 字节)。
|
||||
func ethAddrOf(priv *secp256k1.PrivateKey) string {
|
||||
return eip55(keccak256(priv.PubKey().SerializeUncompressed()[1:])[12:])
|
||||
}
|
||||
|
||||
// testPrivKey 构造确定性的测试私钥;seed 填入 32 字节最低位。
|
||||
func testPrivKey(seed byte) *secp256k1.PrivateKey {
|
||||
raw := make([]byte, 32)
|
||||
raw[31] = seed
|
||||
return secp256k1.PrivKeyFromBytes(raw)
|
||||
}
|
||||
|
||||
// signPersonal 用 dcrec 标准路径生成 personal_sign 签名(r||s||v 的 hex);
|
||||
// legacyV 为 true 时 v 取 27/28,否则 0/1——两种取值链上钱包都会出现。
|
||||
func signPersonal(priv *secp256k1.PrivateKey, message string, legacyV bool) string {
|
||||
compact := ecdsa.SignCompact(priv, personalSignDigest([]byte(message)), false)
|
||||
sig := make([]byte, 65)
|
||||
copy(sig, compact[1:])
|
||||
v := compact[0] // 27 或 28(未压缩公钥)
|
||||
if !legacyV {
|
||||
v -= 27
|
||||
}
|
||||
sig[64] = v
|
||||
return "0x" + hex.EncodeToString(sig)
|
||||
}
|
||||
|
||||
// TestEip55Vectors 用 EIP-55 规范文档中的官方校验和向量。
|
||||
func TestEip55Vectors(t *testing.T) {
|
||||
vectors := []string{
|
||||
"0x5aAeb6053F3E94C9b9A09f33669435E7Ef1BeAed",
|
||||
"0xfB6916095ca1df60bB79Ce92cE3Ea74c37c5d359",
|
||||
"0xdbF03B407c01E7cD3CBea99509d93f8DDDC8C6FB",
|
||||
"0xD1220A0cf47c7B9Be7A2E6BA89F429762e7b9aDb",
|
||||
}
|
||||
for _, want := range vectors {
|
||||
t.Run(want[:10], func(t *testing.T) {
|
||||
for _, input := range []string{strings.ToLower(want), "0x" + strings.ToUpper(want[2:])} {
|
||||
got, err := normalizeEthAddress(input)
|
||||
if err != nil {
|
||||
t.Fatalf("normalizeEthAddress(%s): %v", input, err)
|
||||
}
|
||||
if got != want {
|
||||
t.Errorf("normalizeEthAddress(%s) = %s, want %s", input, got, want)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalizeEthAddressInvalid(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
in string
|
||||
}{
|
||||
{name: "缺前缀", in: "5aAeb6053F3E94C9b9A09f33669435E7Ef1BeAed00"},
|
||||
{name: "长度不足", in: "0x5aAeb6"},
|
||||
{name: "非 hex 字符", in: "0x5aAeb6053F3E94C9b9A09f33669435E7Ef1BeAzz"},
|
||||
{name: "空串", in: ""},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if _, err := normalizeEthAddress(tt.in); !errors.Is(err, ErrWalletAddress) {
|
||||
t.Errorf("normalizeEthAddress(%q) err = %v, want ErrWalletAddress", tt.in, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestKnownPrivKeyAddress 用公开锚点(私钥 0x…01 的以太坊地址)校验派生路径,
|
||||
// 防止 keccak / 公钥序列化环节自洽但整体错误。
|
||||
func TestKnownPrivKeyAddress(t *testing.T) {
|
||||
got := ethAddrOf(testPrivKey(1))
|
||||
want := "0x7E5F4552091A69125d5DfCb7b8C2659029395Bdf"
|
||||
if got != want {
|
||||
t.Fatalf("address of privkey 0x01 = %s, want %s", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecoverEthAddressRoundTrip(t *testing.T) {
|
||||
priv := testPrivKey(7)
|
||||
addr := ethAddrOf(priv)
|
||||
const msg = "demo.example.com wants you to sign in with your Ethereum account:\n0xabc"
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
sig func() string
|
||||
wantErr bool
|
||||
}{
|
||||
{name: "v=27/28", sig: func() string { return signPersonal(priv, msg, true) }},
|
||||
{name: "v=0/1", sig: func() string { return signPersonal(priv, msg, false) }},
|
||||
{name: "签名过短", sig: func() string { return "0x0102" }, wantErr: true},
|
||||
{name: "v 非法", sig: func() string {
|
||||
s := signPersonal(priv, msg, true)
|
||||
return s[:len(s)-2] + "63" // v=99
|
||||
}, wantErr: true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
raw, err := hex.DecodeString(strings.TrimPrefix(tt.sig(), "0x"))
|
||||
if err != nil {
|
||||
t.Fatalf("decode sig: %v", err)
|
||||
}
|
||||
got, err := recoverEthAddress(personalSignDigest([]byte(msg)), raw)
|
||||
if tt.wantErr {
|
||||
if err == nil {
|
||||
t.Fatal("recoverEthAddress: expected error, got nil")
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("recoverEthAddress: %v", err)
|
||||
}
|
||||
if got != addr {
|
||||
t.Errorf("recovered = %s, want %s", got, addr)
|
||||
}
|
||||
})
|
||||
}
|
||||
// 篡改消息:恢复出的地址必然不同
|
||||
raw, _ := hex.DecodeString(strings.TrimPrefix(signPersonal(priv, msg, true), "0x"))
|
||||
got, err := recoverEthAddress(personalSignDigest([]byte(msg+"x")), raw)
|
||||
if err == nil && got == addr {
|
||||
t.Error("tampered message still recovered the signer address")
|
||||
}
|
||||
}
|
||||
+172
-56
@@ -41,9 +41,11 @@ const oauthPendingTTL = 10 * time.Minute
|
||||
// oauthPending 是一次进行中的授权流程上下文;state 一次性使用。
|
||||
type oauthPending struct {
|
||||
provider string
|
||||
mode string // "login" / "bind"
|
||||
username string // bind 模式的绑定目标账号
|
||||
nonce string // OIDC 防 id_token 重放
|
||||
mode string // "login" / "bind"
|
||||
username string // bind 模式的绑定目标账号
|
||||
nonce string // OIDC 防 id_token 重放
|
||||
token string // bind 模式发起时的 Bearer;回调复验,防被盗令牌撤销后仍完成绑定
|
||||
proof TokenProof // 发起时的版本/jti 快照;绑定事务行锁下复核,覆盖撤销全部/注销/定点撤销
|
||||
expires time.Time
|
||||
}
|
||||
|
||||
@@ -68,12 +70,17 @@ type ProviderInfo struct {
|
||||
DisplayName string `json:"displayName"`
|
||||
}
|
||||
|
||||
// Providers 返回可登录的 provider 列表(clientID 非空且未禁用),登录页据此显示按钮。
|
||||
// Providers 返回可实际登录的 provider 列表,登录页据此显示按钮;
|
||||
// 与不变量检查同口径:clientID 与 secret 齐备(oidc 还需 issuer)、未禁用,
|
||||
// 且面板地址已设置(回调地址无从拼接时全部不可登录)——半配置不再暴露必败入口。
|
||||
func (o *OAuthService) Providers(ctx context.Context) []ProviderInfo {
|
||||
out := []ProviderInfo{}
|
||||
if o.settings.EffectiveAppURL() == "" {
|
||||
return out
|
||||
}
|
||||
for _, p := range []string{"oidc", "github"} {
|
||||
id, _, _, err := o.settings.oauthClient(ctx, p)
|
||||
if err != nil || id == "" {
|
||||
id, secret, issuer, err := o.settings.oauthClient(ctx, p)
|
||||
if err != nil || id == "" || secret == "" || (p == "oidc" && issuer == "") {
|
||||
continue
|
||||
}
|
||||
display, disabled, err := o.settings.oauthProviderMeta(ctx, p)
|
||||
@@ -105,7 +112,7 @@ func (o *OAuthService) oauth2Config(ctx context.Context, provider string) (*oaut
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if clientID == "" {
|
||||
if clientID == "" || secret == "" || (provider == "oidc" && issuer == "") {
|
||||
return nil, nil, ErrOAuthNotConfigured
|
||||
}
|
||||
if o.settings.EffectiveAppURL() == "" {
|
||||
@@ -128,12 +135,19 @@ func (o *OAuthService) oauth2Config(ctx context.Context, provider string) (*oaut
|
||||
|
||||
// AuthorizeURL 构造授权跳转 URL 并登记一次性 state;mode 为 bind 时 username 必填。
|
||||
// login 模式拒绝已禁用的 provider;bind 模式不受禁用影响(管理员仍可绑定)。
|
||||
func (o *OAuthService) AuthorizeURL(ctx context.Context, provider, mode, username string) (string, error) {
|
||||
func (o *OAuthService) AuthorizeURL(ctx context.Context, provider, mode, username, bindToken string) (string, error) {
|
||||
if mode == "login" {
|
||||
if _, disabled, err := o.settings.oauthProviderMeta(ctx, provider); err == nil && disabled {
|
||||
return "", ErrOAuthDisabled
|
||||
}
|
||||
}
|
||||
var proof TokenProof
|
||||
if mode == "bind" {
|
||||
var err error
|
||||
if proof, err = o.bindTokenProof(ctx, username, bindToken); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
cfg, _, err := o.oauth2Config(ctx, provider)
|
||||
if err != nil {
|
||||
return "", err
|
||||
@@ -148,7 +162,7 @@ func (o *OAuthService) AuthorizeURL(ctx context.Context, provider, mode, usernam
|
||||
}
|
||||
o.mu.Lock()
|
||||
o.gcPendingLocked()
|
||||
o.pending[state] = oauthPending{provider: provider, mode: mode, username: username, nonce: nonce, expires: time.Now().Add(oauthPendingTTL)}
|
||||
o.pending[state] = oauthPending{provider: provider, mode: mode, username: username, nonce: nonce, token: bindToken, proof: proof, expires: time.Now().Add(oauthPendingTTL)}
|
||||
o.mu.Unlock()
|
||||
opts := []oauth2.AuthCodeOption{}
|
||||
if provider == "oidc" {
|
||||
@@ -157,6 +171,16 @@ func (o *OAuthService) AuthorizeURL(ctx context.Context, provider, mode, usernam
|
||||
return cfg.AuthCodeURL(state, opts...), nil
|
||||
}
|
||||
|
||||
// bindTokenProof 验证 bind 发起令牌的有效性与归属,返回其版本 / jti 快照;
|
||||
// 绑定事务行锁下复核该快照,发起后改密、撤销全部、注销或定点撤销均令绑定作废。
|
||||
func (o *OAuthService) bindTokenProof(ctx context.Context, username, token string) (TokenProof, error) {
|
||||
name, proof, err := o.auth.ParseTokenProof(ctx, token)
|
||||
if err != nil || name != username {
|
||||
return TokenProof{}, ErrOAuthState
|
||||
}
|
||||
return proof, nil
|
||||
}
|
||||
|
||||
// gcPendingLocked 清理过期流程;调用方须持锁。
|
||||
func (o *OAuthService) gcPendingLocked() {
|
||||
now := time.Now()
|
||||
@@ -187,27 +211,52 @@ type externalIdentity struct {
|
||||
|
||||
// HandleCallback 完成授权码回调:换取身份后,bind 模式写绑定、login 模式签发 JWT;
|
||||
// token 仅 login 模式非空;mode 尽力返回(state 无效时为空),供 api 决定错误回跳页面。
|
||||
func (o *OAuthService) HandleCallback(ctx context.Context, provider, state, code string) (token, username, mode string, err error) {
|
||||
func (o *OAuthService) HandleCallback(ctx context.Context, provider, state, code string, meta SessionMeta) (token, username, mode string, err error) {
|
||||
p, err := o.takeState(provider, state)
|
||||
if err != nil {
|
||||
return "", "", "", err
|
||||
}
|
||||
if err := o.validateCallbackPending(ctx, provider, p); err != nil {
|
||||
return "", "", p.mode, err
|
||||
}
|
||||
ident, err := o.fetchIdentity(ctx, provider, code, p.nonce)
|
||||
if err != nil {
|
||||
return "", "", p.mode, err
|
||||
}
|
||||
if p.mode == "bind" {
|
||||
if err := o.bind(ctx, p.username, provider, ident); err != nil {
|
||||
return "", p.username, p.mode, err
|
||||
}
|
||||
// 绑定属敏感变更:版本递增使旧令牌失效,同时为操作者签新令牌随回跳带回
|
||||
token, _, err := o.auth.RevokeSessions(ctx, p.username)
|
||||
// 绑定不改登录方式:接续行保留原 method,旧版无行时新建 method 为空
|
||||
// (该令牌并非经新绑定方式登录,与活跃会话设计一致)
|
||||
token, err := o.bind(ctx, p, ident, meta)
|
||||
return token, p.username, p.mode, err
|
||||
}
|
||||
token, username, err = o.loginByIdentity(ctx, provider, ident)
|
||||
meta.Method = provider
|
||||
token, username, err = o.loginByIdentity(ctx, provider, ident, meta)
|
||||
return token, username, p.mode, err
|
||||
}
|
||||
|
||||
func (o *OAuthService) validateCallbackPending(ctx context.Context, provider string, p oauthPending) error {
|
||||
if p.mode == "login" {
|
||||
return o.ensureProviderLoginEnabled(ctx, provider)
|
||||
}
|
||||
// bind 回调换码前复验发起令牌,撤销后流程立即作废。
|
||||
name, err := o.auth.ParseToken(ctx, p.token)
|
||||
if err != nil || name != p.username {
|
||||
return ErrOAuthState
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (o *OAuthService) ensureProviderLoginEnabled(ctx context.Context, provider string) error {
|
||||
_, disabled, err := o.settings.oauthProviderMeta(ctx, provider)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if disabled {
|
||||
return ErrOAuthDisabled
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// fetchIdentity 用授权码向 provider 换取稳定 subject 与展示名。
|
||||
func (o *OAuthService) fetchIdentity(ctx context.Context, provider, code, nonce string) (externalIdentity, error) {
|
||||
cfg, op, err := o.oauth2Config(ctx, provider)
|
||||
@@ -273,45 +322,110 @@ func oidcIdentity(ctx context.Context, cfg *oauth2.Config, op *oidc.Provider, to
|
||||
return externalIdentity{Subject: idToken.Subject, Display: display}, nil
|
||||
}
|
||||
|
||||
// bind 把外部身份绑定到账号;(provider, subject) 唯一,重复绑定报错。
|
||||
func (o *OAuthService) bind(ctx context.Context, username, provider string, ident externalIdentity) error {
|
||||
user, err := o.auth.findUser(ctx, username)
|
||||
if err != nil {
|
||||
// bind 在单事务内完成绑定全程:行锁下比对发起时令牌版本(发起后被撤销即作废)、
|
||||
// 身份写入、版本递增、原会话行接续换发;任一失败整体回滚,不留半程状态。
|
||||
// 新令牌接续 p.token 的会话行(保留登录方式与创建时间),旧令牌无行时按 meta 新建。
|
||||
func (o *OAuthService) bind(ctx context.Context, p oauthPending, ident externalIdentity, meta SessionMeta) (string, error) {
|
||||
var token string
|
||||
err := o.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, p.username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := o.auth.ensureTokenCurrentTx(tx, user, p.proof); err != nil {
|
||||
return ErrOAuthState
|
||||
}
|
||||
if err := createIdentityTx(tx, user.ID, p.provider, ident.Subject, ident.Display, ErrOAuthBound); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := bumpTokenVersionTx(tx, p.username); err != nil {
|
||||
return err
|
||||
}
|
||||
user.TokenVersion++
|
||||
token, _, err = o.auth.renewSessionTx(tx, user, p.token, meta)
|
||||
return err
|
||||
}
|
||||
})
|
||||
return token, err
|
||||
}
|
||||
|
||||
// createIdentityTx 查重后写入外部身份;(provider,subject) 已存在返回 dupErr。
|
||||
func createIdentityTx(tx *gorm.DB, userID uint, provider, subject, display string, dupErr error) error {
|
||||
var count int64
|
||||
err = o.db.WithContext(ctx).Model(&model.UserIdentity{}).
|
||||
Where("provider = ? AND subject = ?", provider, ident.Subject).Count(&count).Error
|
||||
err := tx.Model(&model.UserIdentity{}).
|
||||
Where("provider = ? AND subject = ?", provider, subject).Count(&count).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("check identity: %w", err)
|
||||
}
|
||||
if count > 0 {
|
||||
return ErrOAuthBound
|
||||
return dupErr
|
||||
}
|
||||
row := model.UserIdentity{UserID: user.ID, Provider: provider, Subject: ident.Subject, Display: ident.Display}
|
||||
if err := o.db.WithContext(ctx).Create(&row).Error; err != nil {
|
||||
row := model.UserIdentity{UserID: userID, Provider: provider, Subject: subject, Display: display}
|
||||
if err := tx.Create(&row).Error; err != nil {
|
||||
return fmt.Errorf("bind identity: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// loginByIdentity 查绑定关系并签发面板 JWT;未绑定一律拒绝(不开放注册)。
|
||||
func (o *OAuthService) loginByIdentity(ctx context.Context, provider string, ident externalIdentity) (string, string, error) {
|
||||
// loginByIdentity 查绑定关系并签发面板 JWT(落地会话);未绑定一律拒绝(不开放注册)。
|
||||
func (o *OAuthService) loginByIdentity(ctx context.Context, provider string, ident externalIdentity, meta SessionMeta) (string, string, error) {
|
||||
row, err := o.findIdentity(ctx, provider, ident.Subject)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
return o.loginIdentityRow(ctx, row, provider, ident.Subject, meta)
|
||||
}
|
||||
|
||||
func (o *OAuthService) loginIdentityRow(
|
||||
ctx context.Context, row *model.UserIdentity, provider, subject string, meta SessionMeta,
|
||||
) (string, string, error) {
|
||||
var token, username string
|
||||
err := o.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserByIDForAuthChange(tx, row.UserID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := identityStillBoundTx(tx, row.ID, user.ID, provider, subject); err != nil {
|
||||
return err
|
||||
}
|
||||
ok, err := oauthProviderUsableTx(tx, provider)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
return ErrOAuthDisabled
|
||||
}
|
||||
token, _, err = o.auth.signSessionTokenTx(tx, user, meta)
|
||||
username = user.Username
|
||||
return err
|
||||
})
|
||||
return token, username, err
|
||||
}
|
||||
|
||||
func (o *OAuthService) findIdentity(ctx context.Context, provider, subject string) (*model.UserIdentity, error) {
|
||||
var row model.UserIdentity
|
||||
err := o.db.WithContext(ctx).
|
||||
Where("provider = ? AND subject = ?", provider, ident.Subject).First(&row).Error
|
||||
Where("provider = ? AND subject = ?", provider, subject).First(&row).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, ErrOAuthNotBound
|
||||
}
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return "", "", ErrOAuthNotBound
|
||||
}
|
||||
return "", "", fmt.Errorf("find identity: %w", err)
|
||||
return nil, fmt.Errorf("find identity: %w", err)
|
||||
}
|
||||
var user model.User
|
||||
if err := o.db.WithContext(ctx).First(&user, row.UserID).Error; err != nil {
|
||||
return "", "", fmt.Errorf("find bound user: %w", err)
|
||||
return &row, nil
|
||||
}
|
||||
|
||||
func identityStillBoundTx(tx *gorm.DB, id, userID uint, provider, subject string) error {
|
||||
var count int64
|
||||
err := tx.Model(&model.UserIdentity{}).
|
||||
Where("id = ? AND user_id = ? AND provider = ? AND subject = ?", id, userID, provider, subject).
|
||||
Count(&count).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("recheck identity: %w", err)
|
||||
}
|
||||
token, _, err := o.auth.signToken(user.Username, user.TokenVersion)
|
||||
return token, user.Username, err
|
||||
if count == 0 {
|
||||
return ErrOAuthNotBound
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Identities 列出账号已绑定的外部身份。
|
||||
@@ -330,13 +444,20 @@ func (o *OAuthService) Identities(ctx context.Context, username string) ([]model
|
||||
|
||||
// Unbind 解绑外部身份(校验归属);密码登录被禁用时不允许解绑最后一个身份,防自锁。
|
||||
// 检查与删除在同一事务内并锁定用户行,防与禁用密码登录并发绕过「至少一种登录方式」。
|
||||
func (o *OAuthService) Unbind(ctx context.Context, username string, id uint) error {
|
||||
func (o *OAuthService) Unbind(ctx context.Context, username string, id uint, proof TokenProof) error {
|
||||
err := o.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensureNotLastLogin(tx, user.ID); err != nil {
|
||||
if err := o.auth.ensureTokenCurrentTx(tx, user, proof); err != nil {
|
||||
return err
|
||||
}
|
||||
origin, err := effectiveOriginTx(tx, o.settings)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensureNotLastLogin(tx, user.ID, id, origin); err != nil {
|
||||
return err
|
||||
}
|
||||
res := tx.Where("id = ? AND user_id = ?", id, user.ID).Delete(&model.UserIdentity{})
|
||||
@@ -346,30 +467,25 @@ func (o *OAuthService) Unbind(ctx context.Context, username string, id uint) err
|
||||
if res.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
return nil
|
||||
// 解绑属敏感变更:版本递增与删除同事务提交,不留「已删而旧令牌仍有效」半程
|
||||
return bumpTokenVersionTx(tx, username)
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
// 解绑属敏感变更:递增令牌版本,已签发会话全部失效
|
||||
return o.auth.bumpTokenVersion(ctx, username)
|
||||
return err
|
||||
}
|
||||
|
||||
// ensureNotLastLogin 事务内校验不变量:仅剩一个身份且密码登录已禁用时拒绝解绑;
|
||||
// ensureNotLastLogin 事务内校验不变量:密码登录已禁用时,解绑该身份后
|
||||
// 须仍存在可实际登录的免密方式(provider 被禁用的身份不算),否则拒绝;
|
||||
// 开关读取失败按失败关闭处理(返回错误),不允许失败放行造成自锁。
|
||||
func ensureNotLastLogin(tx *gorm.DB, userID uint) error {
|
||||
n, err := identityCountTx(tx, userID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n > 1 {
|
||||
return nil
|
||||
}
|
||||
func ensureNotLastLogin(tx *gorm.DB, userID, identityID uint, origin string) error {
|
||||
off, err := settingValueTx(tx, settingSecPasswordLoginOff)
|
||||
if err != nil || off != "1" {
|
||||
return err
|
||||
}
|
||||
ok, err := usablePasswordlessTx(tx, userID, identityID, 0, origin)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if off == "1" {
|
||||
if !ok {
|
||||
return ErrLastIdentity
|
||||
}
|
||||
return nil
|
||||
|
||||
+177
-16
@@ -4,6 +4,8 @@ import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// OAuth provider 配置键;client secret 以 AES-GCM 密文落库。
|
||||
@@ -103,18 +105,83 @@ func flagPtr(p *bool) *string {
|
||||
return &v
|
||||
}
|
||||
|
||||
// UpdateOAuth 部分更新 provider 配置:只落库非 nil 字段;
|
||||
// issuer 规范化去尾斜杠,secret 加密落库(空串清除)。
|
||||
// UpdateOAuth 部分更新 provider 配置:只落库非 nil 字段;issuer 规范化去尾斜杠,
|
||||
// secret 加密落库(空串清除)。预检与写入在同一事务并持有认证变更共用的用户行锁:
|
||||
// 密码登录禁用期间,禁止把最后可实际登录的方式禁用或清空(防自锁),
|
||||
// 且不与「禁用密码 / 删除最后因子」并发交错。
|
||||
func (s *SettingService) UpdateOAuth(ctx context.Context, in UpdateOAuthInput) error {
|
||||
return s.updateOAuth(ctx, in, nil)
|
||||
}
|
||||
|
||||
// UpdateOAuthAuthenticated 在写事务持锁后复核请求令牌,防撤销后的慢请求落库。
|
||||
func (s *SettingService) UpdateOAuthAuthenticated(
|
||||
ctx context.Context, in UpdateOAuthInput, auth *AuthService, username string, proof TokenProof,
|
||||
) error {
|
||||
check := &authenticatedMutation{auth: auth, username: username, proof: proof}
|
||||
return s.updateOAuth(ctx, in, check)
|
||||
}
|
||||
|
||||
func (s *SettingService) updateOAuth(ctx context.Context, in UpdateOAuthInput, check *authenticatedMutation) error {
|
||||
secrets, err := s.encryptOAuthSecrets(in)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if check != nil {
|
||||
err = check.lockAndCheck(tx)
|
||||
} else {
|
||||
err = lockUsersForAuthChange(tx)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
origin, err := effectiveOriginTx(tx, s)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensureLoginRemainsTx(tx, in, origin); err != nil {
|
||||
return err
|
||||
}
|
||||
return writeOAuthTx(tx, in, secrets)
|
||||
})
|
||||
}
|
||||
|
||||
// encryptOAuthSecrets 预先加密补丁中的 secret(nil 沿用,空串清除),事务内直接落库。
|
||||
func (s *SettingService) encryptOAuthSecrets(in UpdateOAuthInput) (map[string]*string, error) {
|
||||
out := map[string]*string{}
|
||||
for key, sec := range map[string]*string{
|
||||
settingOauthOidcClientSecret: in.OidcClientSecret,
|
||||
settingOauthGithubClientSecret: in.GithubClientSecret,
|
||||
} {
|
||||
if sec == nil {
|
||||
continue
|
||||
}
|
||||
v := ""
|
||||
if *sec != "" {
|
||||
enc, err := s.cipher.EncryptString(*sec)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encrypt oauth secret: %w", err)
|
||||
}
|
||||
v = enc
|
||||
}
|
||||
out[key] = &v
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// writeOAuthTx 事务内落库补丁中出现的字段。
|
||||
func writeOAuthTx(tx *gorm.DB, in UpdateOAuthInput, secrets map[string]*string) error {
|
||||
writes := []struct {
|
||||
key string
|
||||
val *string
|
||||
}{
|
||||
{settingOauthOidcIssuer, issuerPtr(in.OidcIssuer)},
|
||||
{settingOauthOidcClientID, trimPtr(in.OidcClientID)},
|
||||
{settingOauthOidcClientSecret, secrets[settingOauthOidcClientSecret]},
|
||||
{settingOauthOidcDisplayName, trimPtr(in.OidcDisplayName)},
|
||||
{settingOauthOidcDisabled, flagPtr(in.OidcDisabled)},
|
||||
{settingOauthGithubClientID, trimPtr(in.GithubClientID)},
|
||||
{settingOauthGithubClientSecret, secrets[settingOauthGithubClientSecret]},
|
||||
{settingOauthGithubDisplayName, trimPtr(in.GithubDisplayName)},
|
||||
{settingOauthGithubDisabled, flagPtr(in.GithubDisabled)},
|
||||
}
|
||||
@@ -122,29 +189,123 @@ func (s *SettingService) UpdateOAuth(ctx context.Context, in UpdateOAuthInput) e
|
||||
if w.val == nil {
|
||||
continue
|
||||
}
|
||||
if err := s.set(ctx, w.key, *w.val); err != nil {
|
||||
if err := saveSettingTx(tx, w.key, *w.val); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
if err := s.saveOAuthSecret(ctx, settingOauthOidcClientSecret, in.OidcClientSecret); err != nil {
|
||||
return err
|
||||
}
|
||||
return s.saveOAuthSecret(ctx, settingOauthGithubClientSecret, in.GithubClientSecret)
|
||||
return nil
|
||||
}
|
||||
|
||||
// saveOAuthSecret 加密保存 secret;nil 沿用,空串清除。
|
||||
func (s *SettingService) saveOAuthSecret(ctx context.Context, key string, secret *string) error {
|
||||
if secret == nil {
|
||||
return nil
|
||||
// ensureLoginRemainsTx 事务内校验补丁生效后仍有可实际登录的方式:密码可登直接放行;
|
||||
// 否则须有「可登录且已绑定身份」的 provider、任一通行密钥或钱包身份
|
||||
// (单管理员面板,不区分账号统计)。开关读取失败按失败关闭处理。
|
||||
func ensureLoginRemainsTx(tx *gorm.DB, in UpdateOAuthInput, origin string) error {
|
||||
off, err := settingValueTx(tx, settingSecPasswordLoginOff)
|
||||
if err != nil || off != "1" {
|
||||
return err
|
||||
}
|
||||
if *secret == "" {
|
||||
return s.set(ctx, key, "")
|
||||
if origin == "" {
|
||||
return ErrProviderLastLogin
|
||||
}
|
||||
enc, err := s.cipher.EncryptString(*secret)
|
||||
view, err := oauthViewTx(tx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encrypt oauth secret: %w", err)
|
||||
return err
|
||||
}
|
||||
return s.set(ctx, key, enc)
|
||||
if ok, err := anyBoundUsableProviderTx(tx, patchedUsable(view, in)); err != nil || ok {
|
||||
return err
|
||||
}
|
||||
// 通行密钥兜底须当前地址下可用(origin 一致);0 表示不排除任何行
|
||||
if n, err := passkeyCountExcludingTx(tx, 0, 0, origin); err != nil || n > 0 {
|
||||
return err
|
||||
}
|
||||
n, err := identityProviderCountTx(tx, "wallet")
|
||||
if err != nil || n > 0 {
|
||||
return err
|
||||
}
|
||||
return ErrProviderLastLogin
|
||||
}
|
||||
|
||||
// oauthViewTx 事务内读 provider 配置视图(secret 只取「已设置」布尔)。
|
||||
func oauthViewTx(tx *gorm.DB) (OAuthProvidersView, error) {
|
||||
var view OAuthProvidersView
|
||||
reads := []struct {
|
||||
key string
|
||||
set func(string)
|
||||
}{
|
||||
{settingOauthOidcIssuer, func(v string) { view.OidcIssuer = v }},
|
||||
{settingOauthOidcClientID, func(v string) { view.OidcClientID = v }},
|
||||
{settingOauthOidcClientSecret, func(v string) { view.OidcSecretSet = v != "" }},
|
||||
{settingOauthOidcDisabled, func(v string) { view.OidcDisabled = v == "1" }},
|
||||
{settingOauthGithubClientID, func(v string) { view.GithubClientID = v }},
|
||||
{settingOauthGithubClientSecret, func(v string) { view.GithubSecretSet = v != "" }},
|
||||
{settingOauthGithubDisabled, func(v string) { view.GithubDisabled = v == "1" }},
|
||||
}
|
||||
for _, r := range reads {
|
||||
v, err := settingValueTx(tx, r.key)
|
||||
if err != nil {
|
||||
return view, err
|
||||
}
|
||||
r.set(v)
|
||||
}
|
||||
return view, nil
|
||||
}
|
||||
|
||||
// patchedOidcUsable 计算补丁生效后 OIDC 是否可登录(clientID/secret/issuer 齐备且未禁用)。
|
||||
func patchedOidcUsable(view OAuthProvidersView, in UpdateOAuthInput) bool {
|
||||
id, sec, iss, off := view.OidcClientID, view.OidcSecretSet, view.OidcIssuer, view.OidcDisabled
|
||||
if v := trimPtr(in.OidcClientID); v != nil {
|
||||
id = *v
|
||||
}
|
||||
if in.OidcClientSecret != nil {
|
||||
sec = *in.OidcClientSecret != ""
|
||||
}
|
||||
if v := issuerPtr(in.OidcIssuer); v != nil {
|
||||
iss = *v
|
||||
}
|
||||
if in.OidcDisabled != nil {
|
||||
off = *in.OidcDisabled
|
||||
}
|
||||
return id != "" && sec && iss != "" && !off
|
||||
}
|
||||
|
||||
// patchedGithubUsable 计算补丁生效后 GitHub 是否可登录(clientID/secret 齐备且未禁用)。
|
||||
func patchedGithubUsable(view OAuthProvidersView, in UpdateOAuthInput) bool {
|
||||
id, sec, off := view.GithubClientID, view.GithubSecretSet, view.GithubDisabled
|
||||
if v := trimPtr(in.GithubClientID); v != nil {
|
||||
id = *v
|
||||
}
|
||||
if in.GithubClientSecret != nil {
|
||||
sec = *in.GithubClientSecret != ""
|
||||
}
|
||||
if in.GithubDisabled != nil {
|
||||
off = *in.GithubDisabled
|
||||
}
|
||||
return id != "" && sec && !off
|
||||
}
|
||||
|
||||
// patchedUsable 汇总补丁生效后各 provider 的可登录性。
|
||||
func patchedUsable(view OAuthProvidersView, in UpdateOAuthInput) map[string]bool {
|
||||
return map[string]bool{
|
||||
"oidc": patchedOidcUsable(view, in),
|
||||
"github": patchedGithubUsable(view, in),
|
||||
}
|
||||
}
|
||||
|
||||
// anyBoundUsableProviderTx 事务内判断是否存在「可登录且已有绑定身份」的 provider。
|
||||
func anyBoundUsableProviderTx(tx *gorm.DB, usable map[string]bool) (bool, error) {
|
||||
for p, ok := range usable {
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
n, err := identityProviderCountTx(tx, p)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
if n > 0 {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// oauthClient 返回 provider 的 clientID/明文 secret/issuer(仅 oidc);未配置时 clientID 为空。
|
||||
|
||||
@@ -0,0 +1,448 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/binary"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/url"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/go-webauthn/webauthn/protocol"
|
||||
"github.com/go-webauthn/webauthn/webauthn"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
|
||||
// 通行密钥流程错误;api 层映射为用户可读提示。
|
||||
var (
|
||||
// ErrPasskeyNoAppURL 表示面板地址缺失,RP ID 无从派生。
|
||||
ErrPasskeyNoAppURL = errors.New("面板地址未设置,请先在「设置 → 安全 → 网络与地址」保存面板地址")
|
||||
// ErrPasskeySession 表示挑战会话无效或已过期(一次性消费)。
|
||||
ErrPasskeySession = errors.New("通行密钥会话无效或已过期,请重新发起")
|
||||
// ErrPasskeyVerify 表示凭据校验失败;不透出库层细节防探测。
|
||||
ErrPasskeyVerify = errors.New("通行密钥校验失败")
|
||||
// ErrPasskeyLimit 表示已达单账号凭据数量上限。
|
||||
ErrPasskeyLimit = errors.New("通行密钥数量已达上限,请先删除不用的")
|
||||
)
|
||||
|
||||
const (
|
||||
// passkeyPendingTTL 是挑战会话有效期(WebAuthn 默认超时同量级)。
|
||||
passkeyPendingTTL = 5 * time.Minute
|
||||
// passkeyMaxPerUser 是单账号凭据上限,防滥用堆积。
|
||||
passkeyMaxPerUser = 10
|
||||
// passkeyGuardUser 是登录守卫的用户名占位:discoverable 登录失败时账号未知。
|
||||
passkeyGuardUser = "__passkey__"
|
||||
)
|
||||
|
||||
// passkeyPending 是一次进行中的 WebAuthn 仪式上下文;sessionId 一次性使用。
|
||||
type passkeyPending struct {
|
||||
session webauthn.SessionData
|
||||
username string // 注册会话的归属账号;登录会话为空串
|
||||
expires time.Time
|
||||
}
|
||||
|
||||
// PasskeyService 承接通行密钥(WebAuthn)注册、登录与凭据管理。
|
||||
type PasskeyService struct {
|
||||
db *gorm.DB
|
||||
settings *SettingService
|
||||
auth *AuthService
|
||||
|
||||
mu sync.Mutex
|
||||
pending map[string]passkeyPending
|
||||
}
|
||||
|
||||
// NewPasskeyService 组装依赖。
|
||||
func NewPasskeyService(db *gorm.DB, settings *SettingService, auth *AuthService) *PasskeyService {
|
||||
return &PasskeyService{db: db, settings: settings, auth: auth, pending: map[string]passkeyPending{}}
|
||||
}
|
||||
|
||||
// rp 按当前面板地址构造 WebAuthn 实例(RP ID = 域名,Origin = 完整来源);
|
||||
// 每次现算使「面板地址」设置变更即时生效,构造仅做配置校验、代价可忽略。
|
||||
func (p *PasskeyService) rp() (*webauthn.WebAuthn, error) {
|
||||
app := p.settings.EffectiveAppURL()
|
||||
if app == "" {
|
||||
return nil, ErrPasskeyNoAppURL
|
||||
}
|
||||
u, err := url.Parse(app)
|
||||
if err != nil || u.Hostname() == "" {
|
||||
return nil, fmt.Errorf("parse app url: %w", err)
|
||||
}
|
||||
return webauthn.New(&webauthn.Config{
|
||||
RPID: u.Hostname(),
|
||||
RPDisplayName: "OCI Portal",
|
||||
RPOrigins: []string{u.Scheme + "://" + u.Host},
|
||||
})
|
||||
}
|
||||
|
||||
// passkeyUserHandle 是 WebAuthn user.id:用户主键的 8 字节大端序(改用户名不漂移)。
|
||||
func passkeyUserHandle(id uint) []byte {
|
||||
b := make([]byte, 8)
|
||||
binary.BigEndian.PutUint64(b, uint64(id))
|
||||
return b
|
||||
}
|
||||
|
||||
// passkeyUser 以账号与其凭据集实现 webauthn.User。
|
||||
type passkeyUser struct {
|
||||
user model.User
|
||||
keys []model.UserPasskey
|
||||
}
|
||||
|
||||
func (u passkeyUser) WebAuthnID() []byte { return passkeyUserHandle(u.user.ID) }
|
||||
func (u passkeyUser) WebAuthnName() string { return u.user.Username }
|
||||
func (u passkeyUser) WebAuthnDisplayName() string { return u.user.Username }
|
||||
|
||||
// WebAuthnCredentials 反序列化各行凭据 JSON;坏行跳过不阻断整体。
|
||||
func (u passkeyUser) WebAuthnCredentials() []webauthn.Credential {
|
||||
out := make([]webauthn.Credential, 0, len(u.keys))
|
||||
for _, k := range u.keys {
|
||||
var c webauthn.Credential
|
||||
if json.Unmarshal([]byte(k.Credential), &c) == nil {
|
||||
out = append(out, c)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// exclusions 生成注册排除清单,阻止同一验证器重复注册。
|
||||
func (u passkeyUser) exclusions() []protocol.CredentialDescriptor {
|
||||
out := []protocol.CredentialDescriptor{}
|
||||
for _, c := range u.WebAuthnCredentials() {
|
||||
out = append(out, c.Descriptor())
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// putPending 登记一次性挑战会话,返回下发给前端的 sessionId。
|
||||
func (p *PasskeyService) putPending(session webauthn.SessionData, username string) (string, error) {
|
||||
id, err := randHex(16)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
p.mu.Lock()
|
||||
p.gcPasskeyLocked()
|
||||
p.pending[id] = passkeyPending{session: session, username: username, expires: time.Now().Add(passkeyPendingTTL)}
|
||||
p.mu.Unlock()
|
||||
return id, nil
|
||||
}
|
||||
|
||||
// takePending 取出并消费会话(一次性);不存在或过期视为无效。
|
||||
func (p *PasskeyService) takePending(id string) (passkeyPending, error) {
|
||||
p.mu.Lock()
|
||||
defer p.mu.Unlock()
|
||||
entry, ok := p.pending[id]
|
||||
delete(p.pending, id)
|
||||
if !ok || time.Now().After(entry.expires) {
|
||||
return passkeyPending{}, ErrPasskeySession
|
||||
}
|
||||
return entry, nil
|
||||
}
|
||||
|
||||
// gcPasskeyLocked 清理过期会话;调用方须持锁。
|
||||
func (p *PasskeyService) gcPasskeyLocked() {
|
||||
now := time.Now()
|
||||
for k, e := range p.pending {
|
||||
if now.After(e.expires) {
|
||||
delete(p.pending, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// loadUser 载入账号与全部凭据。
|
||||
func (p *PasskeyService) loadUser(ctx context.Context, username string) (passkeyUser, error) {
|
||||
user, err := p.auth.findUser(ctx, username)
|
||||
if err != nil {
|
||||
return passkeyUser{}, err
|
||||
}
|
||||
return p.attachKeys(ctx, *user)
|
||||
}
|
||||
|
||||
// attachKeys 挂载账号的凭据行。
|
||||
func (p *PasskeyService) attachKeys(ctx context.Context, user model.User) (passkeyUser, error) {
|
||||
keys := []model.UserPasskey{}
|
||||
err := p.db.WithContext(ctx).Where("user_id = ?", user.ID).Order("id").Find(&keys).Error
|
||||
if err != nil {
|
||||
return passkeyUser{}, fmt.Errorf("list passkeys: %w", err)
|
||||
}
|
||||
return passkeyUser{user: user, keys: keys}, nil
|
||||
}
|
||||
|
||||
// BeginRegister 生成注册 options;ResidentKey 与用户验证均必需——
|
||||
// 可发现凭据是免用户名登录的前提,UV 是 Passkey 登录豁免 TOTP 的前提。
|
||||
func (p *PasskeyService) BeginRegister(ctx context.Context, username string) (string, *protocol.CredentialCreation, error) {
|
||||
w, err := p.rp()
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
u, err := p.loadUser(ctx, username)
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
if len(u.keys) >= passkeyMaxPerUser {
|
||||
return "", nil, ErrPasskeyLimit
|
||||
}
|
||||
opts, session, err := w.BeginRegistration(u,
|
||||
webauthn.WithAuthenticatorSelection(protocol.AuthenticatorSelection{
|
||||
ResidentKey: protocol.ResidentKeyRequirementRequired,
|
||||
UserVerification: protocol.VerificationRequired,
|
||||
}),
|
||||
webauthn.WithExclusions(u.exclusions()),
|
||||
)
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("begin registration: %w", err)
|
||||
}
|
||||
id, err := p.putPending(*session, username)
|
||||
return id, opts, err
|
||||
}
|
||||
|
||||
// FinishRegister 校验注册响应并落库;成功后令牌版本递增(旧会话全部失效)。
|
||||
func (p *PasskeyService) FinishRegister(ctx context.Context, username, sessionID, name string, body io.Reader, proof TokenProof) error {
|
||||
entry, err := p.takePending(sessionID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if entry.username != username {
|
||||
return ErrPasskeySession
|
||||
}
|
||||
w, err := p.rp()
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
u, err := p.loadUser(ctx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
parsed, err := protocol.ParseCredentialCreationResponseBody(body)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%w: %v", ErrPasskeyVerify, err)
|
||||
}
|
||||
cred, err := w.CreateCredential(u, entry.session, parsed)
|
||||
if err != nil {
|
||||
return fmt.Errorf("%w: %v", ErrPasskeyVerify, err)
|
||||
}
|
||||
// 凭据落库与令牌版本递增同事务提交,不留「因子已生效而旧令牌仍有效」的
|
||||
// 半程状态;行锁下复核 proof,请求挂起期间令牌被撤销则整体拒绝
|
||||
return p.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := p.auth.ensureTokenCurrentTx(tx, user, proof); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := saveCredentialTx(tx, user.ID, name, p.currentOrigin(), cred); err != nil {
|
||||
return err
|
||||
}
|
||||
return bumpTokenVersionTx(tx, username)
|
||||
})
|
||||
}
|
||||
|
||||
// currentOrigin 是当前面板地址的 WebAuthn origin(scheme://host);
|
||||
// 注册时随凭据落库,作为该凭据可用性的判定依据。
|
||||
func (p *PasskeyService) currentOrigin() string {
|
||||
return originOf(p.settings.EffectiveAppURL())
|
||||
}
|
||||
|
||||
// passkeyCredHash 是凭据 ID 的 SHA-256 hex:原始 ID 最长可达 1023 字节,
|
||||
// 编码后超出可索引长度,唯一约束与登录反查一律走定长哈希列。
|
||||
func passkeyCredHash(id []byte) string {
|
||||
sum := sha256.Sum256(id)
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// saveCredentialTx 序列化凭据整段落库(事务内);名称缺省给默认值。
|
||||
func saveCredentialTx(tx *gorm.DB, userID uint, name, origin string, cred *webauthn.Credential) error {
|
||||
raw, err := json.Marshal(cred)
|
||||
if err != nil {
|
||||
return fmt.Errorf("marshal credential: %w", err)
|
||||
}
|
||||
if name == "" {
|
||||
name = "通行密钥"
|
||||
}
|
||||
row := model.UserPasskey{
|
||||
UserID: userID,
|
||||
Name: name,
|
||||
CredentialID: base64.RawURLEncoding.EncodeToString(cred.ID),
|
||||
CredentialIDHash: passkeyCredHash(cred.ID),
|
||||
Origin: origin,
|
||||
Credential: string(raw),
|
||||
}
|
||||
if err := tx.Create(&row).Error; err != nil {
|
||||
return fmt.Errorf("save passkey: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// BeginLogin 生成断言 options(公开端点;无凭据也正常下发,不泄露账号状态)。
|
||||
func (p *PasskeyService) BeginLogin(_ context.Context) (string, *protocol.CredentialAssertion, error) {
|
||||
w, err := p.rp()
|
||||
if err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
opts, session, err := w.BeginDiscoverableLogin(
|
||||
webauthn.WithUserVerification(protocol.VerificationRequired))
|
||||
if err != nil {
|
||||
return "", nil, fmt.Errorf("begin login: %w", err)
|
||||
}
|
||||
id, err := p.putPending(*session, "")
|
||||
return id, opts, err
|
||||
}
|
||||
|
||||
// FinishLogin 校验断言并签发 JWT(落地会话),返回登录账号名供留痕;
|
||||
// 失败按 IP 计入登录守卫(账号未知用占位名),锁定期内一律 ErrLoginLocked。
|
||||
func (p *PasskeyService) FinishLogin(ctx context.Context, sessionID string, meta SessionMeta, body io.Reader) (string, time.Time, string, error) {
|
||||
key := guardKey(meta.ClientIP, passkeyGuardUser)
|
||||
now := time.Now()
|
||||
sec := securityOf(p.auth.settings)
|
||||
if p.auth.guard.locked(key, now, time.Duration(sec.LoginLockMinutes)*time.Minute) {
|
||||
return "", time.Time{}, "", ErrLoginLocked
|
||||
}
|
||||
user, err := p.validateLogin(ctx, sessionID, body)
|
||||
if err != nil {
|
||||
if lockErr := p.auth.failLogin(key, now, passkeyGuardUser, meta.ClientIP, sec); errors.Is(lockErr, ErrLoginLocked) {
|
||||
return "", time.Time{}, "", ErrLoginLocked
|
||||
}
|
||||
return "", time.Time{}, "", err
|
||||
}
|
||||
p.auth.guard.success(key)
|
||||
meta.Method = "passkey"
|
||||
token, expires, err := p.auth.signSessionToken(ctx, user, meta)
|
||||
return token, expires, user.Username, err
|
||||
}
|
||||
|
||||
// validateLogin 消费会话并校验断言;成功后回写凭据状态。
|
||||
func (p *PasskeyService) validateLogin(ctx context.Context, sessionID string, body io.Reader) (*model.User, error) {
|
||||
entry, err := p.takePending(sessionID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if entry.username != "" {
|
||||
return nil, ErrPasskeySession // 注册会话不可用于登录
|
||||
}
|
||||
w, err := p.rp()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
parsed, err := protocol.ParseCredentialRequestResponseBody(body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("%w: %v", ErrPasskeyVerify, err)
|
||||
}
|
||||
var owner *model.User
|
||||
handler := func(_, userHandle []byte) (webauthn.User, error) {
|
||||
u, err := p.userByHandle(ctx, userHandle)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
owner = &u.user
|
||||
return u, nil
|
||||
}
|
||||
cred, err := w.ValidateDiscoverableLogin(handler, entry.session, parsed)
|
||||
if err != nil || owner == nil {
|
||||
return nil, ErrPasskeyVerify // 统一文案,不透出细节防探测
|
||||
}
|
||||
p.touchCredential(ctx, cred)
|
||||
return owner, nil
|
||||
}
|
||||
|
||||
// userByHandle 按 userHandle(用户主键大端 8 字节)载入用户与凭据。
|
||||
func (p *PasskeyService) userByHandle(ctx context.Context, handle []byte) (passkeyUser, error) {
|
||||
if len(handle) != 8 {
|
||||
return passkeyUser{}, ErrPasskeyVerify
|
||||
}
|
||||
var user model.User
|
||||
err := p.db.WithContext(ctx).First(&user, uint(binary.BigEndian.Uint64(handle))).Error
|
||||
if err != nil {
|
||||
return passkeyUser{}, fmt.Errorf("find user by handle: %w", err)
|
||||
}
|
||||
return p.attachKeys(ctx, user)
|
||||
}
|
||||
|
||||
// touchCredential 回写校验后的 signCount/flags 与最后使用时间;失败不阻断登录。
|
||||
func (p *PasskeyService) touchCredential(ctx context.Context, cred *webauthn.Credential) {
|
||||
raw, err := json.Marshal(cred)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
p.db.WithContext(ctx).Model(&model.UserPasskey{}).
|
||||
Where("credential_id_hash = ?", passkeyCredHash(cred.ID)).
|
||||
Updates(map[string]any{"credential": string(raw), "last_used_at": time.Now()})
|
||||
}
|
||||
|
||||
// List 列出账号的通行密钥。
|
||||
func (p *PasskeyService) List(ctx context.Context, username string) ([]model.UserPasskey, error) {
|
||||
u, err := p.loadUser(ctx, username)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return u.keys, nil
|
||||
}
|
||||
|
||||
// Remove 删除通行密钥(校验归属)并递增令牌版本;
|
||||
// 密码登录禁用期间通行密钥计入「至少一种登录方式」不变量,删除受防自锁保护。
|
||||
func (p *PasskeyService) Remove(ctx context.Context, username string, id uint, proof TokenProof) error {
|
||||
err := p.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := p.auth.ensureTokenCurrentTx(tx, user, proof); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := ensureNotLastPasskey(tx, user.ID, id, p.currentOrigin()); err != nil {
|
||||
return err
|
||||
}
|
||||
res := tx.Where("id = ? AND user_id = ?", id, user.ID).Delete(&model.UserPasskey{})
|
||||
if res.Error != nil {
|
||||
return fmt.Errorf("remove passkey: %w", res.Error)
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
// 删除属敏感变更:版本递增与删除同事务提交,不留「已删而旧令牌仍有效」半程
|
||||
return bumpTokenVersionTx(tx, username)
|
||||
})
|
||||
return err
|
||||
}
|
||||
|
||||
// ensureNotLastPasskey 事务内校验不变量:密码登录已禁用时,删除该钥匙后
|
||||
// 须仍存在可实际登录的免密方式(provider 被禁用的身份不算),否则拒绝;
|
||||
// 开关读取失败按失败关闭处理,不允许失败放行。
|
||||
func ensureNotLastPasskey(tx *gorm.DB, userID, passkeyID uint, origin string) error {
|
||||
off, err := settingValueTx(tx, settingSecPasswordLoginOff)
|
||||
if err != nil || off != "1" {
|
||||
return err
|
||||
}
|
||||
ok, err := usablePasswordlessTx(tx, userID, 0, passkeyID, origin)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if !ok {
|
||||
return ErrLastIdentity
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// HasAny 报告当前面板 origin 下是否存在可用凭据;地址缺失或查询失败按无处理。
|
||||
func (p *PasskeyService) HasAny(ctx context.Context) bool {
|
||||
if p.settings == nil {
|
||||
return false
|
||||
}
|
||||
origin := p.currentOrigin()
|
||||
if origin == "" {
|
||||
return false
|
||||
}
|
||||
var count int64
|
||||
err := p.db.WithContext(ctx).Model(&model.UserPasskey{}).
|
||||
Where("origin = ?", origin).Count(&count).Error
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return count > 0
|
||||
}
|
||||
@@ -0,0 +1,298 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/go-webauthn/webauthn/webauthn"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
|
||||
"oci-portal/internal/crypto"
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
|
||||
// newTestPasskey 组装内存库上的 PasskeyService(admin 账号已建,面板地址经环境回退注入)。
|
||||
func newTestPasskey(t *testing.T, appURL string) *PasskeyService {
|
||||
t.Helper()
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open in-memory sqlite: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.User{}, &model.UserPasskey{}, &model.UserIdentity{}, &model.UserSession{}, &model.Setting{}); err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
auth := NewAuthService(db, "test-jwt-secret")
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("ensure admin: %v", err)
|
||||
}
|
||||
cipher, err := crypto.NewCipher("test-key")
|
||||
if err != nil {
|
||||
t.Fatalf("new cipher: %v", err)
|
||||
}
|
||||
settings := NewSettingService(db, cipher)
|
||||
settings.SetEnvPublicURL(appURL)
|
||||
return NewPasskeyService(db, settings, auth)
|
||||
}
|
||||
|
||||
func TestPasskeyRPDerivation(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
appURL string
|
||||
wantErr error
|
||||
wantRPID string
|
||||
wantOrigin string
|
||||
}{
|
||||
{name: "无面板地址", appURL: "", wantErr: ErrPasskeyNoAppURL},
|
||||
{name: "https 域名", appURL: "https://demo.example.com", wantRPID: "demo.example.com", wantOrigin: "https://demo.example.com"},
|
||||
{name: "带端口", appURL: "https://panel.example.com:8443", wantRPID: "panel.example.com", wantOrigin: "https://panel.example.com:8443"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
p := newTestPasskey(t, tt.appURL)
|
||||
w, err := p.rp()
|
||||
if tt.wantErr != nil {
|
||||
if !errors.Is(err, tt.wantErr) {
|
||||
t.Fatalf("rp() err = %v, want %v", err, tt.wantErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("rp(): %v", err)
|
||||
}
|
||||
if w.Config.RPID != tt.wantRPID {
|
||||
t.Errorf("RPID = %q, want %q", w.Config.RPID, tt.wantRPID)
|
||||
}
|
||||
if len(w.Config.RPOrigins) != 1 || w.Config.RPOrigins[0] != tt.wantOrigin {
|
||||
t.Errorf("RPOrigins = %v, want [%s]", w.Config.RPOrigins, tt.wantOrigin)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyPendingLifecycle(t *testing.T) {
|
||||
p := newTestPasskey(t, "https://demo.example.com")
|
||||
session := webauthn.SessionData{Challenge: "challenge-1"}
|
||||
|
||||
id, err := p.putPending(session, "admin")
|
||||
if err != nil {
|
||||
t.Fatalf("putPending: %v", err)
|
||||
}
|
||||
got, err := p.takePending(id)
|
||||
if err != nil {
|
||||
t.Fatalf("takePending: %v", err)
|
||||
}
|
||||
if got.session.Challenge != "challenge-1" || got.username != "admin" {
|
||||
t.Errorf("pending = %+v, want challenge-1/admin", got)
|
||||
}
|
||||
// 一次性:再次消费同一 sessionId 必须失效
|
||||
if _, err := p.takePending(id); !errors.Is(err, ErrPasskeySession) {
|
||||
t.Errorf("second take err = %v, want ErrPasskeySession", err)
|
||||
}
|
||||
// 过期条目视为无效
|
||||
expiredID, _ := p.putPending(session, "")
|
||||
p.mu.Lock()
|
||||
e := p.pending[expiredID]
|
||||
e.expires = time.Now().Add(-time.Second)
|
||||
p.pending[expiredID] = e
|
||||
p.mu.Unlock()
|
||||
if _, err := p.takePending(expiredID); !errors.Is(err, ErrPasskeySession) {
|
||||
t.Errorf("expired take err = %v, want ErrPasskeySession", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyBeginRegister(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
appURL string
|
||||
seedKeys int
|
||||
wantErr error
|
||||
}{
|
||||
{name: "正常发起", appURL: "https://demo.example.com"},
|
||||
{name: "无面板地址", appURL: "", wantErr: ErrPasskeyNoAppURL},
|
||||
{name: "数量达上限", appURL: "https://demo.example.com", seedKeys: passkeyMaxPerUser, wantErr: ErrPasskeyLimit},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
p := newTestPasskey(t, tt.appURL)
|
||||
for i := 0; i < tt.seedKeys; i++ {
|
||||
seedPasskeyRow(t, p.db, uint(i+1))
|
||||
}
|
||||
sid, opts, err := p.BeginRegister(context.Background(), "admin")
|
||||
if tt.wantErr != nil {
|
||||
if !errors.Is(err, tt.wantErr) {
|
||||
t.Fatalf("BeginRegister err = %v, want %v", err, tt.wantErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("BeginRegister: %v", err)
|
||||
}
|
||||
if sid == "" || opts == nil || opts.Response.Challenge.String() == "" {
|
||||
t.Errorf("BeginRegister returned empty session/options")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// seedPasskeyRow 直插一行合法凭据(JSON 与 CredentialID 对应)。
|
||||
func seedPasskeyRow(t *testing.T, db *gorm.DB, seq uint) {
|
||||
t.Helper()
|
||||
credID := []byte{byte(seq), 2, 3, 4}
|
||||
cred := webauthn.Credential{ID: credID, PublicKey: []byte{5, 6}}
|
||||
row := model.UserPasskey{
|
||||
UserID: 1,
|
||||
Name: "key",
|
||||
CredentialID: base64.RawURLEncoding.EncodeToString(credID),
|
||||
CredentialIDHash: passkeyCredHash(credID),
|
||||
Credential: mustCredJSON(t, cred),
|
||||
Origin: "https://demo.example.com",
|
||||
}
|
||||
if err := db.Create(&row).Error; err != nil {
|
||||
t.Fatalf("seed passkey: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func mustCredJSON(t *testing.T, cred webauthn.Credential) string {
|
||||
t.Helper()
|
||||
raw, err := json.Marshal(cred)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal credential: %v", err)
|
||||
}
|
||||
return string(raw)
|
||||
}
|
||||
|
||||
func TestPasskeyFinishRegisterSessionChecks(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
sessionUser string
|
||||
finishUser string
|
||||
body string
|
||||
wantErr error
|
||||
}{
|
||||
{name: "会话不存在", sessionUser: "-", finishUser: "admin", wantErr: ErrPasskeySession},
|
||||
{name: "会话归属不符", sessionUser: "other", finishUser: "admin", wantErr: ErrPasskeySession},
|
||||
{name: "登录会话不可注册", sessionUser: "", finishUser: "admin", wantErr: ErrPasskeySession},
|
||||
{name: "凭据体不合法", sessionUser: "admin", finishUser: "admin", body: "{}", wantErr: ErrPasskeyVerify},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
p := newTestPasskey(t, "https://demo.example.com")
|
||||
sid := "missing"
|
||||
if tt.sessionUser != "-" {
|
||||
var err error
|
||||
sid, err = p.putPending(webauthn.SessionData{Challenge: "c"}, tt.sessionUser)
|
||||
if err != nil {
|
||||
t.Fatalf("putPending: %v", err)
|
||||
}
|
||||
}
|
||||
err := p.FinishRegister(context.Background(), tt.finishUser, sid, "名称", strings.NewReader(tt.body), proofOf(t, p.db, "admin"))
|
||||
if !errors.Is(err, tt.wantErr) {
|
||||
t.Fatalf("FinishRegister err = %v, want %v", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyCredentialRoundTrip(t *testing.T) {
|
||||
p := newTestPasskey(t, "https://demo.example.com")
|
||||
cred := &webauthn.Credential{ID: []byte{9, 9, 9}, PublicKey: []byte{1, 2, 3}}
|
||||
cred.Authenticator.SignCount = 7
|
||||
if err := saveCredentialTx(p.db, 1, "", "https://demo.example.com", cred); err != nil {
|
||||
t.Fatalf("saveCredentialTx: %v", err)
|
||||
}
|
||||
u, err := p.loadUser(context.Background(), "admin")
|
||||
if err != nil {
|
||||
t.Fatalf("loadUser: %v", err)
|
||||
}
|
||||
creds := u.WebAuthnCredentials()
|
||||
if len(creds) != 1 || creds[0].Authenticator.SignCount != 7 {
|
||||
t.Fatalf("credentials = %+v, want 1 item signCount 7", creds)
|
||||
}
|
||||
if u.keys[0].Name != "通行密钥" {
|
||||
t.Errorf("default name = %q, want 通行密钥", u.keys[0].Name)
|
||||
}
|
||||
// userHandle 往返:8 字节大端 ID 找回同一账号
|
||||
pu, err := p.userByHandle(context.Background(), passkeyUserHandle(u.user.ID))
|
||||
if err != nil || pu.user.Username != "admin" {
|
||||
t.Errorf("userByHandle = (%+v, %v), want admin", pu.user, err)
|
||||
}
|
||||
if _, err := p.userByHandle(context.Background(), []byte{1, 2}); !errors.Is(err, ErrPasskeyVerify) {
|
||||
t.Errorf("short handle err = %v, want ErrPasskeyVerify", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyRemoveAndVersionBump(t *testing.T) {
|
||||
p := newTestPasskey(t, "https://demo.example.com")
|
||||
seedPasskeyRow(t, p.db, 1)
|
||||
var before model.User
|
||||
if err := p.db.First(&before, 1).Error; err != nil {
|
||||
t.Fatalf("load user: %v", err)
|
||||
}
|
||||
if !p.HasAny(context.Background()) {
|
||||
t.Fatal("HasAny = false, want true after seed")
|
||||
}
|
||||
if err := p.Remove(context.Background(), "admin", 1, proofOf(t, p.db, "admin")); err != nil {
|
||||
t.Fatalf("Remove: %v", err)
|
||||
}
|
||||
if err := p.Remove(context.Background(), "admin", 1, proofOf(t, p.db, "admin")); !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Errorf("second Remove err = %v, want ErrRecordNotFound", err)
|
||||
}
|
||||
if p.HasAny(context.Background()) {
|
||||
t.Error("HasAny = true, want false after remove")
|
||||
}
|
||||
var after model.User
|
||||
if err := p.db.First(&after, 1).Error; err != nil {
|
||||
t.Fatalf("reload user: %v", err)
|
||||
}
|
||||
if after.TokenVersion != before.TokenVersion+1 {
|
||||
t.Errorf("TokenVersion = %d, want %d", after.TokenVersion, before.TokenVersion+1)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyHasAnyRequiresCurrentOrigin(t *testing.T) {
|
||||
p := newTestPasskey(t, "https://demo.example.com")
|
||||
seedPasskeyRow(t, p.db, 1)
|
||||
tests := []struct {
|
||||
name string
|
||||
appURL string
|
||||
want bool
|
||||
}{
|
||||
{name: "registered origin", appURL: "https://demo.example.com", want: true},
|
||||
{name: "missing app url", appURL: ""},
|
||||
{name: "different origin", appURL: "https://other.example.com"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
p.settings.SetEnvPublicURL(tt.appURL)
|
||||
if got := p.HasAny(context.Background()); got != tt.want {
|
||||
t.Fatalf("HasAny = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPasskeyFinishLoginGuard(t *testing.T) {
|
||||
p := newTestPasskey(t, "https://demo.example.com")
|
||||
ctx := context.Background()
|
||||
// 无效 sessionId 反复失败:达到默认阈值后转锁定
|
||||
var lastErr error
|
||||
for i := 0; i < securityDefaults.LoginFailLimit+1; i++ {
|
||||
_, _, _, lastErr = p.FinishLogin(ctx, "bad-session", SessionMeta{ClientIP: "10.0.0.9"}, strings.NewReader("{}"))
|
||||
}
|
||||
if !errors.Is(lastErr, ErrLoginLocked) {
|
||||
t.Fatalf("after %d failures err = %v, want ErrLoginLocked", securityDefaults.LoginFailLimit+1, lastErr)
|
||||
}
|
||||
// 其他 IP 不受连坐
|
||||
if _, _, _, err := p.FinishLogin(ctx, "bad-session", SessionMeta{ClientIP: "10.0.0.10"}, strings.NewReader("{}")); errors.Is(err, ErrLoginLocked) {
|
||||
t.Errorf("different IP got locked prematurely: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -3,9 +3,12 @@ package service
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"strconv"
|
||||
"strings"
|
||||
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// 安全设置键:WAF 参数 / 真实IP请求头 / 面板地址(app_url)。
|
||||
@@ -147,17 +150,49 @@ func (p SecurityPatch) merge(dst *SecuritySettings) map[string]bool {
|
||||
return touched
|
||||
}
|
||||
|
||||
// UpdateSecurity 把补丁合并到现值上整体校验,只落库出现的字段,再重读刷新快照。
|
||||
type securityUpdate struct {
|
||||
next SecuritySettings
|
||||
touched map[string]bool
|
||||
values map[string]string
|
||||
}
|
||||
|
||||
// UpdateSecurity 供启动与内部配置使用;HTTP PATCH 使用带令牌证明的变体。
|
||||
func (s *SettingService) UpdateSecurity(ctx context.Context, p SecurityPatch) error {
|
||||
cur, err := s.Security(ctx)
|
||||
return s.updateSecurity(ctx, p, nil)
|
||||
}
|
||||
|
||||
// UpdateSecurityAuthenticated 在写事务持锁后复核请求令牌,防撤销后的慢请求落库。
|
||||
func (s *SettingService) UpdateSecurityAuthenticated(
|
||||
ctx context.Context, p SecurityPatch, auth *AuthService, username string, proof TokenProof,
|
||||
) error {
|
||||
check := &authenticatedMutation{auth: auth, username: username, proof: proof}
|
||||
return s.updateSecurity(ctx, p, check)
|
||||
}
|
||||
|
||||
func (s *SettingService) updateSecurity(ctx context.Context, p SecurityPatch, check *authenticatedMutation) error {
|
||||
up, err := s.prepareSecurityUpdate(ctx, p)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
touched := p.merge(&cur)
|
||||
if err := validateSecurity(&cur); err != nil {
|
||||
err = s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
return s.applySecurityUpdateTx(tx, up, check)
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
kv := map[string]string{
|
||||
return s.ReloadSecurity(context.WithoutCancel(ctx))
|
||||
}
|
||||
|
||||
func (s *SettingService) prepareSecurityUpdate(ctx context.Context, p SecurityPatch) (securityUpdate, error) {
|
||||
cur, err := s.Security(ctx)
|
||||
if err != nil {
|
||||
return securityUpdate{}, err
|
||||
}
|
||||
touched := p.merge(&cur)
|
||||
if err := validateSecurity(&cur); err != nil {
|
||||
return securityUpdate{}, err
|
||||
}
|
||||
values := map[string]string{
|
||||
settingSecLoginFailLimit: strconv.Itoa(cur.LoginFailLimit),
|
||||
settingSecLoginLockMin: strconv.Itoa(cur.LoginLockMinutes),
|
||||
settingSecIPRateRPS: strconv.Itoa(cur.IPRateRPS),
|
||||
@@ -165,13 +200,100 @@ func (s *SettingService) UpdateSecurity(ctx context.Context, p SecurityPatch) er
|
||||
settingSecRealIPHeader: cur.RealIPHeader,
|
||||
settingSecAppURL: cur.AppURL,
|
||||
}
|
||||
for key := range touched {
|
||||
if err := s.set(ctx, key, kv[key]); err != nil {
|
||||
return securityUpdate{next: cur, touched: touched, values: values}, nil
|
||||
}
|
||||
|
||||
func (s *SettingService) applySecurityUpdateTx(tx *gorm.DB, up securityUpdate, check *authenticatedMutation) error {
|
||||
if check != nil {
|
||||
if err := check.lockAndCheck(tx); err != nil {
|
||||
return err
|
||||
}
|
||||
} else if up.touched[settingSecAppURL] {
|
||||
if err := lockUsersForAuthChange(tx); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
// 重读而非直接 Store 合并值:并发补丁各写各键,重读保证快照收敛到库内最终值
|
||||
return s.ReloadSecurity(ctx)
|
||||
if err := s.checkAppURLUpdateTx(tx, up); err != nil {
|
||||
return err
|
||||
}
|
||||
for key := range up.touched {
|
||||
if err := saveSettingTx(tx, key, up.values[key]); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SettingService) checkAppURLUpdateTx(tx *gorm.DB, up securityUpdate) error {
|
||||
if !up.touched[settingSecAppURL] {
|
||||
return nil
|
||||
}
|
||||
oldApp, err := settingValueTx(tx, settingSecAppURL)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return s.ensureAppURLKeepsLogin(tx, oldApp, up.next.AppURL)
|
||||
}
|
||||
|
||||
// ensureAppURLKeepsLogin 防自锁:密码登录禁用期间,面板地址是所有免密方式的
|
||||
// 运行时依赖——清空(生效值)会使通行密钥/钱包/OAuth 回调全部不可用,一律拒绝;
|
||||
// 域名变更会使通行密钥失效(RP ID 绑定域名),须留有钱包身份或可用外部登录兜底。
|
||||
func (s *SettingService) ensureAppURLKeepsLogin(tx *gorm.DB, oldURL, newURL string) error {
|
||||
off, err := settingValueTx(tx, settingSecPasswordLoginOff)
|
||||
if err != nil || off != "1" {
|
||||
return err
|
||||
}
|
||||
newEff := s.effectiveOf(newURL)
|
||||
if newEff == "" {
|
||||
return ErrProviderLastLogin
|
||||
}
|
||||
// origin(scheme://host)整体比较:scheme 变化同样使 WebAuthn origin 失配
|
||||
if originOf(newEff) == originOf(s.effectiveOf(oldURL)) {
|
||||
return nil
|
||||
}
|
||||
if n, err := identityProviderCountTx(tx, "wallet"); err != nil || n > 0 {
|
||||
return err
|
||||
}
|
||||
view, err := oauthViewTx(tx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
ok, err := anyBoundUsableProviderTx(tx, patchedUsable(view, UpdateOAuthInput{}))
|
||||
if err != nil || ok {
|
||||
return err
|
||||
}
|
||||
return ErrProviderLastLogin
|
||||
}
|
||||
|
||||
// effectiveOf 计算给定 app_url 设置值的生效地址(设置优先,回退 PUBLIC_URL 环境变量)。
|
||||
func (s *SettingService) effectiveOf(v string) string {
|
||||
if v != "" {
|
||||
return v
|
||||
}
|
||||
return s.envPublicURL
|
||||
}
|
||||
|
||||
func effectiveOriginTx(tx *gorm.DB, settings *SettingService) (string, error) {
|
||||
if settings == nil {
|
||||
return "", nil
|
||||
}
|
||||
appURL, err := settingValueTx(tx, settingSecAppURL)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return originOf(settings.effectiveOf(appURL)), nil
|
||||
}
|
||||
|
||||
// originOf 取 URL 的 origin(scheme://host);空串或解析失败按原串返回。
|
||||
func originOf(rawURL string) string {
|
||||
if rawURL == "" {
|
||||
return ""
|
||||
}
|
||||
u, err := url.Parse(rawURL)
|
||||
if err != nil || u.Host == "" {
|
||||
return rawURL
|
||||
}
|
||||
return u.Scheme + "://" + u.Host
|
||||
}
|
||||
|
||||
// ReloadSecurity 从库加载安全设置到内存快照;进程启动时调用一次。
|
||||
|
||||
@@ -0,0 +1,433 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
|
||||
// ErrSessionCurrent 表示试图撤销当前会话;应引导用退出登录。
|
||||
var ErrSessionCurrent = errors.New("不能撤销当前会话,请使用退出登录")
|
||||
|
||||
const (
|
||||
// sessionSeenTTL 是最后活跃时间的回写节流窗口;窗口内命中缓存直接放行。
|
||||
sessionSeenTTL = time.Minute
|
||||
// sessionCleanupTick 是会话行清理周期。
|
||||
sessionCleanupTick = time.Hour
|
||||
// sessionExpiredKeep / sessionRevokedKeep 是失效行的保留期,过后物理删除。
|
||||
sessionExpiredKeep = 24 * time.Hour
|
||||
sessionRevokedKeep = 7 * 24 * time.Hour
|
||||
)
|
||||
|
||||
// jtiTombstones 是定点撤销的负缓存:无容量上限,不存在「淘汰导致已撤销
|
||||
// 令牌复活」;增长由 TTL 清理约束——撤销是认证后的低频人工操作,
|
||||
// 集合尺寸恒小,put 时线性清理过期项即可。
|
||||
type jtiTombstones struct {
|
||||
mu sync.RWMutex
|
||||
m map[string]time.Time // jti → 过期时刻
|
||||
}
|
||||
|
||||
func newJtiTombstones() *jtiTombstones {
|
||||
return &jtiTombstones{m: map[string]time.Time{}}
|
||||
}
|
||||
|
||||
// put 登记撤销标记并顺带清理过期项。
|
||||
func (t *jtiTombstones) put(jti string, ttl time.Duration) {
|
||||
if jti == "" || ttl <= 0 {
|
||||
return
|
||||
}
|
||||
now := time.Now()
|
||||
t.mu.Lock()
|
||||
defer t.mu.Unlock()
|
||||
for k, exp := range t.m {
|
||||
if now.After(exp) {
|
||||
delete(t.m, k)
|
||||
}
|
||||
}
|
||||
t.m[jti] = now.Add(ttl)
|
||||
}
|
||||
|
||||
// has 报告 jti 是否在有效撤销标记中。
|
||||
func (t *jtiTombstones) has(jti string) bool {
|
||||
t.mu.RLock()
|
||||
exp, ok := t.m[jti]
|
||||
t.mu.RUnlock()
|
||||
return ok && time.Now().Before(exp)
|
||||
}
|
||||
|
||||
// SessionMeta 是签发会话时的客户端上下文;Method 由各登录出口的 service 层填写,
|
||||
// api 层只采集 IP/UA。零值 meta 表示不落会话行(内部/测试场景)。
|
||||
type SessionMeta struct {
|
||||
ClientIP string
|
||||
UserAgent string
|
||||
Method string // password / oidc / github / passkey / wallet
|
||||
}
|
||||
|
||||
// empty 报告 meta 是否为「不落行」哨兵。
|
||||
func (m SessionMeta) empty() bool { return m.ClientIP == "" && m.UserAgent == "" }
|
||||
|
||||
// signSessionToken 签发 JWT 并按 meta 落地会话行。
|
||||
func (s *AuthService) signSessionToken(ctx context.Context, user *model.User, meta SessionMeta) (string, time.Time, error) {
|
||||
return s.signSessionTokenDB(s.db.WithContext(ctx), user, meta)
|
||||
}
|
||||
|
||||
// signSessionTokenTx 是身份登录事务内的签发入口。
|
||||
func (s *AuthService) signSessionTokenTx(tx *gorm.DB, user *model.User, meta SessionMeta) (string, time.Time, error) {
|
||||
return s.signSessionTokenDB(tx, user, meta)
|
||||
}
|
||||
|
||||
func (s *AuthService) signSessionTokenDB(db *gorm.DB, user *model.User, meta SessionMeta) (string, time.Time, error) {
|
||||
token, expires, jti, err := s.signToken(user.Username, user.TokenVersion)
|
||||
if err != nil || meta.empty() {
|
||||
return token, expires, err
|
||||
}
|
||||
now := time.Now()
|
||||
row := model.UserSession{
|
||||
UserID: user.ID, TokenID: jti, TokenVer: user.TokenVersion,
|
||||
Method: meta.Method, ClientIP: meta.ClientIP, UserAgent: meta.UserAgent,
|
||||
LastSeenAt: now, ExpiresAt: expires,
|
||||
}
|
||||
if err := db.Create(&row).Error; err != nil {
|
||||
// fail-closed:落行失败拒发令牌,否则产生列表不可见、无法定点撤销的孤儿会话
|
||||
return "", time.Time{}, fmt.Errorf("record session: %w", err)
|
||||
}
|
||||
return token, expires, nil
|
||||
}
|
||||
|
||||
// RenewToken 为敏感变更后的操作者换发新令牌:旧令牌对应的会话行接续
|
||||
// (沿用 jti,同行更新版本/有效期,保留登录方式与创建时间),无行则按 meta 新建。
|
||||
// 旧令牌只验签名不验有效性——版本刚被递增,旧令牌语义上已失效但行仍需接续。
|
||||
func (s *AuthService) RenewToken(ctx context.Context, username, oldToken string, meta SessionMeta) (string, time.Time, error) {
|
||||
proof, ok := s.signedTokenProof(oldToken)
|
||||
if !ok {
|
||||
return "", time.Time{}, ErrTokenStale
|
||||
}
|
||||
var token string
|
||||
var expires time.Time
|
||||
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if user.TokenVersion != proof.Ver+1 {
|
||||
return ErrTokenStale
|
||||
}
|
||||
token, expires, err = s.renewSessionTx(tx, user, oldToken, meta)
|
||||
return err
|
||||
})
|
||||
return token, expires, err
|
||||
}
|
||||
|
||||
// renewSessionTx 事务内签新令牌并接续会话行(RenewToken 的事务内核):
|
||||
// 与认证因子写入 / 版本递增同事务提交,绑定等敏感变更全程要么全成要么全滚;
|
||||
// 调用方须保证 user.TokenVersion 已是递增后的最新值。
|
||||
func (s *AuthService) renewSessionTx(tx *gorm.DB, user *model.User, oldToken string, meta SessionMeta) (string, time.Time, error) {
|
||||
jti := s.signedJti(oldToken)
|
||||
if jti == "" {
|
||||
jti = newTokenID()
|
||||
}
|
||||
token, expires, _, err := s.signTokenWithJTI(user.Username, user.TokenVersion, jti)
|
||||
if err != nil {
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
updates := map[string]any{
|
||||
"token_id": jti, "token_ver": user.TokenVersion,
|
||||
"expires_at": expires, "last_seen_at": time.Now(),
|
||||
}
|
||||
renewed, err := s.renewExistingSessionTx(tx, user.ID, oldToken, updates)
|
||||
if err != nil {
|
||||
return "", time.Time{}, err
|
||||
}
|
||||
if renewed {
|
||||
return token, expires, nil
|
||||
}
|
||||
if meta.empty() {
|
||||
return token, expires, nil
|
||||
}
|
||||
err = createRenewedSessionTx(tx, user, jti, expires, meta)
|
||||
return token, expires, err
|
||||
}
|
||||
|
||||
func createRenewedSessionTx(
|
||||
tx *gorm.DB, user *model.User, jti string, expires time.Time, meta SessionMeta,
|
||||
) error {
|
||||
row := model.UserSession{
|
||||
UserID: user.ID, TokenID: jti, TokenVer: user.TokenVersion,
|
||||
Method: meta.Method, ClientIP: meta.ClientIP, UserAgent: meta.UserAgent,
|
||||
LastSeenAt: time.Now(), ExpiresAt: expires,
|
||||
}
|
||||
if err := tx.Create(&row).Error; err != nil {
|
||||
return fmt.Errorf("record renewed session: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// renewExistingSessionTx 仅在旧行仍有效时接续;有行但已撤销或更新失败均失败关闭。
|
||||
func (s *AuthService) renewExistingSessionTx(tx *gorm.DB, userID uint, oldToken string, updates map[string]any) (bool, error) {
|
||||
oldJTI := s.signedJti(oldToken)
|
||||
if oldJTI == "" {
|
||||
return false, nil
|
||||
}
|
||||
if s.revokedJti.has(oldJTI) {
|
||||
return false, ErrTokenStale
|
||||
}
|
||||
if _, hit := s.revoked.Get(tokenHash(oldToken)); hit {
|
||||
return false, ErrTokenStale
|
||||
}
|
||||
res := tx.Model(&model.UserSession{}).
|
||||
Where("token_id = ? AND user_id = ? AND revoked_at IS NULL", oldJTI, userID).Updates(updates)
|
||||
if res.Error != nil {
|
||||
return false, fmt.Errorf("renew session: %w", res.Error)
|
||||
}
|
||||
if res.RowsAffected > 0 {
|
||||
s.seen.DeletePrefix("seen|" + oldJTI)
|
||||
return true, nil
|
||||
}
|
||||
return false, s.ensureRenewalHasNoOldRow(tx, userID, oldJTI)
|
||||
}
|
||||
|
||||
func (s *AuthService) ensureRenewalHasNoOldRow(tx *gorm.DB, userID uint, oldJTI string) error {
|
||||
var count int64
|
||||
err := tx.Model(&model.UserSession{}).
|
||||
Where("token_id = ? AND user_id = ?", oldJTI, userID).Count(&count).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("check renewed session: %w", err)
|
||||
}
|
||||
if count > 0 {
|
||||
return ErrTokenStale
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// signedJti 校验令牌签名并取出 jti;不验有效期与版本(换发场景旧令牌刚失效)。
|
||||
// 签名必须有效,防止伪造 jti 抢占他人会话行。
|
||||
func (s *AuthService) signedJti(tokenString string) string {
|
||||
proof, ok := s.signedTokenProof(tokenString)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
return proof.Jti
|
||||
}
|
||||
|
||||
func (s *AuthService) signedTokenProof(tokenString string) (TokenProof, bool) {
|
||||
if tokenString == "" {
|
||||
return TokenProof{}, false
|
||||
}
|
||||
claims := &authClaims{}
|
||||
_, err := jwt.ParseWithClaims(tokenString, claims, func(t *jwt.Token) (any, error) {
|
||||
if _, ok := t.Method.(*jwt.SigningMethodHMAC); !ok {
|
||||
return nil, fmt.Errorf("unexpected signing method %v", t.Header["alg"])
|
||||
}
|
||||
return s.jwtSecret, nil
|
||||
}, jwt.WithoutClaimsValidation())
|
||||
if err != nil {
|
||||
return TokenProof{}, false
|
||||
}
|
||||
return TokenProof{Ver: claims.Ver, Jti: claims.ID}, true
|
||||
}
|
||||
|
||||
// checkSession 校验 jti 对应会话未被定点撤销;无行放行(存量令牌兼容)。
|
||||
// 有效会话按节流窗口回写最后活跃时间;撤销动作会清掉节流缓存保证即时生效。
|
||||
func (s *AuthService) checkSession(ctx context.Context, jti string) error {
|
||||
if jti == "" {
|
||||
return nil
|
||||
}
|
||||
// 撤销负缓存优先:防「读库通过→撤销→回写 seen」竞态让已撤销令牌复活
|
||||
if s.revokedJti.has(jti) {
|
||||
return errors.New("session revoked")
|
||||
}
|
||||
if _, hit := s.seen.Get("seen|" + jti); hit {
|
||||
return nil
|
||||
}
|
||||
var row model.UserSession
|
||||
err := s.db.WithContext(ctx).Select("id", "revoked_at").Where("token_id = ?", jti).First(&row).Error
|
||||
if err != nil {
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("find session: %w", err)
|
||||
}
|
||||
if row.RevokedAt != nil {
|
||||
return errors.New("session revoked")
|
||||
}
|
||||
s.seen.Set("seen|"+jti, struct{}{}, sessionSeenTTL)
|
||||
s.db.WithContext(ctx).Model(&model.UserSession{}).
|
||||
Where("id = ?", row.ID).UpdateColumn("last_seen_at", time.Now())
|
||||
return nil
|
||||
}
|
||||
|
||||
// SessionInfo 是会话列表条目;Current 标记请求者自身会话。
|
||||
type SessionInfo struct {
|
||||
model.UserSession
|
||||
Current bool `json:"current"`
|
||||
}
|
||||
|
||||
// ListSessions 列出账号的活跃会话(未撤销、未过期、版本为当前),最近活跃在前。
|
||||
func (s *AuthService) ListSessions(ctx context.Context, username, currentToken string) ([]SessionInfo, error) {
|
||||
user, err := s.findUser(ctx, username)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rows := []model.UserSession{}
|
||||
err = s.db.WithContext(ctx).
|
||||
Where("user_id = ? AND revoked_at IS NULL AND expires_at > ? AND token_ver = ?",
|
||||
user.ID, time.Now(), user.TokenVersion).
|
||||
Order("last_seen_at DESC").Find(&rows).Error
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list sessions: %w", err)
|
||||
}
|
||||
currentJti := s.signedJti(currentToken)
|
||||
out := make([]SessionInfo, 0, len(rows))
|
||||
for _, r := range rows {
|
||||
out = append(out, SessionInfo{UserSession: r, Current: r.TokenID == currentJti})
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// RevokeSession 定点撤销一个会话(校验归属);当前会话拒绝(引导登出),
|
||||
// 不递增令牌版本、不影响其余会话。
|
||||
func (s *AuthService) RevokeSession(ctx context.Context, username, currentToken string, id uint) error {
|
||||
currentJTI := s.signedJti(currentToken)
|
||||
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
row, err := sessionForRevokeTx(tx, user.ID, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if row.TokenID == currentJTI {
|
||||
return ErrSessionCurrent
|
||||
}
|
||||
return s.revokeSessionRowTx(tx, row)
|
||||
})
|
||||
}
|
||||
|
||||
func sessionForRevokeTx(tx *gorm.DB, userID, id uint) (*model.UserSession, error) {
|
||||
var row model.UserSession
|
||||
err := tx.Where("id = ? AND user_id = ? AND revoked_at IS NULL", id, userID).First(&row).Error
|
||||
return &row, err
|
||||
}
|
||||
|
||||
func (s *AuthService) revokeSessionRowTx(tx *gorm.DB, row *model.UserSession) error {
|
||||
res := tx.Model(&model.UserSession{}).
|
||||
Where("id = ? AND revoked_at IS NULL", row.ID).UpdateColumn("revoked_at", time.Now())
|
||||
if res.Error != nil {
|
||||
return fmt.Errorf("revoke session: %w", res.Error)
|
||||
}
|
||||
if res.RowsAffected == 0 {
|
||||
return gorm.ErrRecordNotFound
|
||||
}
|
||||
s.markJTIRevoked(row.TokenID, time.Until(row.ExpiresAt))
|
||||
return nil
|
||||
}
|
||||
|
||||
// revokeSessionByJTI 按 jti 标记会话撤销(登出联动);无行为无害操作。
|
||||
func (s *AuthService) revokeSessionByJTI(ctx context.Context, username, jti string, ttl time.Duration) {
|
||||
s.markJTIRevoked(jti, logoutJTITTL(ttl))
|
||||
if jti == "" {
|
||||
return
|
||||
}
|
||||
userID := s.logoutUserID(ctx, username, jti)
|
||||
if userID == 0 {
|
||||
return
|
||||
}
|
||||
_ = s.revokeJTIForUser(ctx, userID, jti)
|
||||
}
|
||||
|
||||
// logoutJTITTL 覆盖旧令牌剩余窗口内最晚产生的同 JTI 换发令牌。
|
||||
func logoutJTITTL(ttl time.Duration) time.Duration { return ttl + tokenTTL }
|
||||
|
||||
func (s *AuthService) logoutUserID(ctx context.Context, username, jti string) uint {
|
||||
if userID := s.sessionUserID(ctx, jti); userID != 0 {
|
||||
return userID
|
||||
}
|
||||
if userID := s.usernameUserID(ctx, username); userID != 0 {
|
||||
return userID
|
||||
}
|
||||
// legacy 换发行可能正提交;再读一次缩小「无行→建行」窗口。
|
||||
return s.sessionUserID(ctx, jti)
|
||||
}
|
||||
|
||||
func (s *AuthService) sessionUserID(ctx context.Context, jti string) uint {
|
||||
var row model.UserSession
|
||||
err := s.db.WithContext(ctx).Select("user_id").
|
||||
Where("token_id = ?", jti).First(&row).Error
|
||||
if err != nil {
|
||||
return 0
|
||||
}
|
||||
return row.UserID
|
||||
}
|
||||
|
||||
func (s *AuthService) usernameUserID(ctx context.Context, username string) uint {
|
||||
if username == "" {
|
||||
return 0
|
||||
}
|
||||
var user model.User
|
||||
if err := s.db.WithContext(ctx).Select("id").
|
||||
Where("username = ?", username).First(&user).Error; err != nil {
|
||||
return 0
|
||||
}
|
||||
return user.ID
|
||||
}
|
||||
|
||||
func (s *AuthService) revokeJTIForUser(ctx context.Context, userID uint, jti string) error {
|
||||
return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
if _, err := lockUserByIDForAuthChange(tx, userID); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.Model(&model.UserSession{}).
|
||||
Where("token_id = ? AND user_id = ? AND revoked_at IS NULL", jti, userID).
|
||||
UpdateColumn("revoked_at", time.Now()).Error
|
||||
})
|
||||
}
|
||||
|
||||
func (s *AuthService) markJTIRevoked(jti string, ttl time.Duration) {
|
||||
s.seen.DeletePrefix("seen|" + jti)
|
||||
s.revokedJti.put(jti, ttl)
|
||||
}
|
||||
|
||||
// StartSessionCleanup 启动会话行周期清理:启动即清一次,之后每小时一次,
|
||||
// 随 ctx 取消退出;WaitSessionCleanup 可等待其真正结束(并发规范)。
|
||||
func (s *AuthService) StartSessionCleanup(ctx context.Context) {
|
||||
s.cleanupWG.Add(1)
|
||||
go func() {
|
||||
defer s.cleanupWG.Done()
|
||||
s.cleanupSessionsOnce(ctx)
|
||||
ticker := time.NewTicker(sessionCleanupTick)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return
|
||||
case <-ticker.C:
|
||||
s.cleanupSessionsOnce(ctx)
|
||||
}
|
||||
}
|
||||
}()
|
||||
}
|
||||
|
||||
// WaitSessionCleanup 阻塞等待清理 goroutine 退出(取消 ctx 后调用)。
|
||||
func (s *AuthService) WaitSessionCleanup() {
|
||||
s.cleanupWG.Wait()
|
||||
}
|
||||
|
||||
// cleanupSessionsOnce 删除保留期外的失效行;失败只记日志、不中断周期调度。
|
||||
func (s *AuthService) cleanupSessionsOnce(ctx context.Context) {
|
||||
now := time.Now()
|
||||
err := s.db.WithContext(ctx).
|
||||
Where("expires_at < ? OR revoked_at < ?", now.Add(-sessionExpiredKeep), now.Add(-sessionRevokedKeep)).
|
||||
Delete(&model.UserSession{}).Error
|
||||
if err != nil {
|
||||
log.Printf("session cleanup: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,525 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
|
||||
// loginSession 用密码登录建一条会话,返回令牌。
|
||||
func loginSession(t *testing.T, auth *AuthService, ip, ua string) string {
|
||||
t.Helper()
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "",
|
||||
SessionMeta{ClientIP: ip, UserAgent: ua})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
return token
|
||||
}
|
||||
|
||||
// mustProof 从有效令牌解析版本/jti 快照。
|
||||
func mustProof(t *testing.T, auth *AuthService, token string) TokenProof {
|
||||
t.Helper()
|
||||
_, proof, err := auth.ParseTokenProof(context.Background(), token)
|
||||
if err != nil {
|
||||
t.Fatalf("parse token proof: %v", err)
|
||||
}
|
||||
return proof
|
||||
}
|
||||
|
||||
func sessionRows(t *testing.T, auth *AuthService) []model.UserSession {
|
||||
t.Helper()
|
||||
rows := []model.UserSession{}
|
||||
if err := auth.db.Order("id").Find(&rows).Error; err != nil {
|
||||
t.Fatalf("load sessions: %v", err)
|
||||
}
|
||||
return rows
|
||||
}
|
||||
|
||||
func TestSessionRecordedOnLogin(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
loginSession(t, auth, "10.9.0.1", "TestAgent/1.0")
|
||||
rows := sessionRows(t, auth)
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("sessions = %d, want 1", len(rows))
|
||||
}
|
||||
r := rows[0]
|
||||
if r.Method != "password" || r.ClientIP != "10.9.0.1" || r.UserAgent != "TestAgent/1.0" {
|
||||
t.Errorf("row = %+v, want password/10.9.0.1/TestAgent", r)
|
||||
}
|
||||
if r.TokenID == "" || r.ExpiresAt.Before(time.Now()) {
|
||||
t.Errorf("row token/expiry invalid: %+v", r)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionRenewContinuity(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
token := loginSession(t, auth, "10.9.0.2", "UA")
|
||||
before := sessionRows(t, auth)[0]
|
||||
|
||||
// 敏感变更路径:版本递增 + 换发接续(RevokeSessions 即 bump+renew)
|
||||
fresh, _, err := auth.RevokeSessions(ctx, "admin", token, SessionMeta{ClientIP: "10.9.0.2", UserAgent: "UA"}, mustProof(t, auth, token))
|
||||
if err != nil {
|
||||
t.Fatalf("RevokeSessions: %v", err)
|
||||
}
|
||||
rows := sessionRows(t, auth)
|
||||
if len(rows) != 1 {
|
||||
t.Fatalf("sessions after renew = %d, want 1 (continuity, no new row)", len(rows))
|
||||
}
|
||||
after := rows[0]
|
||||
if after.ID != before.ID || after.TokenID != before.TokenID {
|
||||
t.Errorf("renew should keep row id and jti: before %+v after %+v", before, after)
|
||||
}
|
||||
if after.Method != "password" {
|
||||
t.Errorf("method = %q, want inherited password", after.Method)
|
||||
}
|
||||
if _, err := auth.ParseToken(ctx, token); err == nil {
|
||||
t.Error("old token still valid after version bump")
|
||||
}
|
||||
if _, err := auth.ParseToken(ctx, fresh); err != nil {
|
||||
t.Errorf("fresh token invalid: %v", err)
|
||||
}
|
||||
// 列表只剩接续会话且标记 current
|
||||
list, err := auth.ListSessions(ctx, "admin", fresh)
|
||||
if err != nil || len(list) != 1 || !list[0].Current {
|
||||
t.Errorf("ListSessions = (%+v, %v), want single current session", list, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionRevokeSingle(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
token1 := loginSession(t, auth, "10.9.0.3", "Laptop")
|
||||
token2 := loginSession(t, auth, "10.9.0.4", "Phone")
|
||||
// token2 已被 ParseToken 校验过(节流缓存生效)后再撤销,验证缓存被清、即时失效
|
||||
if _, err := auth.ParseToken(ctx, token2); err != nil {
|
||||
t.Fatalf("token2 parse: %v", err)
|
||||
}
|
||||
list, err := auth.ListSessions(ctx, "admin", token1)
|
||||
if err != nil || len(list) != 2 {
|
||||
t.Fatalf("ListSessions = (%d, %v), want 2", len(list), err)
|
||||
}
|
||||
var otherID uint
|
||||
for _, it := range list {
|
||||
if it.Current {
|
||||
continue
|
||||
}
|
||||
otherID = it.ID
|
||||
}
|
||||
tests := []struct {
|
||||
name string
|
||||
id uint
|
||||
wantErr error
|
||||
}{
|
||||
{name: "撤销当前会话被拒", id: currentSessionID(t, list), wantErr: ErrSessionCurrent},
|
||||
{name: "撤销其他会话成功", id: otherID},
|
||||
{name: "重复撤销已不可见", id: otherID, wantErr: gorm.ErrRecordNotFound},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := auth.RevokeSession(ctx, "admin", token1, tt.id)
|
||||
if tt.wantErr != nil {
|
||||
if !errors.Is(err, tt.wantErr) && !(tt.wantErr == gorm.ErrRecordNotFound && err != nil) {
|
||||
t.Fatalf("RevokeSession err = %v, want %v", err, tt.wantErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("RevokeSession: %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
if _, err := auth.ParseToken(ctx, token2); err == nil {
|
||||
t.Error("revoked session token still valid")
|
||||
}
|
||||
if _, err := auth.ParseToken(ctx, token1); err != nil {
|
||||
t.Errorf("current token broken by revoking another session: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionRevokeTombstone(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
token1 := loginSession(t, auth, "10.9.0.7", "Laptop")
|
||||
token2 := loginSession(t, auth, "10.9.0.8", "Phone")
|
||||
list, err := auth.ListSessions(ctx, "admin", token1)
|
||||
if err != nil || len(list) != 2 {
|
||||
t.Fatalf("ListSessions = (%d, %v), want 2", len(list), err)
|
||||
}
|
||||
var otherID uint
|
||||
for _, it := range list {
|
||||
if !it.Current {
|
||||
otherID = it.ID
|
||||
}
|
||||
}
|
||||
if err := auth.RevokeSession(ctx, "admin", token1, otherID); err != nil {
|
||||
t.Fatalf("RevokeSession: %v", err)
|
||||
}
|
||||
// 模拟「校验读库通过→撤销→校验回写 seen」的竞态:手工回写节流缓存,
|
||||
// 撤销负缓存必须仍然盖过它,令牌不得复活
|
||||
jti := auth.signedJti(token2)
|
||||
auth.seen.Set("seen|"+jti, struct{}{}, sessionSeenTTL)
|
||||
if _, err := auth.ParseToken(ctx, token2); err == nil {
|
||||
t.Error("revoked token revived by racing seen-cache write")
|
||||
}
|
||||
}
|
||||
|
||||
// TestSensitiveOpStaleProof 验证在途绕过防线:敏感请求鉴权后挂起,期间发生
|
||||
// 「撤销全部」(版本递增),恢复后的敏感事务复核快照失败,拒绝执行、不发新令牌。
|
||||
// 同一防线也使并发敏感操作串行化(后到者复核失败)。
|
||||
func TestSensitiveOpStaleProof(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
tokenA := loginSession(t, auth, "10.9.1.1", "A")
|
||||
proofA := mustProof(t, auth, tokenA)
|
||||
// 另一设备撤销全部:版本递增,tokenA 的快照随之过期
|
||||
tokenB := loginSession(t, auth, "10.9.1.2", "B")
|
||||
if _, _, err := auth.RevokeSessions(ctx, "admin", tokenB, SessionMeta{ClientIP: "10.9.1.2"}, mustProof(t, auth, tokenB)); err != nil {
|
||||
t.Fatalf("RevokeSessions: %v", err)
|
||||
}
|
||||
// 挂起的旧请求恢复:携带过期快照的敏感操作必须被拒
|
||||
if _, _, err := auth.RevokeSessions(ctx, "admin", tokenA, SessionMeta{ClientIP: "10.9.1.1"}, proofA); !errors.Is(err, ErrTokenStale) {
|
||||
t.Fatalf("stale-proof revoke err = %v, want ErrTokenStale", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestSensitiveOpRevokedJtiProof 验证快照的 jti 维度:定点撤销(不递增版本)
|
||||
// 同样令该令牌的在途敏感请求失效。
|
||||
func TestSensitiveOpRevokedJtiProof(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
token1 := loginSession(t, auth, "10.9.2.1", "Laptop")
|
||||
token2 := loginSession(t, auth, "10.9.2.2", "Phone")
|
||||
proof2 := mustProof(t, auth, token2)
|
||||
list, err := auth.ListSessions(ctx, "admin", token1)
|
||||
if err != nil || len(list) != 2 {
|
||||
t.Fatalf("ListSessions = (%d, %v), want 2", len(list), err)
|
||||
}
|
||||
for _, it := range list {
|
||||
if !it.Current {
|
||||
if err := auth.RevokeSession(ctx, "admin", token1, it.ID); err != nil {
|
||||
t.Fatalf("RevokeSession: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
// token2 已被定点撤销(版本未变):其在途敏感请求恢复后必须被拒
|
||||
if _, _, err := auth.RevokeSessions(ctx, "admin", token2, SessionMeta{ClientIP: "10.9.2.2"}, proof2); !errors.Is(err, ErrTokenStale) {
|
||||
t.Fatalf("revoked-jti proof err = %v, want ErrTokenStale", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckedSessionCannotRenewAfterTargetedRevoke(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
stale := loginSession(t, auth, "10.9.3.1", "stale")
|
||||
current := loginSession(t, auth, "10.9.3.2", "current")
|
||||
proof := mustProof(t, auth, stale)
|
||||
assertProofCurrentTx(t, auth, proof)
|
||||
staleID := otherSessionID(t, auth, current)
|
||||
if err := auth.bumpTokenVersion(ctx, "admin"); err != nil {
|
||||
t.Fatalf("simulate sensitive mutation: %v", err)
|
||||
}
|
||||
if err := auth.RevokeSession(ctx, "admin", current, staleID); err != nil {
|
||||
t.Fatalf("RevokeSession: %v", err)
|
||||
}
|
||||
if _, _, err := auth.RenewToken(
|
||||
ctx, "admin", stale, SessionMeta{ClientIP: "10.9.3.1"},
|
||||
); !errors.Is(err, ErrTokenStale) {
|
||||
t.Fatalf("RenewToken err = %v, want ErrTokenStale", err)
|
||||
}
|
||||
if got := len(sessionRows(t, auth)); got != 2 {
|
||||
t.Fatalf("session rows = %d, want 2 (不得 fallback CREATE)", got)
|
||||
}
|
||||
}
|
||||
|
||||
func otherSessionID(t *testing.T, auth *AuthService, current string) uint {
|
||||
t.Helper()
|
||||
for _, item := range mustSessions(t, auth, current) {
|
||||
if !item.Current {
|
||||
return item.ID
|
||||
}
|
||||
}
|
||||
t.Fatal("no other session")
|
||||
return 0
|
||||
}
|
||||
|
||||
func TestLegacyLogoutBlocksInflightProofAndRenew(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
token, _, _, err := auth.signToken("admin", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("signToken: %v", err)
|
||||
}
|
||||
proof := mustProof(t, auth, token)
|
||||
assertProofCurrentTx(t, auth, proof)
|
||||
auth.Logout(ctx, token)
|
||||
assertProofStaleTx(t, auth, proof)
|
||||
if err := auth.bumpTokenVersion(ctx, "admin"); err != nil {
|
||||
t.Fatalf("simulate already committed mutation: %v", err)
|
||||
}
|
||||
if _, _, err := auth.RenewToken(
|
||||
ctx, "admin", token, SessionMeta{ClientIP: "10.9.4.1"},
|
||||
); !errors.Is(err, ErrTokenStale) {
|
||||
t.Fatalf("RenewToken err = %v, want ErrTokenStale", err)
|
||||
}
|
||||
if got := len(sessionRows(t, auth)); got != 0 {
|
||||
t.Fatalf("session rows = %d, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
type logoutAfterRenameCase struct {
|
||||
name string
|
||||
legacy bool
|
||||
meta SessionMeta
|
||||
}
|
||||
|
||||
func TestLogoutOldTokenAfterCredentialRename(t *testing.T) {
|
||||
tests := []logoutAfterRenameCase{
|
||||
{name: "recorded session", meta: SessionMeta{ClientIP: "10.9.5.1", UserAgent: "recorded"}},
|
||||
{name: "legacy session without row", legacy: true},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
testLogoutOldTokenAfterCredentialRename(t, tt)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func testLogoutOldTokenAfterCredentialRename(t *testing.T, tt logoutAfterRenameCase) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
old, oldExpires := tokenForRenewalTest(t, auth, tt.legacy)
|
||||
fresh, expires := renewAfterCredentialRename(t, auth, old, tt.meta)
|
||||
if auth.signedJti(old) != auth.signedJti(fresh) {
|
||||
t.Fatal("renewal changed jti; logout lineage would be lost")
|
||||
}
|
||||
if tt.legacy && len(sessionRows(t, auth)) != 0 {
|
||||
t.Fatal("legacy renewal unexpectedly created a session row")
|
||||
}
|
||||
auth.Logout(context.Background(), old)
|
||||
if _, err := auth.ParseToken(context.Background(), fresh); err == nil {
|
||||
t.Fatal("fresh token survived logout of its pre-renewal token")
|
||||
}
|
||||
if tt.legacy {
|
||||
assertTombstoneCovers(t, auth, old, oldExpires, expires)
|
||||
return
|
||||
}
|
||||
assertSessionJTIRevoked(t, auth, old)
|
||||
}
|
||||
|
||||
func renewAfterCredentialRename(
|
||||
t *testing.T, auth *AuthService, old string, meta SessionMeta,
|
||||
) (string, time.Time) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
finalName, err := auth.UpdateCredentials(ctx, "admin", UpdateCredentialsInput{
|
||||
NewUsername: "root", CurrentPassword: "pass123",
|
||||
}, mustProof(t, auth, old))
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateCredentials: %v", err)
|
||||
}
|
||||
fresh, expires, err := auth.RenewToken(ctx, finalName, old, meta)
|
||||
if err != nil {
|
||||
t.Fatalf("RenewToken: %v", err)
|
||||
}
|
||||
return fresh, expires
|
||||
}
|
||||
|
||||
func assertSessionJTIRevoked(t *testing.T, auth *AuthService, token string) {
|
||||
t.Helper()
|
||||
var row model.UserSession
|
||||
err := auth.db.Where("token_id = ?", auth.signedJti(token)).First(&row).Error
|
||||
if err != nil {
|
||||
t.Fatalf("find session: %v", err)
|
||||
}
|
||||
if row.RevokedAt == nil {
|
||||
t.Fatal("renamed user's session row was not revoked")
|
||||
}
|
||||
}
|
||||
|
||||
func assertTombstoneCovers(
|
||||
t *testing.T, auth *AuthService, token string, oldExpires, freshExpires time.Time,
|
||||
) {
|
||||
t.Helper()
|
||||
jti := auth.signedJti(token)
|
||||
auth.revokedJti.mu.RLock()
|
||||
tombstoneExpires, ok := auth.revokedJti.m[jti]
|
||||
auth.revokedJti.mu.RUnlock()
|
||||
if !ok {
|
||||
t.Fatal("logout jti tombstone missing")
|
||||
}
|
||||
requiredUntil := oldExpires.Truncate(time.Second).Add(tokenTTL)
|
||||
if tombstoneExpires.Before(requiredUntil) || tombstoneExpires.Before(freshExpires) {
|
||||
t.Fatalf("tombstone expires %v before required horizon %v", tombstoneExpires, requiredUntil)
|
||||
}
|
||||
}
|
||||
|
||||
func tokenForRenewalTest(t *testing.T, auth *AuthService, legacy bool) (string, time.Time) {
|
||||
t.Helper()
|
||||
if !legacy {
|
||||
token, expires, err := auth.Login(context.Background(), "admin", "pass123", "",
|
||||
SessionMeta{ClientIP: "10.9.5.1", UserAgent: "recorded"})
|
||||
if err != nil {
|
||||
t.Fatalf("Login: %v", err)
|
||||
}
|
||||
return token, expires
|
||||
}
|
||||
token, expires, _, err := auth.signToken("admin", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("signToken: %v", err)
|
||||
}
|
||||
return token, expires
|
||||
}
|
||||
|
||||
func assertProofCurrentTx(t *testing.T, auth *AuthService, proof TokenProof) {
|
||||
t.Helper()
|
||||
err := auth.db.Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, "admin")
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return auth.ensureTokenCurrentTx(tx, user, proof)
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("proof should be current: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func assertProofStaleTx(t *testing.T, auth *AuthService, proof TokenProof) {
|
||||
t.Helper()
|
||||
err := auth.db.Transaction(func(tx *gorm.DB) error {
|
||||
user, lockErr := lockUserForAuthChange(tx, "admin")
|
||||
if lockErr != nil {
|
||||
return lockErr
|
||||
}
|
||||
return auth.ensureTokenCurrentTx(tx, user, proof)
|
||||
})
|
||||
if !errors.Is(err, ErrTokenStale) {
|
||||
t.Fatalf("proof err = %v, want ErrTokenStale", err)
|
||||
}
|
||||
}
|
||||
|
||||
func mustSessions(t *testing.T, auth *AuthService, current string) []SessionInfo {
|
||||
t.Helper()
|
||||
items, err := auth.ListSessions(context.Background(), "admin", current)
|
||||
if err != nil {
|
||||
t.Fatalf("ListSessions: %v", err)
|
||||
}
|
||||
return items
|
||||
}
|
||||
|
||||
// TestTombstoneNoEviction 验证撤销负缓存无容量上限:大量撤销标记全部存活,
|
||||
// 不存在「满载淘汰导致已撤销令牌复活」。
|
||||
func TestTombstoneNoEviction(t *testing.T) {
|
||||
ts := newJtiTombstones()
|
||||
for i := 0; i < 600; i++ {
|
||||
ts.put(fmt.Sprintf("jti-%d", i), time.Minute)
|
||||
}
|
||||
for i := 0; i < 600; i++ {
|
||||
if !ts.has(fmt.Sprintf("jti-%d", i)) {
|
||||
t.Fatalf("tombstone jti-%d evicted", i)
|
||||
}
|
||||
}
|
||||
if ts.has("jti-none") {
|
||||
t.Error("unknown jti reported revoked")
|
||||
}
|
||||
}
|
||||
|
||||
// currentSessionID 取列表中 current 条目的 ID。
|
||||
func currentSessionID(t *testing.T, list []SessionInfo) uint {
|
||||
t.Helper()
|
||||
for _, it := range list {
|
||||
if it.Current {
|
||||
return it.ID
|
||||
}
|
||||
}
|
||||
t.Fatal("no current session in list")
|
||||
return 0
|
||||
}
|
||||
|
||||
func TestSessionLegacyTokenAllowed(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
// 直签令牌(无会话行)模拟升级前存量令牌:仍可通过校验
|
||||
token, _, _, err := auth.signToken("admin", 0)
|
||||
if err != nil {
|
||||
t.Fatalf("signToken: %v", err)
|
||||
}
|
||||
if _, err := auth.ParseToken(context.Background(), token); err != nil {
|
||||
t.Errorf("legacy token rejected: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionLogoutMarksRow(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
token := loginSession(t, auth, "10.9.0.5", "UA")
|
||||
auth.Logout(ctx, token)
|
||||
rows := sessionRows(t, auth)
|
||||
if len(rows) != 1 || rows[0].RevokedAt == nil {
|
||||
t.Errorf("logout should mark session revoked: %+v", rows)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSessionCleanup(t *testing.T) {
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
now := time.Now()
|
||||
old := now.Add(-8 * 24 * time.Hour)
|
||||
seed := []model.UserSession{
|
||||
{UserID: 1, TokenID: "expired", ExpiresAt: now.Add(-25 * time.Hour), LastSeenAt: old},
|
||||
{UserID: 1, TokenID: "revoked-old", ExpiresAt: now.Add(time.Hour), RevokedAt: &old, LastSeenAt: old},
|
||||
{UserID: 1, TokenID: "alive", ExpiresAt: now.Add(time.Hour), LastSeenAt: now},
|
||||
}
|
||||
for i := range seed {
|
||||
if err := auth.db.Create(&seed[i]).Error; err != nil {
|
||||
t.Fatalf("seed: %v", err)
|
||||
}
|
||||
}
|
||||
auth.cleanupSessionsOnce(ctx)
|
||||
rows := sessionRows(t, auth)
|
||||
if len(rows) != 1 || rows[0].TokenID != "alive" {
|
||||
t.Errorf("after cleanup rows = %+v, want only alive", rows)
|
||||
}
|
||||
}
|
||||
@@ -29,7 +29,7 @@ func newSettingEnv(t *testing.T) (*SettingService, *gorm.DB) {
|
||||
t.Fatalf("db handle: %v", err)
|
||||
}
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
if err := db.AutoMigrate(&model.Setting{}); err != nil {
|
||||
if err := db.AutoMigrate(&model.Setting{}, &model.User{}, &model.UserIdentity{}, &model.UserPasskey{}); err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
cipher, err := crypto.NewCipher("test-data-key")
|
||||
|
||||
+66
-28
@@ -8,6 +8,7 @@ import (
|
||||
|
||||
"github.com/pquerna/otp/totp"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
@@ -79,52 +80,89 @@ func (s *AuthService) SetupTotp(ctx context.Context, username string) (secret, u
|
||||
return key.Secret(), key.URL(), nil
|
||||
}
|
||||
|
||||
// ActivateTotp 校验暂存密钥的验证码,通过后加密落库启用。
|
||||
func (s *AuthService) ActivateTotp(ctx context.Context, username, code string) error {
|
||||
// ActivateTotp 校验暂存密钥并在同一事务内启用、递增版本、接续当前会话。
|
||||
func (s *AuthService) ActivateTotp(
|
||||
ctx context.Context, username, code, oldToken string, meta SessionMeta, proof TokenProof,
|
||||
) (string, time.Time, error) {
|
||||
s.totpMu.Lock()
|
||||
pending, ok := s.totpPending[username]
|
||||
s.totpMu.Unlock()
|
||||
if !ok || time.Now().After(pending.expires) {
|
||||
return ErrTotpNotSetup
|
||||
return "", time.Time{}, ErrTotpNotSetup
|
||||
}
|
||||
if !totp.Validate(code, pending.secret) {
|
||||
return ErrTotpInvalid
|
||||
return "", time.Time{}, ErrTotpInvalid
|
||||
}
|
||||
enc, err := s.cipher.EncryptString(pending.secret)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encrypt totp secret: %w", err)
|
||||
return "", time.Time{}, fmt.Errorf("encrypt totp secret: %w", err)
|
||||
}
|
||||
err = s.db.WithContext(ctx).Model(&model.User{}).
|
||||
Where("username = ?", username).Update("totp_secret_enc", enc).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("save totp secret: %w", err)
|
||||
token, expires, err := s.changeTotp(ctx, username, enc, oldToken, meta, proof, nil)
|
||||
if err == nil {
|
||||
s.deleteTotpPending(username, pending.secret)
|
||||
}
|
||||
s.totpMu.Lock()
|
||||
delete(s.totpPending, username)
|
||||
s.totpMu.Unlock()
|
||||
// 两步验证形态变更:旧令牌全部失效
|
||||
return s.bumpTokenVersion(ctx, username)
|
||||
return token, expires, err
|
||||
}
|
||||
|
||||
// DisableTotp 停用两步验证;需当前验证码或登录密码任一确认。
|
||||
func (s *AuthService) DisableTotp(ctx context.Context, username, password, code string) error {
|
||||
user, err := s.findUser(ctx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if user.TotpSecretEnc == "" {
|
||||
func (s *AuthService) DisableTotp(
|
||||
ctx context.Context, username, password, code, oldToken string, meta SessionMeta, proof TokenProof,
|
||||
) (string, time.Time, error) {
|
||||
confirm := func(user *model.User) error {
|
||||
if user.TotpSecretEnc != "" && !s.confirmDisable(user, password, code) {
|
||||
return ErrTotpConfirm
|
||||
}
|
||||
return nil
|
||||
}
|
||||
if !s.confirmDisable(user, password, code) {
|
||||
return ErrTotpConfirm
|
||||
return s.changeTotp(ctx, username, "", oldToken, meta, proof, confirm)
|
||||
}
|
||||
|
||||
func (s *AuthService) changeTotp(
|
||||
ctx context.Context, username, secret, oldToken string, meta SessionMeta,
|
||||
proof TokenProof, confirm func(*model.User) error,
|
||||
) (string, time.Time, error) {
|
||||
var token string
|
||||
var expires time.Time
|
||||
err := s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := s.ensureTokenCurrentTx(tx, user, proof); err != nil {
|
||||
return err
|
||||
}
|
||||
if confirm != nil {
|
||||
if err := confirm(user); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return s.persistTotpTx(tx, user, secret, oldToken, meta, &token, &expires)
|
||||
})
|
||||
return token, expires, err
|
||||
}
|
||||
|
||||
func (s *AuthService) persistTotpTx(
|
||||
tx *gorm.DB, user *model.User, secret, oldToken string, meta SessionMeta,
|
||||
token *string, expires *time.Time,
|
||||
) error {
|
||||
if err := tx.Model(user).Update("totp_secret_enc", secret).Error; err != nil {
|
||||
return fmt.Errorf("save totp secret: %w", err)
|
||||
}
|
||||
err = s.db.WithContext(ctx).Model(&model.User{}).
|
||||
Where("username = ?", username).Update("totp_secret_enc", "").Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("clear totp secret: %w", err)
|
||||
if err := bumpTokenVersionTx(tx, user.Username); err != nil {
|
||||
return err
|
||||
}
|
||||
user.TokenVersion++
|
||||
var err error
|
||||
*token, *expires, err = s.renewSessionTx(tx, user, oldToken, meta)
|
||||
return err
|
||||
}
|
||||
|
||||
func (s *AuthService) deleteTotpPending(username, secret string) {
|
||||
s.totpMu.Lock()
|
||||
defer s.totpMu.Unlock()
|
||||
if p, ok := s.totpPending[username]; ok && p.secret == secret {
|
||||
delete(s.totpPending, username)
|
||||
}
|
||||
// 两步验证形态变更:旧令牌全部失效
|
||||
return s.bumpTokenVersion(ctx, username)
|
||||
}
|
||||
|
||||
// confirmDisable 校验停用凭证:验证码或密码任一通过即可。
|
||||
|
||||
@@ -29,7 +29,7 @@ func newTotpEnv(t *testing.T) (*AuthService, *gorm.DB) {
|
||||
t.Fatalf("db handle: %v", err)
|
||||
}
|
||||
sqlDB.SetMaxOpenConns(1)
|
||||
if err := db.AutoMigrate(&model.User{}, &model.UserIdentity{}, &model.Setting{}); err != nil {
|
||||
if err := db.AutoMigrate(&model.User{}, &model.UserIdentity{}, &model.UserPasskey{}, &model.UserSession{}, &model.Setting{}); err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
auth := NewAuthService(db, "test-secret")
|
||||
@@ -44,10 +44,11 @@ func newTotpEnv(t *testing.T) (*AuthService, *gorm.DB) {
|
||||
return auth, db
|
||||
}
|
||||
|
||||
// enableTotp 走完整 setup→activate 流程,返回明文密钥供测试生成验证码。
|
||||
func enableTotp(t *testing.T, auth *AuthService) string {
|
||||
// enableTotp 走完整 login→setup→activate 流程,返回明文密钥与接续令牌。
|
||||
func enableTotp(t *testing.T, auth *AuthService) (string, string) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
oldToken := loginSession(t, auth, "127.0.0.1", "totp-test")
|
||||
secret, uri, err := auth.SetupTotp(ctx, "admin")
|
||||
if err != nil {
|
||||
t.Fatalf("SetupTotp: %v", err)
|
||||
@@ -59,10 +60,12 @@ func enableTotp(t *testing.T, auth *AuthService) string {
|
||||
if err != nil {
|
||||
t.Fatalf("generate code: %v", err)
|
||||
}
|
||||
if err := auth.ActivateTotp(ctx, "admin", code); err != nil {
|
||||
token, _, err := auth.ActivateTotp(
|
||||
ctx, "admin", code, oldToken, SessionMeta{ClientIP: "127.0.0.1"}, mustProof(t, auth, oldToken))
|
||||
if err != nil {
|
||||
t.Fatalf("ActivateTotp: %v", err)
|
||||
}
|
||||
return secret
|
||||
return secret, token
|
||||
}
|
||||
|
||||
func TestTotpLifecycle(t *testing.T) {
|
||||
@@ -72,7 +75,7 @@ func TestTotpLifecycle(t *testing.T) {
|
||||
if on, _ := auth.TotpStatus(ctx, "admin"); on {
|
||||
t.Fatal("初始不应启用")
|
||||
}
|
||||
secret := enableTotp(t, auth)
|
||||
secret, token := enableTotp(t, auth)
|
||||
if on, _ := auth.TotpStatus(ctx, "admin"); !on {
|
||||
t.Fatal("激活后应为启用")
|
||||
}
|
||||
@@ -89,10 +92,12 @@ func TestTotpLifecycle(t *testing.T) {
|
||||
t.Errorf("重复 setup err = %v, want ErrTotpAlreadyOn", err)
|
||||
}
|
||||
// 停用:无凭证拒绝,密码通过
|
||||
if err := auth.DisableTotp(ctx, "admin", "", ""); !errors.Is(err, ErrTotpConfirm) {
|
||||
if _, _, err := auth.DisableTotp(
|
||||
ctx, "admin", "", "", token, SessionMeta{}, mustProof(t, auth, token)); !errors.Is(err, ErrTotpConfirm) {
|
||||
t.Errorf("空凭证停用 err = %v, want ErrTotpConfirm", err)
|
||||
}
|
||||
if err := auth.DisableTotp(ctx, "admin", "pass123", ""); err != nil {
|
||||
if _, _, err := auth.DisableTotp(
|
||||
ctx, "admin", "pass123", "", token, SessionMeta{}, mustProof(t, auth, token)); err != nil {
|
||||
t.Fatalf("密码停用: %v", err)
|
||||
}
|
||||
if on, _ := auth.TotpStatus(ctx, "admin"); on {
|
||||
@@ -104,29 +109,59 @@ func TestActivateTotpRejects(t *testing.T) {
|
||||
auth, _ := newTotpEnv(t)
|
||||
ctx := context.Background()
|
||||
// 未 setup 直接激活
|
||||
if err := auth.ActivateTotp(ctx, "admin", "123456"); !errors.Is(err, ErrTotpNotSetup) {
|
||||
if _, _, err := auth.ActivateTotp(
|
||||
ctx, "admin", "123456", "", SessionMeta{}, TokenProof{}); !errors.Is(err, ErrTotpNotSetup) {
|
||||
t.Errorf("err = %v, want ErrTotpNotSetup", err)
|
||||
}
|
||||
// setup 后错误验证码
|
||||
if _, _, err := auth.SetupTotp(ctx, "admin"); err != nil {
|
||||
t.Fatalf("SetupTotp: %v", err)
|
||||
}
|
||||
if err := auth.ActivateTotp(ctx, "admin", "000000"); !errors.Is(err, ErrTotpInvalid) {
|
||||
if _, _, err := auth.ActivateTotp(
|
||||
ctx, "admin", "000000", "", SessionMeta{}, TokenProof{}); !errors.Is(err, ErrTotpInvalid) {
|
||||
t.Errorf("err = %v, want ErrTotpInvalid", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestActivateTotpRejectsRevokedProof(t *testing.T) {
|
||||
auth, _ := newTotpEnv(t)
|
||||
ctx := context.Background()
|
||||
staleToken := loginSession(t, auth, "127.0.0.1", "stale")
|
||||
staleProof := mustProof(t, auth, staleToken)
|
||||
secret, _, err := auth.SetupTotp(ctx, "admin")
|
||||
if err != nil {
|
||||
t.Fatalf("SetupTotp: %v", err)
|
||||
}
|
||||
code, err := totp.GenerateCode(secret, time.Now())
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateCode: %v", err)
|
||||
}
|
||||
current := loginSession(t, auth, "127.0.0.2", "current")
|
||||
if _, _, err := auth.RevokeSessions(
|
||||
ctx, "admin", current, SessionMeta{ClientIP: "127.0.0.2"}, mustProof(t, auth, current)); err != nil {
|
||||
t.Fatalf("RevokeSessions: %v", err)
|
||||
}
|
||||
if _, _, err := auth.ActivateTotp(
|
||||
ctx, "admin", code, staleToken, SessionMeta{ClientIP: "127.0.0.1"}, staleProof,
|
||||
); !errors.Is(err, ErrTokenStale) {
|
||||
t.Fatalf("ActivateTotp err = %v, want ErrTokenStale", err)
|
||||
}
|
||||
if on, err := auth.TotpStatus(ctx, "admin"); err != nil || on {
|
||||
t.Fatalf("TotpStatus = %v, %v; want disabled", on, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginWithTotp(t *testing.T) {
|
||||
auth, _ := newTotpEnv(t)
|
||||
ctx := context.Background()
|
||||
secret := enableTotp(t, auth)
|
||||
secret, _ := enableTotp(t, auth)
|
||||
|
||||
// 缺验证码:密码对也返回 ErrTotpRequired(不计失败)
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", ""); !errors.Is(err, ErrTotpRequired) {
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "127.0.0.1"}); !errors.Is(err, ErrTotpRequired) {
|
||||
t.Fatalf("err = %v, want ErrTotpRequired", err)
|
||||
}
|
||||
// 错误验证码:按失败处理
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", "000000"); !errors.Is(err, ErrInvalidCredentials) {
|
||||
if _, _, err := auth.Login(ctx, "admin", "pass123", "000000", SessionMeta{ClientIP: "127.0.0.1"}); !errors.Is(err, ErrInvalidCredentials) {
|
||||
t.Fatalf("err = %v, want ErrInvalidCredentials", err)
|
||||
}
|
||||
// 正确验证码:登录成功
|
||||
@@ -134,7 +169,7 @@ func TestLoginWithTotp(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("generate code: %v", err)
|
||||
}
|
||||
token, _, err := auth.Login(ctx, "admin", "pass123", "127.0.0.1", code)
|
||||
token, _, err := auth.Login(ctx, "admin", "pass123", code, SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil || token == "" {
|
||||
t.Fatalf("带验证码登录失败: %v", err)
|
||||
}
|
||||
@@ -143,11 +178,200 @@ func TestLoginWithTotp(t *testing.T) {
|
||||
// fakeBindIdentity 直接把外部身份写入待测服务(绕过真实 OAuth flow)。
|
||||
func fakeBindIdentity(t *testing.T, o *OAuthService, username, provider, subject, display string) {
|
||||
t.Helper()
|
||||
if err := o.bind(context.Background(), username, provider, externalIdentity{Subject: subject, Display: display}); err != nil {
|
||||
p := fakePending(t, o, username, provider)
|
||||
if _, err := o.bind(context.Background(), p, externalIdentity{Subject: subject, Display: display}, SessionMeta{}); err != nil {
|
||||
t.Fatalf("bind: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// fakePending 构造与当前令牌版本一致的 bind 流程上下文(跳过外部授权码交换)。
|
||||
func fakePending(t *testing.T, o *OAuthService, username, provider string) oauthPending {
|
||||
t.Helper()
|
||||
var user model.User
|
||||
if err := o.db.Where("username = ?", username).First(&user).Error; err != nil {
|
||||
t.Fatalf("find user: %v", err)
|
||||
}
|
||||
return oauthPending{provider: provider, mode: "bind", username: username, proof: TokenProof{Ver: user.TokenVersion}}
|
||||
}
|
||||
|
||||
// TestOAuthBindStaleToken 验证 bind 回调复验发起时令牌:撤销全部(版本递增)后,
|
||||
// 已登记的绑定流程随之作废,不再触发外部换码。
|
||||
func TestOAuthBindStaleToken(t *testing.T) {
|
||||
o, auth := newOAuthEnv(t)
|
||||
ctx := context.Background()
|
||||
o.settings.SetEnvPublicURL("https://demo.example.com")
|
||||
secret := "secret"
|
||||
if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{
|
||||
GithubClientID: strPtr("cid"), GithubClientSecret: &secret,
|
||||
}); err != nil {
|
||||
t.Fatalf("UpdateOAuth: %v", err)
|
||||
}
|
||||
token, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "10.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
if _, err := o.AuthorizeURL(ctx, "github", "bind", "admin", token); err != nil {
|
||||
t.Fatalf("AuthorizeURL(bind): %v", err)
|
||||
}
|
||||
var state string
|
||||
o.mu.Lock()
|
||||
for k := range o.pending {
|
||||
state = k
|
||||
}
|
||||
o.mu.Unlock()
|
||||
if err := auth.bumpTokenVersion(ctx, "admin"); err != nil {
|
||||
t.Fatalf("bump token version: %v", err)
|
||||
}
|
||||
if _, _, _, err := o.HandleCallback(ctx, "github", state, "code", SessionMeta{ClientIP: "10.0.0.1"}); !errors.Is(err, ErrOAuthState) {
|
||||
t.Fatalf("stale-token callback err = %v, want ErrOAuthState", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestUpdateOAuthKeepsLogin 验证 provider 配置防自锁:密码登录禁用期间,
|
||||
// 禁用或清空最后可用的登录方式被拒;有通行密钥兜底后放行。
|
||||
func TestUpdateOAuthKeepsLogin(t *testing.T) {
|
||||
o, auth := newOAuthEnv(t)
|
||||
ctx := context.Background()
|
||||
o.settings.SetEnvPublicURL("https://demo.example.com")
|
||||
secret := "gh-secret"
|
||||
if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubClientID: strPtr("cid"), GithubClientSecret: &secret}); err != nil {
|
||||
t.Fatalf("UpdateOAuth: %v", err)
|
||||
}
|
||||
fakeBindIdentity(t, o, "admin", "github", "77", "octo")
|
||||
if err := saveSettingTx(auth.db, settingSecPasswordLoginOff, "1"); err != nil {
|
||||
t.Fatalf("disable password login: %v", err)
|
||||
}
|
||||
off := true
|
||||
if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubDisabled: &off}); !errors.Is(err, ErrProviderLastLogin) {
|
||||
t.Fatalf("disable last provider err = %v, want ErrProviderLastLogin", err)
|
||||
}
|
||||
if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubClientID: strPtr("")}); !errors.Is(err, ErrProviderLastLogin) {
|
||||
t.Fatalf("clear last clientID err = %v, want ErrProviderLastLogin", err)
|
||||
}
|
||||
pk := model.UserPasskey{
|
||||
UserID: 1, Name: "k", CredentialID: "c", CredentialIDHash: "h",
|
||||
Credential: "{}", Origin: "https://demo.example.com",
|
||||
}
|
||||
if err := auth.db.Create(&pk).Error; err != nil {
|
||||
t.Fatalf("seed passkey: %v", err)
|
||||
}
|
||||
if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubDisabled: &off}); err != nil {
|
||||
t.Fatalf("disable provider with passkey fallback: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestUpdateSecurityAppURLKeepsLogin 验证面板地址防自锁:密码禁用期间
|
||||
// 清空地址一律拒绝;域名变更须有钱包身份兜底(通行密钥随 RP ID 失效)。
|
||||
func TestUpdateSecurityAppURLKeepsLogin(t *testing.T) {
|
||||
o, auth := newOAuthEnv(t)
|
||||
ctx := context.Background()
|
||||
if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://a.example.com")}); err != nil {
|
||||
t.Fatalf("seed app url: %v", err)
|
||||
}
|
||||
pk := model.UserPasskey{
|
||||
UserID: 1, Name: "k", CredentialID: "c", CredentialIDHash: "h",
|
||||
Credential: "{}", Origin: "https://a.example.com",
|
||||
}
|
||||
if err := auth.db.Create(&pk).Error; err != nil {
|
||||
t.Fatalf("seed passkey: %v", err)
|
||||
}
|
||||
if err := saveSettingTx(auth.db, settingSecPasswordLoginOff, "1"); err != nil {
|
||||
t.Fatalf("disable password login: %v", err)
|
||||
}
|
||||
if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("")}); !errors.Is(err, ErrProviderLastLogin) {
|
||||
t.Fatalf("clear app url err = %v, want ErrProviderLastLogin", err)
|
||||
}
|
||||
if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://b.example.com")}); !errors.Is(err, ErrProviderLastLogin) {
|
||||
t.Fatalf("change host err = %v, want ErrProviderLastLogin", err)
|
||||
}
|
||||
ident := model.UserIdentity{UserID: 1, Provider: "wallet", Subject: "0xaa", Display: "w"}
|
||||
if err := auth.db.Create(&ident).Error; err != nil {
|
||||
t.Fatalf("seed wallet identity: %v", err)
|
||||
}
|
||||
if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://b.example.com")}); err != nil {
|
||||
t.Fatalf("change host with wallet fallback: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPasswordDisableStalePasskeyOrigin 验证逆序自锁防线:地址 A 注册的
|
||||
// 通行密钥在改到地址 B 后不再计入「可用免密方式」,禁用密码被拒。
|
||||
func TestPasswordDisableStalePasskeyOrigin(t *testing.T) {
|
||||
o, auth := newOAuthEnv(t)
|
||||
ctx := context.Background()
|
||||
auth.SetNotifier(nil, o.settings)
|
||||
if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://a.example.com")}); err != nil {
|
||||
t.Fatalf("seed app url: %v", err)
|
||||
}
|
||||
pk := model.UserPasskey{UserID: 1, Name: "k", CredentialID: "c", CredentialIDHash: "h",
|
||||
Origin: "https://a.example.com", Credential: "{}"}
|
||||
if err := auth.db.Create(&pk).Error; err != nil {
|
||||
t.Fatalf("seed passkey: %v", err)
|
||||
}
|
||||
// 密码未禁用:改地址不受限
|
||||
if err := o.settings.UpdateSecurity(ctx, SecurityPatch{AppURL: strPtr("https://b.example.com")}); err != nil {
|
||||
t.Fatalf("change app url: %v", err)
|
||||
}
|
||||
// 旧地址的通行密钥不计入可用方式:禁用密码被拒,不再自锁
|
||||
if err := auth.SetPasswordLoginDisabled(ctx, "admin", true, proofOf(t, auth.db, "admin")); !errors.Is(err, ErrNeedIdentity) {
|
||||
t.Fatalf("disable with stale-origin passkey err = %v, want ErrNeedIdentity", err)
|
||||
}
|
||||
// 当前地址重新注册(origin 一致)后可禁用
|
||||
pk2 := model.UserPasskey{UserID: 1, Name: "k2", CredentialID: "c2", CredentialIDHash: "h2",
|
||||
Origin: "https://b.example.com", Credential: "{}"}
|
||||
if err := auth.db.Create(&pk2).Error; err != nil {
|
||||
t.Fatalf("seed passkey2: %v", err)
|
||||
}
|
||||
if err := auth.SetPasswordLoginDisabled(ctx, "admin", true, proofOf(t, auth.db, "admin")); err != nil {
|
||||
t.Fatalf("disable with current-origin passkey: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthenticatedSettingsRejectRevokedProof(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
apply func(*OAuthService, *AuthService, TokenProof) error
|
||||
}{
|
||||
{
|
||||
name: "security patch",
|
||||
apply: func(o *OAuthService, auth *AuthService, proof TokenProof) error {
|
||||
return o.settings.UpdateSecurityAuthenticated(
|
||||
context.Background(), SecurityPatch{LoginFailLimit: intPtr(9)}, auth, "admin", proof)
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "oauth patch",
|
||||
apply: func(o *OAuthService, auth *AuthService, proof TokenProof) error {
|
||||
return o.settings.UpdateOAuthAuthenticated(
|
||||
context.Background(), UpdateOAuthInput{GithubDisplayName: strPtr("late")}, auth, "admin", proof)
|
||||
},
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
o, auth := newOAuthEnv(t)
|
||||
proof := revokeSettingsToken(t, auth)
|
||||
if err := tt.apply(o, auth, proof); !errors.Is(err, ErrTokenStale) {
|
||||
t.Fatalf("patch err = %v, want ErrTokenStale", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func revokeSettingsToken(t *testing.T, auth *AuthService) TokenProof {
|
||||
t.Helper()
|
||||
stale := loginSession(t, auth, "10.20.0.1", "stale")
|
||||
current := loginSession(t, auth, "10.20.0.2", "current")
|
||||
proof := mustProof(t, auth, stale)
|
||||
for _, item := range mustSessions(t, auth, current) {
|
||||
if !item.Current {
|
||||
if err := auth.RevokeSession(context.Background(), "admin", current, item.ID); err != nil {
|
||||
t.Fatalf("RevokeSession: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
return proof
|
||||
}
|
||||
|
||||
func newOAuthEnv(t *testing.T) (*OAuthService, *AuthService) {
|
||||
t.Helper()
|
||||
auth, db := newTotpEnv(t)
|
||||
@@ -159,17 +383,86 @@ func newOAuthEnv(t *testing.T) (*OAuthService, *AuthService) {
|
||||
return NewOAuthService(db, settings, auth), auth
|
||||
}
|
||||
|
||||
func TestOAuthLoginRechecksIdentityAndProvider(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
mutate func(*testing.T, *OAuthService, *model.UserIdentity)
|
||||
want error
|
||||
}{
|
||||
{
|
||||
name: "identity unbound",
|
||||
mutate: func(t *testing.T, o *OAuthService, row *model.UserIdentity) {
|
||||
if err := o.db.Delete(row).Error; err != nil {
|
||||
t.Fatalf("delete identity: %v", err)
|
||||
}
|
||||
},
|
||||
want: ErrOAuthNotBound,
|
||||
},
|
||||
{
|
||||
name: "provider disabled",
|
||||
mutate: func(t *testing.T, o *OAuthService, _ *model.UserIdentity) {
|
||||
patch := UpdateOAuthInput{GithubDisabled: boolPtr(true)}
|
||||
if err := o.settings.UpdateOAuth(context.Background(), patch); err != nil {
|
||||
t.Fatalf("disable provider: %v", err)
|
||||
}
|
||||
},
|
||||
want: ErrOAuthDisabled,
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) { testOAuthLoginRecheck(t, tt.mutate, tt.want) })
|
||||
}
|
||||
}
|
||||
|
||||
func testOAuthLoginRecheck(
|
||||
t *testing.T, mutate func(*testing.T, *OAuthService, *model.UserIdentity), want error,
|
||||
) {
|
||||
t.Helper()
|
||||
o, _ := newOAuthEnv(t)
|
||||
o.settings.SetEnvPublicURL("https://demo.example.com")
|
||||
secret := "secret"
|
||||
if err := o.settings.UpdateOAuth(context.Background(), UpdateOAuthInput{
|
||||
GithubClientID: strPtr("cid"), GithubClientSecret: &secret,
|
||||
}); err != nil {
|
||||
t.Fatalf("UpdateOAuth: %v", err)
|
||||
}
|
||||
fakeBindIdentity(t, o, "admin", "github", "10086", "octocat")
|
||||
row, err := o.findIdentity(context.Background(), "github", "10086")
|
||||
if err != nil {
|
||||
t.Fatalf("findIdentity: %v", err)
|
||||
}
|
||||
mutate(t, o, row)
|
||||
_, _, err = o.loginIdentityRow(
|
||||
context.Background(), row, "github", "10086", SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if !errors.Is(err, want) {
|
||||
t.Fatalf("login err = %v, want %v", err, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthBindLoginUnbind(t *testing.T) {
|
||||
o, auth := newOAuthEnv(t)
|
||||
ctx := context.Background()
|
||||
o.settings.SetEnvPublicURL("https://demo.example.com")
|
||||
secret := "secret"
|
||||
if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{
|
||||
GithubClientID: strPtr("cid"), GithubClientSecret: &secret,
|
||||
}); err != nil {
|
||||
t.Fatalf("UpdateOAuth: %v", err)
|
||||
}
|
||||
|
||||
fakeBindIdentity(t, o, "admin", "github", "10086", "octocat")
|
||||
// 重复绑定同一身份拒绝
|
||||
if err := o.bind(ctx, "admin", "github", externalIdentity{Subject: "10086", Display: "octocat"}); !errors.Is(err, ErrOAuthBound) {
|
||||
if _, err := o.bind(ctx, fakePending(t, o, "admin", "github"), externalIdentity{Subject: "10086", Display: "octocat"}, SessionMeta{}); !errors.Is(err, ErrOAuthBound) {
|
||||
t.Errorf("重复绑定 err = %v, want ErrOAuthBound", err)
|
||||
}
|
||||
// 发起后版本已变(改密/撤销全部):事务内比对拒绝,身份不落库
|
||||
stale := fakePending(t, o, "admin", "github")
|
||||
stale.proof.Ver--
|
||||
if _, err := o.bind(ctx, stale, externalIdentity{Subject: "20250", Display: "x"}, SessionMeta{}); !errors.Is(err, ErrOAuthState) {
|
||||
t.Errorf("stale-version bind err = %v, want ErrOAuthState", err)
|
||||
}
|
||||
// 已绑定身份可登录并拿到有效 JWT
|
||||
token, loginUser, err := o.loginByIdentity(ctx, "github", externalIdentity{Subject: "10086", Display: "octocat"})
|
||||
token, loginUser, err := o.loginByIdentity(ctx, "github", externalIdentity{Subject: "10086", Display: "octocat"}, SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil || token == "" {
|
||||
t.Fatalf("loginByIdentity: %v", err)
|
||||
}
|
||||
@@ -180,7 +473,7 @@ func TestOAuthBindLoginUnbind(t *testing.T) {
|
||||
t.Errorf("token 应属 admin, got %q (%v)", username, err)
|
||||
}
|
||||
// 未绑定身份拒绝登录
|
||||
if _, _, err := o.loginByIdentity(ctx, "github", externalIdentity{Subject: "999"}); !errors.Is(err, ErrOAuthNotBound) {
|
||||
if _, _, err := o.loginByIdentity(ctx, "github", externalIdentity{Subject: "999"}, SessionMeta{ClientIP: "127.0.0.1"}); !errors.Is(err, ErrOAuthNotBound) {
|
||||
t.Errorf("未绑定登录 err = %v, want ErrOAuthNotBound", err)
|
||||
}
|
||||
// 列表与解绑
|
||||
@@ -188,7 +481,7 @@ func TestOAuthBindLoginUnbind(t *testing.T) {
|
||||
if err != nil || len(items) != 1 {
|
||||
t.Fatalf("identities = %d (%v), want 1", len(items), err)
|
||||
}
|
||||
if err := o.Unbind(ctx, "admin", items[0].ID); err != nil {
|
||||
if err := o.Unbind(ctx, "admin", items[0].ID, proofOf(t, o.db, "admin")); err != nil {
|
||||
t.Fatalf("Unbind: %v", err)
|
||||
}
|
||||
if items, _ = o.Identities(ctx, "admin"); len(items) != 0 {
|
||||
@@ -228,6 +521,11 @@ func TestOAuthProvidersListsConfigured(t *testing.T) {
|
||||
if err != nil {
|
||||
t.Fatalf("UpdateOAuth: %v", err)
|
||||
}
|
||||
// 面板地址缺失时回调无从拼接:半配置不暴露必败入口
|
||||
if got := o.Providers(ctx); len(got) != 0 {
|
||||
t.Fatalf("无面板地址 providers = %v, want empty", got)
|
||||
}
|
||||
o.settings.SetEnvPublicURL("https://demo.example.com")
|
||||
got := o.Providers(ctx)
|
||||
if len(got) != 1 || got[0].Provider != "github" {
|
||||
t.Errorf("providers = %v, want [github]", got)
|
||||
@@ -245,26 +543,34 @@ func TestOAuthAuthorizeConfigErrors(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// clientID 缺失
|
||||
if _, err := o.AuthorizeURL(ctx, "github", "login", ""); !errors.Is(err, ErrOAuthNotConfigured) {
|
||||
if _, err := o.AuthorizeURL(ctx, "github", "login", "", ""); !errors.Is(err, ErrOAuthNotConfigured) {
|
||||
t.Errorf("无 clientID err = %v, want ErrOAuthNotConfigured", err)
|
||||
}
|
||||
// clientID 已配但面板地址未设置
|
||||
// 仅 clientID 的半配置仍按未配置拒绝(不能生成必败授权 URL)
|
||||
if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubClientID: strPtr("Iv1.test")}); err != nil {
|
||||
t.Fatalf("UpdateOAuth: %v", err)
|
||||
}
|
||||
if _, err := o.AuthorizeURL(ctx, "github", "login", ""); !errors.Is(err, ErrOAuthNoAppURL) {
|
||||
if _, err := o.AuthorizeURL(ctx, "github", "login", "", ""); !errors.Is(err, ErrOAuthNotConfigured) {
|
||||
t.Errorf("无 secret err = %v, want ErrOAuthNotConfigured", err)
|
||||
}
|
||||
// clientID + secret 已配但面板地址未设置
|
||||
secret := "secret"
|
||||
if err := o.settings.UpdateOAuth(ctx, UpdateOAuthInput{GithubClientSecret: &secret}); err != nil {
|
||||
t.Fatalf("UpdateOAuth secret: %v", err)
|
||||
}
|
||||
if _, err := o.AuthorizeURL(ctx, "github", "login", "", ""); !errors.Is(err, ErrOAuthNoAppURL) {
|
||||
t.Errorf("无面板地址 err = %v, want ErrOAuthNoAppURL", err)
|
||||
}
|
||||
// 面板地址就绪后正常返回授权 URL
|
||||
o.settings.SetEnvPublicURL("https://demo.example.com")
|
||||
url, err := o.AuthorizeURL(ctx, "github", "login", "")
|
||||
url, err := o.AuthorizeURL(ctx, "github", "login", "", "")
|
||||
if err != nil || url == "" {
|
||||
t.Fatalf("AuthorizeURL: %q, %v", url, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOAuthProvidersAndDisabled(t *testing.T) {
|
||||
o, _ := newOAuthEnv(t)
|
||||
o, auth := newOAuthEnv(t)
|
||||
ctx := context.Background()
|
||||
o.settings.SetEnvPublicURL("https://demo.example.com")
|
||||
|
||||
@@ -272,11 +578,19 @@ func TestOAuthProvidersAndDisabled(t *testing.T) {
|
||||
if got := o.Providers(ctx); len(got) != 0 {
|
||||
t.Fatalf("Providers = %v, want empty", got)
|
||||
}
|
||||
// 配置 github(无显示名称)→ 默认名 GitHub
|
||||
// 配置 github(无显示名称)→ 默认名 GitHub;仅 clientID 的半配置不暴露
|
||||
in := UpdateOAuthInput{GithubClientID: strPtr("Iv1.test")}
|
||||
if err := o.settings.UpdateOAuth(ctx, in); err != nil {
|
||||
t.Fatalf("UpdateOAuth: %v", err)
|
||||
}
|
||||
if got := o.Providers(ctx); len(got) != 0 {
|
||||
t.Fatalf("仅 clientID 半配置 Providers = %v, want empty", got)
|
||||
}
|
||||
ghSecret := "gh-secret"
|
||||
in.GithubClientSecret = &ghSecret
|
||||
if err := o.settings.UpdateOAuth(ctx, in); err != nil {
|
||||
t.Fatalf("UpdateOAuth: %v", err)
|
||||
}
|
||||
got := o.Providers(ctx)
|
||||
if len(got) != 1 || got[0].Provider != "github" || got[0].DisplayName != "GitHub" {
|
||||
t.Fatalf("Providers = %+v, want [github/GitHub]", got)
|
||||
@@ -297,12 +611,20 @@ func TestOAuthProvidersAndDisabled(t *testing.T) {
|
||||
if got = o.Providers(ctx); len(got) != 0 {
|
||||
t.Fatalf("禁用后 Providers = %v, want empty", got)
|
||||
}
|
||||
if _, err := o.AuthorizeURL(ctx, "github", "login", ""); !errors.Is(err, ErrOAuthDisabled) {
|
||||
if _, err := o.AuthorizeURL(ctx, "github", "login", "", ""); !errors.Is(err, ErrOAuthDisabled) {
|
||||
t.Errorf("禁用 login err = %v, want ErrOAuthDisabled", err)
|
||||
}
|
||||
if url, err := o.AuthorizeURL(ctx, "github", "bind", "admin"); err != nil || url == "" {
|
||||
// bind 模式须携带有效令牌(发起即验),但不受 provider 禁用影响
|
||||
token, _, err := auth.Login(ctx, "admin", "pass123", "", SessionMeta{ClientIP: "10.0.0.9"})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
if url, err := o.AuthorizeURL(ctx, "github", "bind", "admin", token); err != nil || url == "" {
|
||||
t.Errorf("禁用 bind = %q, %v, want 正常返回", url, err)
|
||||
}
|
||||
if _, err := o.AuthorizeURL(ctx, "github", "bind", "admin", ""); !errors.Is(err, ErrOAuthState) {
|
||||
t.Errorf("bind 无令牌 err = %v, want ErrOAuthState", err)
|
||||
}
|
||||
// view 回读禁用态与显示名称
|
||||
view, err := o.settings.OAuthView(ctx)
|
||||
if err != nil || !view.GithubDisabled || view.GithubDisplayName != "公司账号" {
|
||||
|
||||
@@ -84,8 +84,9 @@ func (s *OciConfigService) DeleteTenantUserApiKey(ctx context.Context, id uint,
|
||||
}
|
||||
|
||||
// ActivateApiKey 把刚创建的 key 设为本配置的签名凭据:验证可用后落库,不删除旧 key。
|
||||
// 私钥由前端回传(仅创建时下发过);归属无需显式校验——非本签名用户的 key 验证必失败。
|
||||
func (s *OciConfigService) ActivateApiKey(ctx context.Context, id uint, fingerprint, privateKey string) error {
|
||||
// userID 非空时以该用户身份验证并一并切换签名用户(空串沿用当前用户);
|
||||
// 私钥由前端回传(仅创建时下发过);归属无需显式校验——非归属用户的 key 验证必失败。
|
||||
func (s *OciConfigService) ActivateApiKey(ctx context.Context, id uint, userID, fingerprint, privateKey string) error {
|
||||
cfg, err := s.Get(ctx, id)
|
||||
if err != nil {
|
||||
return err
|
||||
@@ -96,6 +97,9 @@ func (s *OciConfigService) ActivateApiKey(ctx context.Context, id uint, fingerpr
|
||||
}
|
||||
newCred := cred
|
||||
newCred.Fingerprint, newCred.PrivateKey, newCred.Passphrase = fingerprint, privateKey, ""
|
||||
if userID != "" {
|
||||
newCred.UserOCID = userID
|
||||
}
|
||||
if err := newCred.Validate(); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -121,13 +125,16 @@ func (s *OciConfigService) waitApiKeyUsable(ctx context.Context, cred oci.Creden
|
||||
return fmt.Errorf("new api key not usable: %w", err)
|
||||
}
|
||||
|
||||
// persistSigningKey 加密新私钥,更新配置指纹并清空口令密文(面板生成的 key 无口令)。
|
||||
// persistSigningKey 加密新私钥,更新配置签名用户与指纹并清空口令密文(面板生成的 key 无口令)。
|
||||
func (s *OciConfigService) persistSigningKey(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)
|
||||
}
|
||||
updates := map[string]any{"fingerprint": newCred.Fingerprint, "private_key_enc": enc, "passphrase_enc": ""}
|
||||
updates := map[string]any{
|
||||
"user_oc_id": newCred.UserOCID, "fingerprint": newCred.Fingerprint,
|
||||
"private_key_enc": enc, "passphrase_enc": "",
|
||||
}
|
||||
if err := s.db.Model(cfg).Updates(updates).Error; err != nil {
|
||||
return fmt.Errorf("persist rotated key: %w", err)
|
||||
}
|
||||
|
||||
@@ -17,13 +17,14 @@ import (
|
||||
type apiKeyClient struct {
|
||||
oci.Client
|
||||
|
||||
keys []oci.TenantUserApiKey
|
||||
uploadFp string
|
||||
uploadErr error
|
||||
uploadedPub string
|
||||
deleted []string
|
||||
validateErr error
|
||||
validated []string // ValidateKey 收到的指纹序列
|
||||
keys []oci.TenantUserApiKey
|
||||
uploadFp string
|
||||
uploadErr error
|
||||
uploadedPub string
|
||||
deleted []string
|
||||
validateErr error
|
||||
validated []string // ValidateKey 收到的指纹序列
|
||||
validatedUser string // ValidateKey 最近一次使用的签名用户
|
||||
}
|
||||
|
||||
func (f *apiKeyClient) ListTenantUserApiKeys(ctx context.Context, cred oci.Credentials, homeRegion, userID string) ([]oci.TenantUserApiKey, error) {
|
||||
@@ -45,6 +46,7 @@ func (f *apiKeyClient) DeleteTenantUserApiKey(ctx context.Context, cred oci.Cred
|
||||
|
||||
func (f *apiKeyClient) ValidateKey(ctx context.Context, cred oci.Credentials) (oci.TenancyInfo, error) {
|
||||
f.validated = append(f.validated, cred.Fingerprint)
|
||||
f.validatedUser = cred.UserOCID
|
||||
if f.validateErr != nil {
|
||||
return oci.TenancyInfo{}, f.validateErr
|
||||
}
|
||||
@@ -149,14 +151,17 @@ func TestActivateApiKey(t *testing.T) {
|
||||
const newKey = "-----BEGIN RSA PRIVATE KEY-----\nnew\n-----END RSA PRIVATE KEY-----"
|
||||
tests := []struct {
|
||||
name string
|
||||
userID string
|
||||
privateKey string
|
||||
validateErr error
|
||||
wantErr bool
|
||||
wantFp string // 期望落库指纹
|
||||
wantUser string // 期望落库签名用户
|
||||
}{
|
||||
{"成功:验证通过后落库,不删任何 key", newKey, nil, false, "11:22"},
|
||||
{"验证失败:配置不动", newKey, errors.New("401"), true, "aa:bb"},
|
||||
{"私钥非 PEM:直接拒绝", "not-a-pem", nil, true, "aa:bb"},
|
||||
{"成功:验证通过后落库,不删任何 key", "", newKey, nil, false, "11:22", "ocid1.user.oc1..me"},
|
||||
{"切换用户:以新用户验证并一并落库", "ocid1.user.oc1..other", newKey, nil, false, "11:22", "ocid1.user.oc1..other"},
|
||||
{"验证失败:配置不动", "", newKey, errors.New("401"), true, "aa:bb", "ocid1.user.oc1..me"},
|
||||
{"私钥非 PEM:直接拒绝", "", "not-a-pem", nil, true, "aa:bb", "ocid1.user.oc1..me"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
@@ -164,7 +169,7 @@ func TestActivateApiKey(t *testing.T) {
|
||||
s := newTestService(t, fc)
|
||||
cfg := seedApiKeyConfig(t, s)
|
||||
|
||||
err := s.ActivateApiKey(context.Background(), cfg.ID, "11:22", tt.privateKey)
|
||||
err := s.ActivateApiKey(context.Background(), cfg.ID, tt.userID, "11:22", tt.privateKey)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Fatalf("err = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
@@ -175,13 +180,16 @@ func TestActivateApiKey(t *testing.T) {
|
||||
if got.Fingerprint != tt.wantFp {
|
||||
t.Fatalf("fingerprint = %q, want %q", got.Fingerprint, tt.wantFp)
|
||||
}
|
||||
if got.UserOCID != tt.wantUser {
|
||||
t.Fatalf("userOCID = %q, want %q", got.UserOCID, tt.wantUser)
|
||||
}
|
||||
if len(fc.deleted) != 0 {
|
||||
t.Fatalf("deleted = %v, want none", fc.deleted)
|
||||
}
|
||||
if tt.wantErr {
|
||||
return
|
||||
}
|
||||
// 成功路径:落库私钥可解密且与回传一致,验证调用用的是新指纹
|
||||
// 成功路径:落库私钥可解密且与回传一致,验证调用用的是新指纹与目标用户
|
||||
plain, err := s.cipher.DecryptString(got.PrivateKeyEnc)
|
||||
if err != nil || plain != newKey {
|
||||
t.Fatalf("persisted key mismatch (err=%v)", err)
|
||||
@@ -189,6 +197,9 @@ func TestActivateApiKey(t *testing.T) {
|
||||
if len(fc.validated) == 0 || fc.validated[0] != "11:22" {
|
||||
t.Fatalf("validated = %v", fc.validated)
|
||||
}
|
||||
if fc.validatedUser != tt.wantUser {
|
||||
t.Fatalf("validatedUser = %q, want %q", fc.validatedUser, tt.wantUser)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,308 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
|
||||
// 钱包(SIWE)流程错误;api 层映射为用户可读提示。
|
||||
var (
|
||||
// ErrWalletNoAppURL 表示面板地址缺失,EIP-4361 的 domain/URI 无从派生。
|
||||
ErrWalletNoAppURL = errors.New("面板地址未设置,请先在「设置 → 安全 → 网络与地址」保存面板地址")
|
||||
// ErrWalletAddress 表示地址格式非法(非 0x+40 hex)。
|
||||
ErrWalletAddress = errors.New("钱包地址格式不正确")
|
||||
// ErrWalletChallenge 表示挑战无效或已过期(一次性消费)。
|
||||
ErrWalletChallenge = errors.New("签名挑战无效或已过期,请重新发起")
|
||||
// ErrWalletSig 表示签名校验失败或恢复地址不匹配;统一文案防探测。
|
||||
ErrWalletSig = errors.New("钱包签名校验失败")
|
||||
// ErrWalletNotBound 表示地址未绑定任何账号,拒绝登录(不开放注册)。
|
||||
ErrWalletNotBound = errors.New("该钱包地址未绑定面板账号,请先登录后在设置中绑定")
|
||||
// ErrWalletBound 表示地址已被绑定(重复绑定)。
|
||||
ErrWalletBound = errors.New("该钱包地址已绑定过")
|
||||
)
|
||||
|
||||
// walletProvider 是钱包身份在 UserIdentity 表中的 provider 取值。
|
||||
const walletProvider = "wallet"
|
||||
|
||||
// walletPendingTTL 是签名挑战有效期,与消息中的 Expiration Time 一致。
|
||||
const walletPendingTTL = 5 * time.Minute
|
||||
|
||||
// walletPending 是一次进行中的签名挑战;nonce 一次性使用。
|
||||
type walletPending struct {
|
||||
message string // EIP-4361 消息全文,验签对象
|
||||
address string // EIP-55 规范地址
|
||||
mode string // "login" / "bind"
|
||||
username string // bind 模式的绑定目标账号
|
||||
token string // bind 模式发起时的 Bearer;verify 复验,防被盗令牌撤销后仍完成绑定
|
||||
proof TokenProof // 发起时的版本/jti 快照;绑定事务行锁下复核,覆盖撤销全部/注销/定点撤销
|
||||
expires time.Time
|
||||
}
|
||||
|
||||
// WalletService 承接以太坊钱包(EIP-4361)绑定与登录;身份复用 UserIdentity 通道。
|
||||
type WalletService struct {
|
||||
db *gorm.DB
|
||||
settings *SettingService
|
||||
auth *AuthService
|
||||
|
||||
mu sync.Mutex
|
||||
pending map[string]walletPending
|
||||
}
|
||||
|
||||
// NewWalletService 组装依赖。
|
||||
func NewWalletService(db *gorm.DB, settings *SettingService, auth *AuthService) *WalletService {
|
||||
return &WalletService{db: db, settings: settings, auth: auth, pending: map[string]walletPending{}}
|
||||
}
|
||||
|
||||
// Challenge 生成 EIP-4361 消息全文并登记一次性 nonce;
|
||||
// 消息由后端持有并作为最终验签对象,前端仅负责原样签名。
|
||||
func (w *WalletService) Challenge(ctx context.Context, address, mode, username, bindToken string) (nonce, message string, err error) {
|
||||
app := w.settings.EffectiveAppURL()
|
||||
if app == "" {
|
||||
return "", "", ErrWalletNoAppURL
|
||||
}
|
||||
addr, err := normalizeEthAddress(address)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
var proof TokenProof
|
||||
if mode == "bind" {
|
||||
if proof, err = w.bindTokenProof(ctx, username, bindToken); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
}
|
||||
nonce, err = randHex(16)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
message = buildSiweMessage(app, addr, mode, nonce, time.Now())
|
||||
w.mu.Lock()
|
||||
w.gcWalletLocked()
|
||||
w.pending[nonce] = walletPending{message: message, address: addr, mode: mode, username: username, token: bindToken, proof: proof, expires: time.Now().Add(walletPendingTTL)}
|
||||
w.mu.Unlock()
|
||||
return nonce, message, nil
|
||||
}
|
||||
|
||||
// bindTokenProof 验证 bind 发起令牌的有效性与归属,返回其版本 / jti 快照;
|
||||
// 绑定事务行锁下复核该快照,发起后改密、撤销全部、注销或定点撤销均令绑定作废。
|
||||
func (w *WalletService) bindTokenProof(ctx context.Context, username, token string) (TokenProof, error) {
|
||||
name, proof, err := w.auth.ParseTokenProof(ctx, token)
|
||||
if err != nil || name != username {
|
||||
return TokenProof{}, ErrWalletChallenge
|
||||
}
|
||||
return proof, nil
|
||||
}
|
||||
|
||||
// buildSiweMessage 按 EIP-4361 模板拼装消息;Chain ID 固定 1
|
||||
// (personal_sign 与链无关,字段仅为满足标准的人类可读格式)。
|
||||
func buildSiweMessage(appURL, address, mode, nonce string, now time.Time) string {
|
||||
domain := appURL
|
||||
if u, err := url.Parse(appURL); err == nil && u.Host != "" {
|
||||
domain = u.Host
|
||||
}
|
||||
statement := "登录 OCI Portal 面板"
|
||||
if mode == "bind" {
|
||||
statement = "将此钱包绑定到 OCI Portal 账号"
|
||||
}
|
||||
return domain + " wants you to sign in with your Ethereum account:\n" +
|
||||
address + "\n\n" +
|
||||
statement + "\n\n" +
|
||||
"URI: " + appURL + "\n" +
|
||||
"Version: 1\n" +
|
||||
"Chain ID: 1\n" +
|
||||
"Nonce: " + nonce + "\n" +
|
||||
"Issued At: " + now.UTC().Format(time.RFC3339) + "\n" +
|
||||
"Expiration Time: " + now.Add(walletPendingTTL).UTC().Format(time.RFC3339)
|
||||
}
|
||||
|
||||
// gcWalletLocked 清理过期挑战;调用方须持锁。
|
||||
func (w *WalletService) gcWalletLocked() {
|
||||
now := time.Now()
|
||||
for k, p := range w.pending {
|
||||
if now.After(p.expires) {
|
||||
delete(w.pending, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// takeChallenge 取出并消费挑战(一次性);不存在或过期视为无效。
|
||||
func (w *WalletService) takeChallenge(nonce string) (walletPending, error) {
|
||||
w.mu.Lock()
|
||||
defer w.mu.Unlock()
|
||||
p, ok := w.pending[nonce]
|
||||
delete(w.pending, nonce)
|
||||
if !ok || time.Now().After(p.expires) {
|
||||
return walletPending{}, ErrWalletChallenge
|
||||
}
|
||||
return p, nil
|
||||
}
|
||||
|
||||
// WalletVerifyResult 是校验成功后的会话信息;Mode 在失败时也尽力携带,
|
||||
// 供 api 层留痕与决定响应形态。
|
||||
type WalletVerifyResult struct {
|
||||
Token string
|
||||
ExpiresAt time.Time
|
||||
Username string
|
||||
Mode string
|
||||
}
|
||||
|
||||
// Verify 校验签名并完成绑定或登录(落地会话);登录失败按「IP+地址」计入登录守卫,
|
||||
// 锁定期内一律 ErrLoginLocked。
|
||||
func (w *WalletService) Verify(ctx context.Context, nonce, signature string, meta SessionMeta) (WalletVerifyResult, error) {
|
||||
p, err := w.takeChallenge(nonce)
|
||||
if err != nil {
|
||||
return WalletVerifyResult{}, err
|
||||
}
|
||||
res := WalletVerifyResult{Mode: p.mode}
|
||||
if p.mode == "login" {
|
||||
meta.Method = walletProvider // 绑定不改登录方式,见 OAuth 侧同款注释
|
||||
}
|
||||
key := guardKey(meta.ClientIP, strings.ToLower(p.address))
|
||||
now := time.Now()
|
||||
sec := securityOf(w.auth.settings)
|
||||
if p.mode == "login" && w.auth.guard.locked(key, now, time.Duration(sec.LoginLockMinutes)*time.Minute) {
|
||||
return res, ErrLoginLocked
|
||||
}
|
||||
if err := verifyWalletSig(p, signature); err != nil {
|
||||
if p.mode == "login" {
|
||||
if lockErr := w.auth.failLogin(key, now, p.address, meta.ClientIP, sec); errors.Is(lockErr, ErrLoginLocked) {
|
||||
return res, ErrLoginLocked
|
||||
}
|
||||
}
|
||||
return res, err
|
||||
}
|
||||
if p.mode == "bind" {
|
||||
// 复验发起挑战时的令牌仍有效且归属一致:改密/撤销全部后挑战随之作废
|
||||
if name, tokenErr := w.auth.ParseToken(ctx, p.token); tokenErr != nil || name != p.username {
|
||||
return res, ErrWalletChallenge
|
||||
}
|
||||
res.Username = p.username
|
||||
res.Token, res.ExpiresAt, err = w.bind(ctx, p, meta)
|
||||
return res, err
|
||||
}
|
||||
res.Token, res.ExpiresAt, res.Username, err = w.login(ctx, p, key, meta)
|
||||
return res, err
|
||||
}
|
||||
|
||||
// verifyWalletSig 验签并比对恢复地址(大小写不敏感);失败统一 ErrWalletSig。
|
||||
func verifyWalletSig(p walletPending, signature string) error {
|
||||
raw, err := hex.DecodeString(strings.TrimPrefix(signature, "0x"))
|
||||
if err != nil {
|
||||
return ErrWalletSig
|
||||
}
|
||||
recovered, err := recoverEthAddress(personalSignDigest([]byte(p.message)), raw)
|
||||
if err != nil {
|
||||
return ErrWalletSig
|
||||
}
|
||||
if !strings.EqualFold(recovered, p.address) {
|
||||
return ErrWalletSig
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// bind 在单事务内完成绑定全程:行锁下比对发起时令牌版本(发起后被撤销即作废)、
|
||||
// 身份写入、版本递增、原会话行接续换发;任一失败整体回滚,不留半程状态。
|
||||
// 新令牌接续 p.token 的会话行(保留登录方式与创建时间),旧令牌无行时按 meta 新建。
|
||||
func (w *WalletService) bind(ctx context.Context, p walletPending, meta SessionMeta) (string, time.Time, error) {
|
||||
var token string
|
||||
var expires time.Time
|
||||
err := w.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserForAuthChange(tx, p.username)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := w.auth.ensureTokenCurrentTx(tx, user, p.proof); err != nil {
|
||||
return ErrWalletChallenge
|
||||
}
|
||||
if err := createIdentityTx(tx, user.ID, walletProvider, p.address, shortEthAddress(p.address), ErrWalletBound); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := bumpTokenVersionTx(tx, p.username); err != nil {
|
||||
return err
|
||||
}
|
||||
user.TokenVersion++
|
||||
token, expires, err = w.auth.renewSessionTx(tx, user, p.token, meta)
|
||||
return err
|
||||
})
|
||||
return token, expires, err
|
||||
}
|
||||
|
||||
// login 查绑定关系并签发 JWT(落地会话);未绑定一律拒绝(不开放注册)。
|
||||
func (w *WalletService) login(ctx context.Context, p walletPending, key string, meta SessionMeta) (string, time.Time, string, error) {
|
||||
row, err := w.findWalletIdentity(ctx, p.address)
|
||||
if err != nil {
|
||||
return "", time.Time{}, "", err
|
||||
}
|
||||
token, expires, username, err := w.loginWalletIdentity(ctx, row, p.address, meta)
|
||||
if err == nil {
|
||||
w.auth.guard.success(key)
|
||||
}
|
||||
return token, expires, username, err
|
||||
}
|
||||
|
||||
func (w *WalletService) loginWalletIdentity(
|
||||
ctx context.Context, row *model.UserIdentity, address string, meta SessionMeta,
|
||||
) (string, time.Time, string, error) {
|
||||
var token, username string
|
||||
var expires time.Time
|
||||
err := w.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error {
|
||||
user, err := lockUserByIDForAuthChange(tx, row.UserID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err := identityStillBoundTx(tx, row.ID, user.ID, walletProvider, address); err != nil {
|
||||
if errors.Is(err, ErrOAuthNotBound) {
|
||||
return ErrWalletNotBound
|
||||
}
|
||||
return err
|
||||
}
|
||||
token, expires, err = w.auth.signSessionTokenTx(tx, user, meta)
|
||||
username = user.Username
|
||||
return err
|
||||
})
|
||||
return token, expires, username, err
|
||||
}
|
||||
|
||||
func (w *WalletService) findWalletIdentity(ctx context.Context, address string) (*model.UserIdentity, error) {
|
||||
var row model.UserIdentity
|
||||
err := w.db.WithContext(ctx).
|
||||
Where("provider = ? AND subject = ?", walletProvider, address).First(&row).Error
|
||||
if errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
return nil, ErrWalletNotBound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("find wallet identity: %w", err)
|
||||
}
|
||||
return &row, nil
|
||||
}
|
||||
|
||||
// HasAny 报告是否存在可登录的钱包身份;SIWE 依赖面板地址,缺失时不暴露入口。
|
||||
func (w *WalletService) HasAny(ctx context.Context) bool {
|
||||
if w.settings == nil || w.settings.EffectiveAppURL() == "" {
|
||||
return false
|
||||
}
|
||||
var count int64
|
||||
err := w.db.WithContext(ctx).Model(&model.UserIdentity{}).
|
||||
Where("provider = ?", walletProvider).Count(&count).Error
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
return count > 0
|
||||
}
|
||||
|
||||
// shortEthAddress 生成列表展示用的缩写地址(0x1234…abcd)。
|
||||
func shortEthAddress(addr string) string {
|
||||
if len(addr) < 12 {
|
||||
return addr
|
||||
}
|
||||
return addr[:6] + "…" + addr[len(addr)-4:]
|
||||
}
|
||||
@@ -0,0 +1,381 @@
|
||||
package service
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/glebarez/sqlite"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
|
||||
"oci-portal/internal/crypto"
|
||||
"oci-portal/internal/model"
|
||||
)
|
||||
|
||||
// newTestWallet 组装内存库上的 WalletService(admin 已建,面板地址经环境回退注入)。
|
||||
func newTestWallet(t *testing.T, appURL string) *WalletService {
|
||||
t.Helper()
|
||||
db, err := gorm.Open(sqlite.Open(":memory:"), &gorm.Config{
|
||||
Logger: logger.Default.LogMode(logger.Silent),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("open in-memory sqlite: %v", err)
|
||||
}
|
||||
if err := db.AutoMigrate(&model.User{}, &model.UserIdentity{}, &model.UserSession{}, &model.Setting{}); err != nil {
|
||||
t.Fatalf("auto migrate: %v", err)
|
||||
}
|
||||
auth := NewAuthService(db, "test-jwt-secret")
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("ensure admin: %v", err)
|
||||
}
|
||||
cipher, err := crypto.NewCipher("test-key")
|
||||
if err != nil {
|
||||
t.Fatalf("new cipher: %v", err)
|
||||
}
|
||||
settings := NewSettingService(db, cipher)
|
||||
settings.SetEnvPublicURL(appURL)
|
||||
return NewWalletService(db, settings, auth)
|
||||
}
|
||||
|
||||
func TestWalletChallenge(t *testing.T) {
|
||||
priv := testPrivKey(9)
|
||||
addr := ethAddrOf(priv)
|
||||
tests := []struct {
|
||||
name string
|
||||
appURL string
|
||||
address string
|
||||
wantErr error
|
||||
}{
|
||||
{name: "正常下发", appURL: "https://demo.example.com", address: strings.ToLower(addr)},
|
||||
{name: "无面板地址", appURL: "", address: addr, wantErr: ErrWalletNoAppURL},
|
||||
{name: "地址非法", appURL: "https://demo.example.com", address: "0x12", wantErr: ErrWalletAddress},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
w := newTestWallet(t, tt.appURL)
|
||||
nonce, message, err := w.Challenge(context.Background(), tt.address, "login", "", "")
|
||||
if tt.wantErr != nil {
|
||||
if !errors.Is(err, tt.wantErr) {
|
||||
t.Fatalf("Challenge err = %v, want %v", err, tt.wantErr)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge: %v", err)
|
||||
}
|
||||
// 消息含 EIP-55 规范地址、nonce 与 domain;地址大小写已被规范化
|
||||
for _, part := range []string{addr, "Nonce: " + nonce, "demo.example.com wants you"} {
|
||||
if !strings.Contains(message, part) {
|
||||
t.Errorf("message missing %q:\n%s", part, message)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// adminToken 以口令登录换取 admin 的有效会话令牌(bind 挑战的复验凭据)。
|
||||
func adminToken(t *testing.T, w *WalletService) string {
|
||||
t.Helper()
|
||||
token, _, err := w.auth.Login(context.Background(), "admin", "pass123", "", SessionMeta{ClientIP: "10.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("admin login: %v", err)
|
||||
}
|
||||
return token
|
||||
}
|
||||
|
||||
// bindWallet 走完整 bind 流程(挑战 → 签名 → 校验),返回校验结果。
|
||||
func bindWallet(t *testing.T, w *WalletService, seed byte) WalletVerifyResult {
|
||||
t.Helper()
|
||||
priv := testPrivKey(seed)
|
||||
nonce, message, err := w.Challenge(context.Background(), ethAddrOf(priv), "bind", "admin", adminToken(t, w))
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge(bind): %v", err)
|
||||
}
|
||||
res, err := w.Verify(context.Background(), nonce, signPersonal(priv, message, true), SessionMeta{ClientIP: "10.1.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("Verify(bind): %v", err)
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func TestWalletBindAndLogin(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
ctx := context.Background()
|
||||
priv := testPrivKey(9)
|
||||
addr := ethAddrOf(priv)
|
||||
|
||||
res := bindWallet(t, w, 9)
|
||||
if res.Mode != "bind" || res.Token == "" || res.Username != "admin" {
|
||||
t.Fatalf("bind result = %+v, want bind/admin with token", res)
|
||||
}
|
||||
var row model.UserIdentity
|
||||
if err := w.db.Where("provider = ? AND subject = ?", walletProvider, addr).First(&row).Error; err != nil {
|
||||
t.Fatalf("identity row: %v", err)
|
||||
}
|
||||
if row.Display != shortEthAddress(addr) {
|
||||
t.Errorf("display = %q, want %q", row.Display, shortEthAddress(addr))
|
||||
}
|
||||
// 钱包身份计入外部身份数(禁用密码登录的门槛)
|
||||
if n, err := identityCountTx(w.db, row.UserID); err != nil || n != 1 {
|
||||
t.Errorf("identityCountTx = (%d, %v), want 1", n, err)
|
||||
}
|
||||
if !w.HasAny(ctx) {
|
||||
t.Error("HasAny = false, want true after bind")
|
||||
}
|
||||
|
||||
// 重复绑定同一地址被拒
|
||||
nonce, message, err := w.Challenge(ctx, addr, "bind", "admin", adminToken(t, w))
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge(rebind): %v", err)
|
||||
}
|
||||
if _, err := w.Verify(ctx, nonce, signPersonal(priv, message, true), SessionMeta{ClientIP: "10.1.0.1"}); !errors.Is(err, ErrWalletBound) {
|
||||
t.Fatalf("rebind err = %v, want ErrWalletBound", err)
|
||||
}
|
||||
|
||||
// 登录:全大写地址输入被规范化,v=0/1 形态签名同样有效
|
||||
nonce, message, err = w.Challenge(ctx, "0x"+strings.ToUpper(addr[2:]), "login", "", "")
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge(login): %v", err)
|
||||
}
|
||||
got, err := w.Verify(ctx, nonce, signPersonal(priv, message, false), SessionMeta{ClientIP: "10.1.0.2"})
|
||||
if err != nil {
|
||||
t.Fatalf("Verify(login): %v", err)
|
||||
}
|
||||
if got.Mode != "login" || got.Username != "admin" || got.Token == "" || got.ExpiresAt.IsZero() {
|
||||
t.Fatalf("login result = %+v, want login/admin with token+expiry", got)
|
||||
}
|
||||
if _, err := w.auth.ParseToken(ctx, got.Token); err != nil {
|
||||
t.Errorf("issued token invalid: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWalletVerifyRejections(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
ctx := context.Background()
|
||||
bound := testPrivKey(9)
|
||||
bindWallet(t, w, 9)
|
||||
|
||||
t.Run("nonce 一次性", func(t *testing.T) {
|
||||
nonce, message, _ := w.Challenge(ctx, ethAddrOf(bound), "login", "", "")
|
||||
sig := signPersonal(bound, message, true)
|
||||
if _, err := w.Verify(ctx, nonce, sig, SessionMeta{ClientIP: "10.2.0.1"}); err != nil {
|
||||
t.Fatalf("first verify: %v", err)
|
||||
}
|
||||
if _, err := w.Verify(ctx, nonce, sig, SessionMeta{ClientIP: "10.2.0.1"}); !errors.Is(err, ErrWalletChallenge) {
|
||||
t.Errorf("replayed nonce err = %v, want ErrWalletChallenge", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("过期挑战", func(t *testing.T) {
|
||||
nonce, message, _ := w.Challenge(ctx, ethAddrOf(bound), "login", "", "")
|
||||
w.mu.Lock()
|
||||
p := w.pending[nonce]
|
||||
p.expires = time.Now().Add(-time.Second)
|
||||
w.pending[nonce] = p
|
||||
w.mu.Unlock()
|
||||
if _, err := w.Verify(ctx, nonce, signPersonal(bound, message, true), SessionMeta{ClientIP: "10.2.0.2"}); !errors.Is(err, ErrWalletChallenge) {
|
||||
t.Errorf("expired nonce err = %v, want ErrWalletChallenge", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("他人签名(地址不匹配)", func(t *testing.T) {
|
||||
nonce, message, _ := w.Challenge(ctx, ethAddrOf(bound), "login", "", "")
|
||||
other := testPrivKey(13)
|
||||
if _, err := w.Verify(ctx, nonce, signPersonal(other, message, true), SessionMeta{ClientIP: "10.2.0.3"}); !errors.Is(err, ErrWalletSig) {
|
||||
t.Errorf("foreign signature err = %v, want ErrWalletSig", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("未绑定地址拒登", func(t *testing.T) {
|
||||
stranger := testPrivKey(21)
|
||||
nonce, message, _ := w.Challenge(ctx, ethAddrOf(stranger), "login", "", "")
|
||||
if _, err := w.Verify(ctx, nonce, signPersonal(stranger, message, true), SessionMeta{ClientIP: "10.2.0.4"}); !errors.Is(err, ErrWalletNotBound) {
|
||||
t.Errorf("unbound address err = %v, want ErrWalletNotBound", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
func TestWalletBindStaleToken(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
ctx := context.Background()
|
||||
priv := testPrivKey(9)
|
||||
|
||||
// 挑战发起后全量撤销(TokenVersion 自增),verify 复验须拒绝绑定
|
||||
nonce, message, err := w.Challenge(ctx, ethAddrOf(priv), "bind", "admin", adminToken(t, w))
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge(bind): %v", err)
|
||||
}
|
||||
if err := w.auth.bumpTokenVersion(ctx, "admin"); err != nil {
|
||||
t.Fatalf("bump token version: %v", err)
|
||||
}
|
||||
if _, err := w.Verify(ctx, nonce, signPersonal(priv, message, true), SessionMeta{ClientIP: "10.4.0.1"}); !errors.Is(err, ErrWalletChallenge) {
|
||||
t.Fatalf("stale-token bind err = %v, want ErrWalletChallenge", err)
|
||||
}
|
||||
if w.HasAny(ctx) {
|
||||
t.Error("HasAny = true, want false: stale-token bind must not persist identity")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWalletBindStaleVersion 验证事务内版本比对:令牌复验通过但发起后版本
|
||||
// 已变(极窄竞态窗口)的绑定同样被拒,身份不落库。
|
||||
func TestWalletBindStaleVersion(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
ctx := context.Background()
|
||||
priv := testPrivKey(9)
|
||||
nonce, message, err := w.Challenge(ctx, ethAddrOf(priv), "bind", "admin", adminToken(t, w))
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge(bind): %v", err)
|
||||
}
|
||||
w.mu.Lock()
|
||||
p := w.pending[nonce]
|
||||
p.proof.Ver--
|
||||
w.pending[nonce] = p
|
||||
w.mu.Unlock()
|
||||
if _, err := w.Verify(ctx, nonce, signPersonal(priv, message, true), SessionMeta{ClientIP: "10.4.0.2"}); !errors.Is(err, ErrWalletChallenge) {
|
||||
t.Fatalf("stale-version bind err = %v, want ErrWalletChallenge", err)
|
||||
}
|
||||
if w.HasAny(ctx) {
|
||||
t.Error("HasAny = true, want false: stale-version bind must not persist identity")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWalletBindKeepsSession 验证绑定接续当前会话:原行同行更新,
|
||||
// 保留登录方式、行 ID 与 jti,且新令牌有效。
|
||||
func TestWalletBindKeepsSession(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
ctx := context.Background()
|
||||
priv := testPrivKey(9)
|
||||
token := adminToken(t, w)
|
||||
before := sessionRows(t, w.auth)
|
||||
if len(before) != 1 || before[0].Method != "password" {
|
||||
t.Fatalf("seed session = %+v, want single password row", before)
|
||||
}
|
||||
nonce, message, err := w.Challenge(ctx, ethAddrOf(priv), "bind", "admin", token)
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge(bind): %v", err)
|
||||
}
|
||||
res, err := w.Verify(ctx, nonce, signPersonal(priv, message, true), SessionMeta{ClientIP: "10.1.0.1", UserAgent: "UA"})
|
||||
if err != nil {
|
||||
t.Fatalf("Verify(bind): %v", err)
|
||||
}
|
||||
after := sessionRows(t, w.auth)
|
||||
if len(after) != 1 || after[0].ID != before[0].ID {
|
||||
t.Fatalf("sessions after bind = %+v, want same single row (continuity)", after)
|
||||
}
|
||||
if after[0].Method != "password" || after[0].TokenID != before[0].TokenID {
|
||||
t.Errorf("renewed row = %+v, want method and jti kept", after[0])
|
||||
}
|
||||
if _, err := w.auth.ParseToken(ctx, res.Token); err != nil {
|
||||
t.Errorf("issued token invalid: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestWalletBindRevokedSession 验证绑定事务的 jti 复核:挑战发起后该设备
|
||||
// 会话被定点撤销(版本未变),验签通过的绑定仍须拒绝且身份不落库。
|
||||
func TestWalletBindRevokedSession(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
ctx := context.Background()
|
||||
priv := testPrivKey(9)
|
||||
token := adminToken(t, w)
|
||||
nonce, message, err := w.Challenge(ctx, ethAddrOf(priv), "bind", "admin", token)
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge(bind): %v", err)
|
||||
}
|
||||
// 模拟入口复验之后、事务之前的撤销:直接标记该 jti 的会话行 revoked
|
||||
// (入口 ParseToken 走缓存路径时可能尚未察觉,事务内查库是最终防线)
|
||||
w.auth.revokeSessionByJTI(ctx, "admin", w.auth.signedJti(token), tokenTTL)
|
||||
if _, err := w.Verify(ctx, nonce, signPersonal(priv, message, true), SessionMeta{ClientIP: "10.5.0.1"}); err == nil {
|
||||
t.Fatal("bind after session revoke succeeded, want rejection")
|
||||
}
|
||||
if w.HasAny(ctx) {
|
||||
t.Error("HasAny = true, want false: revoked-session bind must not persist identity")
|
||||
}
|
||||
}
|
||||
|
||||
// TestWalletBindNoRowMethodEmpty 验证旧版无会话行令牌绑定时,
|
||||
// 新建会话行的登录方式为空(该令牌并非经钱包登录,与活跃会话设计一致)。
|
||||
func TestWalletBindNoRowMethodEmpty(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
ctx := context.Background()
|
||||
priv := testPrivKey(9)
|
||||
token := adminToken(t, w)
|
||||
// 删除会话行,模拟升级前签发的存量令牌(无行)
|
||||
if err := w.db.Where("1 = 1").Delete(&model.UserSession{}).Error; err != nil {
|
||||
t.Fatalf("clear sessions: %v", err)
|
||||
}
|
||||
nonce, message, err := w.Challenge(ctx, ethAddrOf(priv), "bind", "admin", token)
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge(bind): %v", err)
|
||||
}
|
||||
res, err := w.Verify(ctx, nonce, signPersonal(priv, message, true), SessionMeta{ClientIP: "10.5.0.2", UserAgent: "UA"})
|
||||
if err != nil {
|
||||
t.Fatalf("Verify(bind): %v", err)
|
||||
}
|
||||
rows := sessionRows(t, w.auth)
|
||||
if len(rows) != 1 || rows[0].Method != "" {
|
||||
t.Fatalf("rows = %+v, want single row with empty method", rows)
|
||||
}
|
||||
if _, err := w.auth.ParseToken(ctx, res.Token); err != nil {
|
||||
t.Errorf("issued token invalid: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWalletLoginRechecksIdentity(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
bindWallet(t, w, 9)
|
||||
var row model.UserIdentity
|
||||
if err := w.db.Where("provider = ?", walletProvider).First(&row).Error; err != nil {
|
||||
t.Fatalf("find identity: %v", err)
|
||||
}
|
||||
if err := w.db.Delete(&row).Error; err != nil {
|
||||
t.Fatalf("delete identity: %v", err)
|
||||
}
|
||||
_, _, _, err := w.loginWalletIdentity(
|
||||
context.Background(), &row, row.Subject, SessionMeta{ClientIP: "10.6.0.1"})
|
||||
if !errors.Is(err, ErrWalletNotBound) {
|
||||
t.Fatalf("login err = %v, want ErrWalletNotBound", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWalletHasAnyRequiresAppURL(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
row := model.UserIdentity{UserID: 1, Provider: walletProvider, Subject: "0xabc"}
|
||||
if err := w.db.Create(&row).Error; err != nil {
|
||||
t.Fatalf("seed identity: %v", err)
|
||||
}
|
||||
if !w.HasAny(context.Background()) {
|
||||
t.Fatal("HasAny = false, want true with app url")
|
||||
}
|
||||
w.settings.SetEnvPublicURL("")
|
||||
if w.HasAny(context.Background()) {
|
||||
t.Fatal("HasAny = true without app url")
|
||||
}
|
||||
}
|
||||
|
||||
func TestWalletLoginGuardLock(t *testing.T) {
|
||||
w := newTestWallet(t, "https://demo.example.com")
|
||||
ctx := context.Background()
|
||||
priv := testPrivKey(9)
|
||||
addr := ethAddrOf(priv)
|
||||
other := testPrivKey(13)
|
||||
|
||||
var lastErr error
|
||||
for i := 0; i < securityDefaults.LoginFailLimit+1; i++ {
|
||||
nonce, message, err := w.Challenge(ctx, addr, "login", "", "")
|
||||
if err != nil {
|
||||
t.Fatalf("Challenge #%d: %v", i, err)
|
||||
}
|
||||
_, lastErr = w.Verify(ctx, nonce, signPersonal(other, message, true), SessionMeta{ClientIP: "10.3.0.1"})
|
||||
}
|
||||
if !errors.Is(lastErr, ErrLoginLocked) {
|
||||
t.Fatalf("after %d bad signatures err = %v, want ErrLoginLocked", securityDefaults.LoginFailLimit+1, lastErr)
|
||||
}
|
||||
// 锁定针对「IP+地址」:另一 IP 不连坐
|
||||
nonce, message, _ := w.Challenge(ctx, addr, "login", "", "")
|
||||
if _, err := w.Verify(ctx, nonce, signPersonal(other, message, true), SessionMeta{ClientIP: "10.3.0.2"}); errors.Is(err, ErrLoginLocked) {
|
||||
t.Errorf("different IP got locked prematurely: %v", err)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user