滑动续期、网络错误归类、API Key 激活乐观锁与生效提示
This commit is contained in:
@@ -39,10 +39,30 @@ func RequireAuth(auth *service.AuthService) gin.HandlerFunc {
|
||||
c.Set(usernameKey, username)
|
||||
c.Set(tokenVerKey, proof.Ver)
|
||||
c.Set(tokenJtiKey, proof.Jti)
|
||||
if allowsSessionRenewal(c.Request.Method) {
|
||||
maybeRenewToken(c, auth, token)
|
||||
}
|
||||
c.Next()
|
||||
}
|
||||
}
|
||||
|
||||
// allowsSessionRenewal 仅允许不改变认证状态的只读请求续期。写请求可能在
|
||||
// handler 内撤销会话或递增令牌版本,预生成的续期头到响应时会已经失效。
|
||||
func allowsSessionRenewal(method string) bool {
|
||||
return method == http.MethodGet || method == http.MethodHead
|
||||
}
|
||||
|
||||
// maybeRenewToken 滑动续期:临近过期的令牌换发同会话新令牌,经响应头透出,
|
||||
// 前端读到后无感替换本地会话;未到阈值时不加头。
|
||||
func maybeRenewToken(c *gin.Context, auth *service.AuthService, token string) {
|
||||
newToken, expires, ok := auth.MaybeRenew(c.Request.Context(), token)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
c.Header("X-Renewed-Token", newToken)
|
||||
c.Header("X-Renewed-Expires-At", expires.Format(time.RFC3339))
|
||||
}
|
||||
|
||||
// tokenProofOf 取出鉴权时的令牌快照,交给敏感 service 事务复核。
|
||||
func tokenProofOf(c *gin.Context) service.TokenProof {
|
||||
v, _ := c.Get(tokenVerKey)
|
||||
|
||||
@@ -283,11 +283,28 @@ func respondError(c *gin.Context, err error) {
|
||||
c.JSON(http.StatusUnauthorized, gin.H{"error": err.Error()})
|
||||
return
|
||||
}
|
||||
// 代理/网络类连接失败以 502 透出明确文案,不落无信息量的 500
|
||||
if msg, ok := oci.UpstreamNetworkError(err); ok {
|
||||
respondNetworkError(c, msg)
|
||||
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})
|
||||
}
|
||||
|
||||
// respondNetworkError 以 502 透出上游连接失败;服务端留日志,requestId 供关联。
|
||||
func respondNetworkError(c *gin.Context, msg string) {
|
||||
id := newRequestID()
|
||||
log.Printf("[NET %s] %s %s: %s", id, c.Request.Method, requestPath(c), msg)
|
||||
c.JSON(http.StatusBadGateway, gin.H{
|
||||
"error": msg,
|
||||
"hint": "上游连接失败(代理或网络),请检查该租户关联代理的可用性",
|
||||
"code": "UpstreamNetwork",
|
||||
"requestId": id,
|
||||
})
|
||||
}
|
||||
|
||||
// newRequestID 生成错误关联 ID(8 字节随机 hex),响应与服务端日志据此对应。
|
||||
func newRequestID() string {
|
||||
b := make([]byte, 8)
|
||||
|
||||
@@ -1,8 +1,14 @@
|
||||
package api
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
@@ -52,3 +58,30 @@ func TestRespondErrorOCIStatus(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRespondNetworkErrorHidesProxyCredentials(t *testing.T) {
|
||||
gin.SetMode(gin.TestMode)
|
||||
rec := httptest.NewRecorder()
|
||||
c, _ := gin.CreateTestContext(rec)
|
||||
c.Request = httptest.NewRequest(http.MethodGet, "/api/v1/oci-configs/1/instances", nil)
|
||||
proxyErr := &url.Error{
|
||||
Op: "Dial", URL: "http://proxy-user:proxy-pass@proxy.example.com:8080",
|
||||
Err: errors.New("connect: connection refused"),
|
||||
}
|
||||
err := &url.Error{Op: "Get", URL: "https://iaas.example.com/instances",
|
||||
Err: fmt.Errorf("proxyconnect tcp: %w", proxyErr)}
|
||||
var logs bytes.Buffer
|
||||
oldWriter := log.Writer()
|
||||
log.SetOutput(&logs)
|
||||
t.Cleanup(func() { log.SetOutput(oldWriter) })
|
||||
|
||||
respondError(c, err)
|
||||
if rec.Code != http.StatusBadGateway || !strings.Contains(logs.String(), "[NET ") {
|
||||
t.Fatalf("response/log = %d %q / %q", rec.Code, rec.Body.String(), logs.String())
|
||||
}
|
||||
for _, output := range []string{rec.Body.String(), logs.String()} {
|
||||
if strings.Contains(output, "proxy-user") || strings.Contains(output, "proxy-pass") {
|
||||
t.Fatalf("proxy credentials leaked: %q", output)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -8,9 +8,11 @@ import (
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gin-gonic/gin"
|
||||
"github.com/glebarez/sqlite"
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"gorm.io/gorm"
|
||||
"gorm.io/gorm/logger"
|
||||
|
||||
@@ -122,6 +124,30 @@ func doRequest(t *testing.T, r *gin.Engine, method, path, token, body string) *h
|
||||
return w
|
||||
}
|
||||
|
||||
type apiTestClaims struct {
|
||||
jwt.RegisteredClaims
|
||||
Ver uint `json:"ver"`
|
||||
}
|
||||
|
||||
func shortAPIToken(t *testing.T, token string) string {
|
||||
t.Helper()
|
||||
claims := &apiTestClaims{}
|
||||
_, err := jwt.ParseWithClaims(token, claims, func(*jwt.Token) (any, error) {
|
||||
return []byte("test-secret"), nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("parse login token: %v", err)
|
||||
}
|
||||
now := time.Now()
|
||||
claims.IssuedAt = jwt.NewNumericDate(now)
|
||||
claims.ExpiresAt = jwt.NewNumericDate(now.Add(time.Hour))
|
||||
short, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString([]byte("test-secret"))
|
||||
if err != nil {
|
||||
t.Fatalf("sign short token: %v", err)
|
||||
}
|
||||
return short
|
||||
}
|
||||
|
||||
func TestLoginEndpoint(t *testing.T) {
|
||||
r, _, _ := newTestRouter(t)
|
||||
tests := []struct {
|
||||
@@ -171,6 +197,26 @@ func TestSecuredRoutesRequireToken(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestReadRequestRenewsSession(t *testing.T) {
|
||||
r, auth, _ := newTestRouter(t)
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "",
|
||||
service.SessionMeta{ClientIP: "127.0.0.1"})
|
||||
if err != nil {
|
||||
t.Fatalf("login: %v", err)
|
||||
}
|
||||
w := doRequest(t, r, http.MethodGet, "/api/v1/auth/credentials", shortAPIToken(t, token), "")
|
||||
if w.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200, body %s", w.Code, w.Body.String())
|
||||
}
|
||||
renewed := w.Header().Get("X-Renewed-Token")
|
||||
if renewed == "" {
|
||||
t.Fatal("X-Renewed-Token 为空, want 只读请求续期")
|
||||
}
|
||||
if _, err := auth.ParseToken(context.Background(), renewed); err != nil {
|
||||
t.Errorf("renewed token invalid: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSystemLogsEndpoint(t *testing.T) {
|
||||
r, auth, logs := newTestRouter(t)
|
||||
token, _, err := auth.Login(context.Background(), "admin", "pass123", "", service.SessionMeta{ClientIP: "127.0.0.1"})
|
||||
@@ -245,11 +291,17 @@ func TestRevokeSessionsEndpoint(t *testing.T) {
|
||||
if err := json.Unmarshal(login.Body.Bytes(), &sess); err != nil || sess.Token == "" {
|
||||
t.Fatalf("login: %s", login.Body.String())
|
||||
}
|
||||
w := doRequest(t, r, http.MethodPost, "/api/v1/auth/revoke-sessions", sess.Token, "")
|
||||
short := shortAPIToken(t, sess.Token)
|
||||
w := doRequest(t, r, http.MethodPost, "/api/v1/auth/revoke-sessions", short, "")
|
||||
if w.Code != http.StatusOK || !strings.Contains(w.Body.String(), "token") {
|
||||
t.Fatalf("revoke = %d %s", w.Code, w.Body.String())
|
||||
}
|
||||
w = doRequest(t, r, http.MethodGet, "/api/v1/auth/credentials", sess.Token, "")
|
||||
for _, name := range []string{"X-Renewed-Token", "X-Renewed-Expires-At"} {
|
||||
if got := w.Header().Get(name); got != "" {
|
||||
t.Errorf("写请求 %s = %q, want empty", name, got)
|
||||
}
|
||||
}
|
||||
w = doRequest(t, r, http.MethodGet, "/api/v1/auth/credentials", short, "")
|
||||
if w.Code != http.StatusUnauthorized {
|
||||
t.Errorf("撤销后旧 token 访问 = %d, want 401", w.Code)
|
||||
}
|
||||
|
||||
@@ -209,6 +209,10 @@ type OciConfig struct {
|
||||
PrivateKeyEnc string `gorm:"type:text" json:"-"`
|
||||
PassphraseEnc string `gorm:"type:text" json:"-"`
|
||||
|
||||
// KeyActivatedAt 是最近一次启用新签名 key(面板轮换激活或手工替换私钥)
|
||||
// 的时刻;OCI 公钥全球传播为分钟级,前端据此在窗口期内提示,nil 表示未记录
|
||||
KeyActivatedAt *time.Time `json:"keyActivatedAt"`
|
||||
|
||||
TenancyName string `json:"tenancyName"`
|
||||
HomeRegionKey string `json:"homeRegionKey"`
|
||||
|
||||
|
||||
@@ -1,7 +1,11 @@
|
||||
package oci
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"regexp"
|
||||
"strings"
|
||||
|
||||
@@ -16,6 +20,7 @@ var ociErrorHints = map[string]string{
|
||||
"QuotaExceeded": "compartment 配额不足",
|
||||
"OutOfHostCapacity": "该可用域容量不足,可稍后重试或改用抢机任务",
|
||||
"TooManyRequests": "请求过于频繁,请稍后再试",
|
||||
"NotAuthenticated": "OCI 拒绝了请求签名:密钥无效或尚未生效。若刚替换过 API Key,全球生效需数分钟,请稍后重试",
|
||||
"InvalidParameter": "请求参数无效",
|
||||
"InternalError": "OCI 服务内部错误,请稍后重试",
|
||||
"Conflict": "资源正在变更中,请稍后重试",
|
||||
@@ -24,6 +29,9 @@ var ociErrorHints = map[string]string{
|
||||
// ocidRe 匹配消息中的完整 OCID(unique 段 20 位以上才压缩,避免误伤短标识)。
|
||||
var ocidRe = regexp.MustCompile(`ocid1\.([a-z0-9]+)\.[a-z0-9]*\.[a-z0-9-]*\.{1,2}[a-z0-9]{20,}`)
|
||||
|
||||
// proxyUserinfoRe 只匹配网络错误文本里位于 token / URL 起点的 userinfo@。
|
||||
var proxyUserinfoRe = regexp.MustCompile(`(^|[\s/("'=])[^@\s/]+@`)
|
||||
|
||||
// shortenOcids 把消息里的长 OCID 压缩为「ocid1.<类型>…<尾6位>」,
|
||||
// 保留资源类型与可比对的尾部,避免整条错误被 OCID 撑爆。
|
||||
func shortenOcids(s string) string {
|
||||
@@ -68,6 +76,67 @@ func ErrorHint(err error) string {
|
||||
return ociErrorHints[svcErr.GetCode()]
|
||||
}
|
||||
|
||||
// UpstreamNetworkError 判定错误链是否为上游连接层失败(代理拨号 / SOCKS 握手 /
|
||||
// 超时 / DNS 等),命中返回给用户看的精简描述。已拿到 OCI 响应的 ServiceError
|
||||
// 不在此列;判定保守:宁可漏归 500,不把业务错误误标为网络问题。
|
||||
func UpstreamNetworkError(err error) (string, bool) {
|
||||
var svcErr common.ServiceError
|
||||
if errors.As(err, &svcErr) {
|
||||
return "", false
|
||||
}
|
||||
var uerr *url.Error
|
||||
if errors.As(err, &uerr) {
|
||||
return classifyNetDetail(redactURLError(uerr)), true
|
||||
}
|
||||
if errors.Is(err, context.DeadlineExceeded) {
|
||||
return "上游连接失败: 请求超时", true
|
||||
}
|
||||
var nerr net.Error
|
||||
if errors.As(err, &nerr) && nerr.Timeout() {
|
||||
return "上游连接失败: 连接超时", true
|
||||
}
|
||||
return "", false
|
||||
}
|
||||
|
||||
// redactURLError 把 url.Error 的完整 URL 压缩为主机,并递归处理内层 url.Error。
|
||||
func redactURLError(uerr *url.Error) string {
|
||||
host := "<redacted-url>"
|
||||
if u, err := url.Parse(uerr.URL); err == nil && u.Host != "" {
|
||||
host = u.Host
|
||||
}
|
||||
return fmt.Sprintf("%s %s: %s", uerr.Op, host, redactURLCause(uerr.Err))
|
||||
}
|
||||
|
||||
// redactURLCause 保留普通包装前缀,但用脱敏文本替换其中嵌套的 url.Error。
|
||||
func redactURLCause(err error) string {
|
||||
if err == nil {
|
||||
return "unknown network error"
|
||||
}
|
||||
var nested *url.Error
|
||||
if !errors.As(err, &nested) {
|
||||
return redactProxyUserinfo(err.Error())
|
||||
}
|
||||
raw, nestedRaw := err.Error(), nested.Error()
|
||||
if strings.Contains(raw, nestedRaw) {
|
||||
raw = strings.Replace(raw, nestedRaw, redactURLError(nested), 1)
|
||||
return redactProxyUserinfo(raw)
|
||||
}
|
||||
return redactURLError(nested)
|
||||
}
|
||||
|
||||
// redactProxyUserinfo 清除存量非法代理 host 可能写入普通错误文本的 userinfo。
|
||||
func redactProxyUserinfo(detail string) string {
|
||||
return proxyUserinfoRe.ReplaceAllString(detail, "$1")
|
||||
}
|
||||
|
||||
// classifyNetDetail 按细节特征加中文分类前缀:代理链路失败与一般上游失败分开表述。
|
||||
func classifyNetDetail(detail string) string {
|
||||
if strings.Contains(detail, "proxyconnect") || strings.Contains(detail, "socks connect") {
|
||||
return "代理连接失败: " + detail
|
||||
}
|
||||
return "上游连接失败: " + detail
|
||||
}
|
||||
|
||||
// ServiceStatus 返回错误链中 OCI ServiceError 的 HTTP 状态码;非服务端错误 ok=false。
|
||||
func ServiceStatus(err error) (int, bool) {
|
||||
var svcErr common.ServiceError
|
||||
|
||||
@@ -1,8 +1,10 @@
|
||||
package oci
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net/url"
|
||||
"testing"
|
||||
)
|
||||
|
||||
@@ -95,6 +97,11 @@ func TestErrorHint(t *testing.T) {
|
||||
err: fmt.Errorf("launch: %w", fakeServiceError{500, "InternalError", "Out of host capacity."}),
|
||||
want: ociErrorHints["OutOfHostCapacity"],
|
||||
},
|
||||
{
|
||||
name: "NotAuthenticated 提示密钥无效或传播中",
|
||||
err: fmt.Errorf("get instance: %w", fakeServiceError{401, "NotAuthenticated", "The required information ..."}),
|
||||
want: ociErrorHints["NotAuthenticated"],
|
||||
},
|
||||
{
|
||||
name: "未知错误码返回空串",
|
||||
err: fmt.Errorf("x: %w", fakeServiceError{400, "SomethingNew", "boom"}),
|
||||
@@ -158,3 +165,106 @@ func TestIsModelUnavailable(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// fakeTimeoutErr 模拟实现 net.Error 的超时错误(如响应体读取超时)。
|
||||
type fakeTimeoutErr struct{}
|
||||
|
||||
func (fakeTimeoutErr) Error() string { return "read tcp 10.0.0.1:443: i/o timeout" }
|
||||
func (fakeTimeoutErr) Timeout() bool { return true }
|
||||
func (fakeTimeoutErr) Temporary() bool { return false }
|
||||
|
||||
type upstreamNetworkErrorCase struct {
|
||||
name string
|
||||
err error
|
||||
want string
|
||||
wantOK bool
|
||||
}
|
||||
|
||||
var upstreamNetworkErrorCases = []upstreamNetworkErrorCase{
|
||||
{
|
||||
name: "OCI 服务错误不归网络",
|
||||
err: fmt.Errorf("get instance: %w", fakeServiceError{status: 401, code: "NotAuthenticated", message: "x"}),
|
||||
wantOK: false,
|
||||
},
|
||||
{
|
||||
name: "普通业务错误不归网络",
|
||||
err: errors.New("parse payload: bad json"),
|
||||
wantOK: false,
|
||||
},
|
||||
{
|
||||
name: "代理 CONNECT 失败带中文分类且 URL 缩为主机",
|
||||
err: fmt.Errorf("list instances: %w", &url.Error{
|
||||
Op: "Get",
|
||||
URL: "https://iaas.uk-london-1.oraclecloud.com/20160918/instances?limit=100",
|
||||
Err: errors.New("proxyconnect tcp: dial tcp 1.2.3.4:8080: connect: connection refused"),
|
||||
}),
|
||||
want: "代理连接失败: Get iaas.uk-london-1.oraclecloud.com: proxyconnect tcp: dial tcp 1.2.3.4:8080: connect: connection refused",
|
||||
wantOK: true,
|
||||
},
|
||||
{
|
||||
name: "SOCKS 握手失败归代理类",
|
||||
err: &url.Error{
|
||||
Op: "Post",
|
||||
URL: "https://identity.us-sanjose-1.oci.oraclecloud.com/20160918/users",
|
||||
Err: errors.New("socks connect tcp 5.6.7.8:1080->identity: dial refused"),
|
||||
},
|
||||
want: "代理连接失败: Post identity.us-sanjose-1.oci.oraclecloud.com: socks connect tcp 5.6.7.8:1080->identity: dial refused",
|
||||
wantOK: true,
|
||||
},
|
||||
{
|
||||
name: "URL 内嵌 userinfo 不回显",
|
||||
err: &url.Error{
|
||||
Op: "Get",
|
||||
URL: "https://user:pass@example.com/path",
|
||||
Err: errors.New("EOF"),
|
||||
},
|
||||
want: "上游连接失败: Get example.com: EOF",
|
||||
wantOK: true,
|
||||
},
|
||||
{
|
||||
name: "嵌套 URL 的代理 userinfo 不回显",
|
||||
err: &url.Error{
|
||||
Op: "Get",
|
||||
URL: "https://identity.us-ashburn-1.oraclecloud.com/20160918/tenancies/x",
|
||||
Err: fmt.Errorf("proxyconnect tcp: %w", &url.Error{
|
||||
Op: "Dial", URL: "http://proxy-user:proxy-pass@proxy.example.com:8080",
|
||||
Err: errors.New("connect: connection refused"),
|
||||
}),
|
||||
},
|
||||
want: "代理连接失败: Get identity.us-ashburn-1.oraclecloud.com: proxyconnect tcp: Dial proxy.example.com:8080: connect: connection refused",
|
||||
wantOK: true,
|
||||
},
|
||||
{
|
||||
name: "普通内层错误中的存量代理 userinfo 不回显",
|
||||
err: &url.Error{
|
||||
Op: "Get",
|
||||
URL: "https://identity.us-phoenix-1.oraclecloud.com/20160918/tenancies/x",
|
||||
Err: errors.New("proxyconnect tcp: dial tcp alice:secret@proxy.example.com:8080: connect: connection refused"),
|
||||
},
|
||||
want: "代理连接失败: Get identity.us-phoenix-1.oraclecloud.com: proxyconnect tcp: dial tcp proxy.example.com:8080: connect: connection refused",
|
||||
wantOK: true,
|
||||
},
|
||||
{
|
||||
name: "context 超时归网络",
|
||||
err: fmt.Errorf("summarize costs: %w", context.DeadlineExceeded),
|
||||
want: "上游连接失败: 请求超时",
|
||||
wantOK: true,
|
||||
},
|
||||
{
|
||||
name: "net.Error 超时归网络",
|
||||
err: fmt.Errorf("read body: %w", fakeTimeoutErr{}),
|
||||
want: "上游连接失败: 连接超时",
|
||||
wantOK: true,
|
||||
},
|
||||
}
|
||||
|
||||
func TestUpstreamNetworkError(t *testing.T) {
|
||||
for _, tt := range upstreamNetworkErrorCases {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, ok := UpstreamNetworkError(tt.err)
|
||||
if ok != tt.wantOK || got != tt.want {
|
||||
t.Errorf("UpstreamNetworkError() = (%q, %v), want (%q, %v)", got, ok, tt.want, tt.wantOK)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -25,8 +25,12 @@ var ErrInvalidCredentials = errors.New("invalid username or password")
|
||||
// ErrLoginLocked 表示该 IP+用户名组合因连续失败被锁定;不提示剩余次数与时长细节。
|
||||
var ErrLoginLocked = errors.New("too many failed attempts, try again later")
|
||||
|
||||
// tokenTTL 是登录令牌有效期。
|
||||
const tokenTTL = 24 * time.Hour
|
||||
// tokenTTL 是登录令牌有效期;renewThreshold 是滑动续期阈值——剩余有效期
|
||||
// 低于该值的令牌在鉴权响应中自动换发同会话新令牌(见 MaybeRenew)。
|
||||
const (
|
||||
tokenTTL = 24 * time.Hour
|
||||
renewThreshold = tokenTTL / 2
|
||||
)
|
||||
|
||||
// authClaims 在标准声明外携带令牌版本;版本落后于账号当前值即失效。
|
||||
// 存量令牌无 ver 字段解析为 0,与存量账号的零值版本兼容(升级不强制登出)。
|
||||
|
||||
@@ -190,6 +190,9 @@ func (s *OciConfigService) applyCredentialUpdate(cfg *model.OciConfig, in Update
|
||||
return fmt.Errorf("encrypt private key: %w", err)
|
||||
}
|
||||
cfg.PrivateKeyEnc = enc
|
||||
// 手工替换私钥同样进入 OCI 公钥传播窗口,记录激活时刻供前端提示
|
||||
now := time.Now()
|
||||
cfg.KeyActivatedAt = &now
|
||||
}
|
||||
if in.Passphrase == nil {
|
||||
return nil
|
||||
|
||||
@@ -84,6 +84,9 @@ func validateProxyInput(in ProxyInput) error {
|
||||
if strings.TrimSpace(in.Host) == "" {
|
||||
return fmt.Errorf("主机不能为空: %w", ErrProxyInvalid)
|
||||
}
|
||||
if strings.Contains(in.Host, "@") {
|
||||
return fmt.Errorf("主机不可包含用户凭据: %w", ErrProxyInvalid)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -101,6 +101,7 @@ func TestProxyValidateAndDeleteInUse(t *testing.T) {
|
||||
{Name: "x", Type: "ss", Host: "h", Port: 1080},
|
||||
{Name: "x", Type: "http", Host: "h", Port: 0},
|
||||
{Name: "x", Type: "http", Host: " ", Port: 8080},
|
||||
{Name: "x", Type: "http", Host: "user:pass@proxy.example.com", Port: 8080},
|
||||
} {
|
||||
if _, err := svc.Create(ctx, in); err == nil {
|
||||
t.Fatalf("Create(%+v) accepted invalid input", in)
|
||||
|
||||
@@ -431,3 +431,42 @@ func (s *AuthService) cleanupSessionsOnce(ctx context.Context) {
|
||||
log.Printf("session cleanup: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// MaybeRenew 滑动续期:对剩余有效期不足 renewThreshold 的令牌换发同会话
|
||||
// (同 jti / 同版本)新令牌,并把会话行有效期延长到新过期点,不新建会话行。
|
||||
// 仅供鉴权通过后的请求调用(令牌有效性已由 ParseTokenProof 保证);
|
||||
// 无需换发或换发失败返回 ok=false,调用方跳过即可。
|
||||
func (s *AuthService) MaybeRenew(ctx context.Context, tokenString string) (string, time.Time, bool) {
|
||||
claims := &authClaims{}
|
||||
if _, err := jwt.ParseWithClaims(tokenString, claims, func(*jwt.Token) (any, error) {
|
||||
return s.jwtSecret, nil
|
||||
}); err != nil || claims.ExpiresAt == nil {
|
||||
return "", time.Time{}, false
|
||||
}
|
||||
if time.Until(claims.ExpiresAt.Time) >= renewThreshold {
|
||||
return "", time.Time{}, false
|
||||
}
|
||||
token, expires, _, err := s.signTokenWithJTI(claims.Subject, claims.Ver, claims.ID)
|
||||
if err != nil {
|
||||
return "", time.Time{}, false
|
||||
}
|
||||
if err := s.extendSessionExpiry(ctx, claims.ID, expires); err != nil {
|
||||
log.Printf("[WARN] %v", err)
|
||||
return "", time.Time{}, false
|
||||
}
|
||||
return token, expires, true
|
||||
}
|
||||
|
||||
// extendSessionExpiry 把会话行有效期延长到新过期点,保证「活跃会话」展示与
|
||||
// 清理任务看到真实过期时间;无行(存量令牌)静默跳过。
|
||||
func (s *AuthService) extendSessionExpiry(ctx context.Context, jti string, expires time.Time) error {
|
||||
if jti == "" {
|
||||
return nil
|
||||
}
|
||||
err := s.db.WithContext(ctx).Model(&model.UserSession{}).
|
||||
Where("token_id = ?", jti).UpdateColumn("expires_at", expires).Error
|
||||
if err != nil {
|
||||
return fmt.Errorf("extend session expiry: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang-jwt/jwt/v5"
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
@@ -523,3 +524,124 @@ func TestSessionCleanup(t *testing.T) {
|
||||
t.Errorf("after cleanup rows = %+v, want only alive", rows)
|
||||
}
|
||||
}
|
||||
|
||||
// signShortToken 用服务同款密钥手工签指定 TTL 的令牌,构造临近过期态。
|
||||
func signShortToken(t *testing.T, auth *AuthService, username string, ver uint, jti string, ttl time.Duration) string {
|
||||
t.Helper()
|
||||
claims := &authClaims{
|
||||
RegisteredClaims: jwt.RegisteredClaims{
|
||||
Subject: username,
|
||||
ID: jti,
|
||||
IssuedAt: jwt.NewNumericDate(time.Now()),
|
||||
ExpiresAt: jwt.NewNumericDate(time.Now().Add(ttl)),
|
||||
},
|
||||
Ver: ver,
|
||||
}
|
||||
tok, err := jwt.NewWithClaims(jwt.SigningMethodHS256, claims).SignedString(auth.jwtSecret)
|
||||
if err != nil {
|
||||
t.Fatalf("sign short token: %v", err)
|
||||
}
|
||||
return tok
|
||||
}
|
||||
|
||||
func renewalFixture(t *testing.T) (*AuthService, string, TokenProof) {
|
||||
t.Helper()
|
||||
auth := newTestAuth(t)
|
||||
if err := auth.EnsureAdmin("admin", "pass123"); err != nil {
|
||||
t.Fatalf("EnsureAdmin: %v", err)
|
||||
}
|
||||
ctx := context.Background()
|
||||
loginTok, _, err := auth.Login(ctx, "admin", "pass123", "",
|
||||
SessionMeta{ClientIP: "127.0.0.1", UserAgent: "t"})
|
||||
if err != nil {
|
||||
t.Fatalf("Login: %v", err)
|
||||
}
|
||||
_, proof, err := auth.ParseTokenProof(ctx, loginTok)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseTokenProof: %v", err)
|
||||
}
|
||||
return auth, loginTok, proof
|
||||
}
|
||||
|
||||
func setSessionExpiry(t *testing.T, auth *AuthService, jti string, expires time.Time) model.UserSession {
|
||||
t.Helper()
|
||||
if err := auth.db.Model(&model.UserSession{}).Where("token_id = ?", jti).
|
||||
UpdateColumn("expires_at", expires).Error; err != nil {
|
||||
t.Fatalf("set session expiry: %v", err)
|
||||
}
|
||||
var row model.UserSession
|
||||
if err := auth.db.Where("token_id = ?", jti).First(&row).Error; err != nil {
|
||||
t.Fatalf("find session row: %v", err)
|
||||
}
|
||||
return row
|
||||
}
|
||||
|
||||
func TestMaybeRenewEligibility(t *testing.T) {
|
||||
auth, loginTok, _ := renewalFixture(t)
|
||||
tests := []struct {
|
||||
name string
|
||||
token string
|
||||
}{
|
||||
{name: "剩余时间高于阈值", token: loginTok},
|
||||
{name: "非法令牌", token: "not.a.token"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if _, _, ok := auth.MaybeRenew(context.Background(), tt.token); ok {
|
||||
t.Error("MaybeRenew ok = true, want false")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaybeRenewExtendsSession(t *testing.T) {
|
||||
auth, _, proof := renewalFixture(t)
|
||||
before := setSessionExpiry(t, auth, proof.Jti, time.Now().Add(time.Hour))
|
||||
short := signShortToken(t, auth, "admin", proof.Ver, proof.Jti, time.Hour)
|
||||
renewed, expires, ok := auth.MaybeRenew(context.Background(), short)
|
||||
if !ok {
|
||||
t.Fatal("剩余 1h 的令牌应换发")
|
||||
}
|
||||
_, newProof, err := auth.ParseTokenProof(context.Background(), renewed)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseTokenProof(renewed): %v", err)
|
||||
}
|
||||
if newProof.Jti != proof.Jti || newProof.Ver != proof.Ver {
|
||||
t.Errorf("proof = %+v, want %+v", newProof, proof)
|
||||
}
|
||||
var after model.UserSession
|
||||
if err := auth.db.Where("token_id = ?", proof.Jti).First(&after).Error; err != nil {
|
||||
t.Fatalf("reload session row: %v", err)
|
||||
}
|
||||
if delta := after.ExpiresAt.Sub(before.ExpiresAt); delta < 22*time.Hour {
|
||||
t.Errorf("session expiry delta = %v, want >= 22h", delta)
|
||||
}
|
||||
if gap := after.ExpiresAt.Sub(expires); gap < -time.Second || gap > time.Second {
|
||||
t.Errorf("session expiry = %v, token expiry = %v", after.ExpiresAt, expires)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaybeRenewedTokenFollowsSessionRevocation(t *testing.T) {
|
||||
auth, loginTok, proof := renewalFixture(t)
|
||||
short := signShortToken(t, auth, "admin", proof.Ver, proof.Jti, time.Hour)
|
||||
renewed, _, ok := auth.MaybeRenew(context.Background(), short)
|
||||
if !ok {
|
||||
t.Fatal("剩余 1h 的令牌应换发")
|
||||
}
|
||||
auth.Logout(context.Background(), loginTok)
|
||||
if _, err := auth.ParseToken(context.Background(), renewed); err == nil {
|
||||
t.Error("会话撤销后换发令牌仍有效, want 失效")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMaybeRenewUpdateFailureDoesNotIssueToken(t *testing.T) {
|
||||
auth, _, proof := renewalFixture(t)
|
||||
short := signShortToken(t, auth, "admin", proof.Ver, proof.Jti, time.Hour)
|
||||
if err := auth.db.Migrator().DropTable(&model.UserSession{}); err != nil {
|
||||
t.Fatalf("drop sessions table: %v", err)
|
||||
}
|
||||
token, expires, ok := auth.MaybeRenew(context.Background(), short)
|
||||
if ok || token != "" || !expires.IsZero() {
|
||||
t.Errorf("MaybeRenew = (%q, %v, %v), want empty result", token, expires, ok)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -10,6 +10,8 @@ import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
"oci-portal/internal/oci"
|
||||
)
|
||||
@@ -106,7 +108,7 @@ func (s *OciConfigService) ActivateApiKey(ctx context.Context, id uint, userID,
|
||||
if err := s.waitApiKeyUsable(ctx, newCred); err != nil {
|
||||
return err
|
||||
}
|
||||
return s.persistSigningKey(cfg, newCred)
|
||||
return s.persistSigningKey(ctx, cfg, newCred)
|
||||
}
|
||||
|
||||
// waitApiKeyUsable 用新凭据测活,等待上传的公钥在 OCI 侧生效。
|
||||
@@ -126,7 +128,7 @@ func (s *OciConfigService) waitApiKeyUsable(ctx context.Context, cred oci.Creden
|
||||
}
|
||||
|
||||
// persistSigningKey 加密新私钥,更新配置签名用户与指纹并清空口令密文(面板生成的 key 无口令)。
|
||||
func (s *OciConfigService) persistSigningKey(cfg *model.OciConfig, newCred oci.Credentials) error {
|
||||
func (s *OciConfigService) persistSigningKey(ctx context.Context, cfg *model.OciConfig, newCred oci.Credentials) error {
|
||||
enc, err := s.cipher.EncryptString(newCred.PrivateKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encrypt private key: %w", err)
|
||||
@@ -134,9 +136,16 @@ func (s *OciConfigService) persistSigningKey(cfg *model.OciConfig, newCred oci.C
|
||||
updates := map[string]any{
|
||||
"user_oc_id": newCred.UserOCID, "fingerprint": newCred.Fingerprint,
|
||||
"private_key_enc": enc, "passphrase_enc": "",
|
||||
// 记录激活时刻:OCI 公钥全球传播为分钟级,前端据此做窗口期提示
|
||||
"key_activated_at": time.Now(),
|
||||
}
|
||||
if err := s.db.Model(cfg).Updates(updates).Error; err != nil {
|
||||
return fmt.Errorf("persist rotated key: %w", err)
|
||||
res := s.db.WithContext(ctx).Model(&model.OciConfig{}).
|
||||
Where("id = ? AND updated_at = ?", cfg.ID, cfg.UpdatedAt).Updates(updates)
|
||||
if res.Error != nil {
|
||||
return fmt.Errorf("persist rotated key: %w", res.Error)
|
||||
}
|
||||
if res.RowsAffected != 1 {
|
||||
return fmt.Errorf("persist rotated key: %w", gorm.ErrRecordNotFound)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -9,6 +9,8 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"gorm.io/gorm"
|
||||
|
||||
"oci-portal/internal/model"
|
||||
"oci-portal/internal/oci"
|
||||
)
|
||||
@@ -194,6 +196,9 @@ func TestActivateApiKey(t *testing.T) {
|
||||
if err != nil || plain != newKey {
|
||||
t.Fatalf("persisted key mismatch (err=%v)", err)
|
||||
}
|
||||
if got.KeyActivatedAt == nil || time.Since(*got.KeyActivatedAt) > time.Minute {
|
||||
t.Fatalf("keyActivatedAt = %v, want 刚写入的时间", got.KeyActivatedAt)
|
||||
}
|
||||
if len(fc.validated) == 0 || fc.validated[0] != "11:22" {
|
||||
t.Fatalf("validated = %v", fc.validated)
|
||||
}
|
||||
@@ -203,3 +208,88 @@ func TestActivateApiKey(t *testing.T) {
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPersistSigningKeyRejectsStaleSnapshot(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
deleted bool
|
||||
}{
|
||||
{name: "租户已删除", deleted: true},
|
||||
{name: "凭据已被并发更新"},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
testStaleSigningKeyPersistence(t, tt.deleted)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestPersistSigningKeyAdvancesSnapshotVersion(t *testing.T) {
|
||||
s := newTestService(t, &apiKeyClient{})
|
||||
cfg := seedApiKeyConfig(t, s)
|
||||
stale := *cfg
|
||||
first := oci.Credentials{UserOCID: cfg.UserOCID, Fingerprint: "11:22", PrivateKey: "first-key"}
|
||||
if err := s.persistSigningKey(context.Background(), cfg, first); err != nil {
|
||||
t.Fatalf("first persist: %v", err)
|
||||
}
|
||||
second := oci.Credentials{UserOCID: cfg.UserOCID, Fingerprint: "33:44", PrivateKey: "second-key"}
|
||||
if err := s.persistSigningKey(context.Background(), &stale, second); !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("stale persist err = %v, want ErrRecordNotFound", err)
|
||||
}
|
||||
var got model.OciConfig
|
||||
if err := s.db.First(&got, cfg.ID).Error; err != nil || got.Fingerprint != "11:22" {
|
||||
t.Fatalf("persisted config = %+v, err %v", got, err)
|
||||
}
|
||||
}
|
||||
|
||||
func testStaleSigningKeyPersistence(t *testing.T, deleted bool) {
|
||||
t.Helper()
|
||||
s := newTestService(t, &apiKeyClient{})
|
||||
cfg := seedApiKeyConfig(t, s)
|
||||
invalidateSigningSnapshot(t, s, cfg, deleted)
|
||||
cred := oci.Credentials{UserOCID: cfg.UserOCID, Fingerprint: "11:22", PrivateKey: "new-key"}
|
||||
err := s.persistSigningKey(context.Background(), cfg, cred)
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("err = %v, want ErrRecordNotFound", err)
|
||||
}
|
||||
assertSigningSnapshotPreserved(t, s, cfg.ID, deleted)
|
||||
}
|
||||
|
||||
func invalidateSigningSnapshot(t *testing.T, s *OciConfigService, cfg *model.OciConfig, deleted bool) {
|
||||
t.Helper()
|
||||
if deleted {
|
||||
if err := s.db.Delete(&model.OciConfig{}, cfg.ID).Error; err != nil {
|
||||
t.Fatalf("delete config: %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
enc, err := s.cipher.EncryptString("concurrent-key")
|
||||
if err != nil {
|
||||
t.Fatalf("encrypt concurrent key: %v", err)
|
||||
}
|
||||
res := s.db.Model(&model.OciConfig{}).Where("id = ?", cfg.ID).UpdateColumns(map[string]any{
|
||||
"fingerprint": "cc:dd", "private_key_enc": enc, "updated_at": cfg.UpdatedAt.Add(time.Second),
|
||||
})
|
||||
if res.Error != nil || res.RowsAffected != 1 {
|
||||
t.Fatalf("mutate config = rows %d, err %v", res.RowsAffected, res.Error)
|
||||
}
|
||||
}
|
||||
|
||||
func assertSigningSnapshotPreserved(t *testing.T, s *OciConfigService, id uint, deleted bool) {
|
||||
t.Helper()
|
||||
var got model.OciConfig
|
||||
err := s.db.First(&got, id).Error
|
||||
if deleted {
|
||||
if !errors.Is(err, gorm.ErrRecordNotFound) {
|
||||
t.Fatalf("deleted config reload err = %v", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil || got.Fingerprint != "cc:dd" {
|
||||
t.Fatalf("concurrent config = %+v, err %v", got, err)
|
||||
}
|
||||
plain, err := s.cipher.DecryptString(got.PrivateKeyEnc)
|
||||
if err != nil || plain != "concurrent-key" {
|
||||
t.Fatalf("concurrent key = %q, err %v", plain, err)
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user