Files
easyai-ai-gateway/apps/api/internal/store/api_key_auth_test.go
T
wangbo 98820378b7 perf(store): 消除同步复制下的鉴权与心跳锁阻塞
将 Worker 心跳与容量分配事务设为本地异步提交,避免同步副本延迟期间长期持有全局分配锁。API Key 使用时间改为异步提交并按分钟合并,消除高并发请求对同一热行的锁排队。业务任务、钱包、结算、并发租约等关键数据仍保持同步复制。验证通过完整 Go 测试、go vet、竞态测试、迁移安全检查及 PostgreSQL 18 集成测试。
2026-07-29 23:25:54 +08:00

179 lines
5.0 KiB
Go

package store
import (
"context"
"errors"
"strings"
"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 !strings.Contains(database.updateSQL, "set_config('synchronous_commit', 'off', true)") {
t.Fatalf("API key usage update did not disable synchronous commit: %s", database.updateSQL)
}
if !strings.Contains(database.updateSQL, "interval '1 minute'") {
t.Fatalf("API key usage update did not coalesce hot-key writes: %s", database.updateSQL)
}
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
updateSQL string
}
func (database *fakeLocalAPIKeyDatabase) Query(context.Context, string, ...any) (pgx.Rows, error) {
return database.rows, nil
}
func (database *fakeLocalAPIKeyDatabase) Exec(_ context.Context, sql 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.updateSQL = sql
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
}