fix(auth): 消除 API Key 校验连接池死锁
先完整收集同前缀候选项并关闭查询结果,再执行 bcrypt 比对和 last_used_at 更新,避免小连接池下查询与更新相互等待。 新增 Rows 关闭顺序、前缀碰撞、MaxConns=1 和 8 并发真实 PostgreSQL 回归测试。
This commit is contained in:
@@ -0,0 +1,169 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/easyai/easyai-ai-gateway/apps/api/internal/auth"
|
||||
"github.com/jackc/pgx/v5"
|
||||
"github.com/jackc/pgx/v5/pgconn"
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
|
||||
func TestVerifyLocalAPIKeyClosesCandidateRowsBeforeUpdatingUsage(t *testing.T) {
|
||||
secret := "sk-gw-matching-secret"
|
||||
wrongHash, err := bcrypt.GenerateFromPassword([]byte("sk-gw-different-secret"), bcrypt.MinCost)
|
||||
if err != nil {
|
||||
t.Fatalf("hash non-matching API key: %v", err)
|
||||
}
|
||||
matchingHash, err := bcrypt.GenerateFromPassword([]byte(secret), bcrypt.MinCost)
|
||||
if err != nil {
|
||||
t.Fatalf("hash matching API key: %v", err)
|
||||
}
|
||||
rows := &fakeLocalAPIKeyRows{candidates: []localAPIKeyCandidate{
|
||||
{apiKeyID: "wrong-key", hash: string(wrongHash), keyPrefix: apiKeyPrefix(secret)},
|
||||
{
|
||||
apiKeyID: "matching-key",
|
||||
hash: string(matchingHash),
|
||||
keyPrefix: apiKeyPrefix(secret),
|
||||
keyName: "Matching key",
|
||||
scopesBytes: []byte(`["chat"]`),
|
||||
userGroupID: "group-id",
|
||||
gatewayUserID: "user-id",
|
||||
username: "api-key-user",
|
||||
rolesBytes: []byte(`["user"]`),
|
||||
gatewayTenantID: "gateway-tenant-id",
|
||||
tenantID: "tenant-id",
|
||||
tenantKey: "tenant-key",
|
||||
},
|
||||
}}
|
||||
database := &fakeLocalAPIKeyDatabase{rows: rows}
|
||||
|
||||
user, err := verifyLocalAPIKey(context.Background(), database, secret)
|
||||
if err != nil {
|
||||
t.Fatalf("verify local API key: %v", err)
|
||||
}
|
||||
if !rows.closed {
|
||||
t.Fatal("candidate rows remained open after API key verification")
|
||||
}
|
||||
if database.updatedAPIKeyID != "matching-key" {
|
||||
t.Fatalf("updated API key = %q, want matching-key", database.updatedAPIKeyID)
|
||||
}
|
||||
if user.APIKeyID != "matching-key" || user.GatewayUserID != "user-id" {
|
||||
t.Fatalf("verified user = %+v", user)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyLocalAPIKeyReturnsUnauthorizedAfterClosingCandidateRows(t *testing.T) {
|
||||
secret := "sk-gw-unknown-secret"
|
||||
wrongHash, err := bcrypt.GenerateFromPassword([]byte("sk-gw-different-secret"), bcrypt.MinCost)
|
||||
if err != nil {
|
||||
t.Fatalf("hash non-matching API key: %v", err)
|
||||
}
|
||||
rows := &fakeLocalAPIKeyRows{candidates: []localAPIKeyCandidate{{
|
||||
apiKeyID: "wrong-key",
|
||||
hash: string(wrongHash),
|
||||
}}}
|
||||
database := &fakeLocalAPIKeyDatabase{rows: rows}
|
||||
|
||||
_, err = verifyLocalAPIKey(context.Background(), database, secret)
|
||||
if !errors.Is(err, auth.ErrUnauthorized) {
|
||||
t.Fatalf("verify error = %v, want unauthorized", err)
|
||||
}
|
||||
if !rows.closed {
|
||||
t.Fatal("candidate rows remained open after unsuccessful API key verification")
|
||||
}
|
||||
if database.updatedAPIKeyID != "" {
|
||||
t.Fatalf("unexpected API key usage update for %q", database.updatedAPIKeyID)
|
||||
}
|
||||
}
|
||||
|
||||
type fakeLocalAPIKeyDatabase struct {
|
||||
rows *fakeLocalAPIKeyRows
|
||||
updatedAPIKeyID string
|
||||
}
|
||||
|
||||
func (database *fakeLocalAPIKeyDatabase) Query(context.Context, string, ...any) (pgx.Rows, error) {
|
||||
return database.rows, nil
|
||||
}
|
||||
|
||||
func (database *fakeLocalAPIKeyDatabase) Exec(_ context.Context, _ string, arguments ...any) (pgconn.CommandTag, error) {
|
||||
if !database.rows.closed {
|
||||
return pgconn.CommandTag{}, errors.New("API key usage update started before candidate rows closed")
|
||||
}
|
||||
database.updatedAPIKeyID, _ = arguments[0].(string)
|
||||
return pgconn.NewCommandTag("UPDATE 1"), nil
|
||||
}
|
||||
|
||||
type fakeLocalAPIKeyRows struct {
|
||||
candidates []localAPIKeyCandidate
|
||||
current int
|
||||
closed bool
|
||||
}
|
||||
|
||||
func (rows *fakeLocalAPIKeyRows) Close() {
|
||||
rows.closed = true
|
||||
}
|
||||
|
||||
func (rows *fakeLocalAPIKeyRows) Err() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (rows *fakeLocalAPIKeyRows) CommandTag() pgconn.CommandTag {
|
||||
return pgconn.CommandTag{}
|
||||
}
|
||||
|
||||
func (rows *fakeLocalAPIKeyRows) FieldDescriptions() []pgconn.FieldDescription {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (rows *fakeLocalAPIKeyRows) Next() bool {
|
||||
if rows.current >= len(rows.candidates) {
|
||||
rows.Close()
|
||||
return false
|
||||
}
|
||||
rows.current++
|
||||
return true
|
||||
}
|
||||
|
||||
func (rows *fakeLocalAPIKeyRows) Scan(destinations ...any) error {
|
||||
candidate := rows.candidates[rows.current-1]
|
||||
values := []any{
|
||||
candidate.apiKeyID,
|
||||
candidate.hash,
|
||||
candidate.keyPrefix,
|
||||
candidate.keyName,
|
||||
candidate.scopesBytes,
|
||||
candidate.userGroupID,
|
||||
candidate.gatewayUserID,
|
||||
candidate.username,
|
||||
candidate.rolesBytes,
|
||||
candidate.gatewayTenantID,
|
||||
candidate.tenantID,
|
||||
candidate.tenantKey,
|
||||
}
|
||||
for index, value := range values {
|
||||
switch destination := destinations[index].(type) {
|
||||
case *string:
|
||||
*destination = value.(string)
|
||||
case *[]byte:
|
||||
*destination = value.([]byte)
|
||||
default:
|
||||
return errors.New("unsupported fake row destination")
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (rows *fakeLocalAPIKeyRows) Values() ([]any, error) {
|
||||
return nil, errors.New("not implemented")
|
||||
}
|
||||
|
||||
func (rows *fakeLocalAPIKeyRows) RawValues() [][]byte {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (rows *fakeLocalAPIKeyRows) Conn() *pgx.Conn {
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user