diff --git a/apps/api/docs/swagger.json b/apps/api/docs/swagger.json index 4a77aad..bc5222c 100644 --- a/apps/api/docs/swagger.json +++ b/apps/api/docs/swagger.json @@ -5208,6 +5208,12 @@ "schema": { "$ref": "#/definitions/httpapi.ErrorEnvelope" } + }, + "503": { + "description": "Service Unavailable", + "schema": { + "$ref": "#/definitions/httpapi.ErrorEnvelope" + } } } } diff --git a/apps/api/docs/swagger.yaml b/apps/api/docs/swagger.yaml index cb7e255..b7271b8 100644 --- a/apps/api/docs/swagger.yaml +++ b/apps/api/docs/swagger.yaml @@ -6441,6 +6441,10 @@ paths: description: Internal Server Error schema: $ref: '#/definitions/httpapi.ErrorEnvelope' + "503": + description: Service Unavailable + schema: + $ref: '#/definitions/httpapi.ErrorEnvelope' summary: 本地登录 tags: - auth diff --git a/apps/api/internal/httpapi/availability_timeout_integration_test.go b/apps/api/internal/httpapi/availability_timeout_integration_test.go new file mode 100644 index 0000000..964627c --- /dev/null +++ b/apps/api/internal/httpapi/availability_timeout_integration_test.go @@ -0,0 +1,111 @@ +package httpapi + +import ( + "bytes" + "context" + "encoding/json" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "net/url" + "os" + "strings" + "testing" + "time" + + "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" +) + +func TestReadyReturnsPostgresUnavailableWithinTwoSeconds(t *testing.T) { + db := newExhaustedPostgresStore(t) + server := &Server{store: db, logger: slog.New(slog.NewJSONHandler(io.Discard, nil))} + requestContext, cancel := context.WithTimeout(context.Background(), 4*time.Second) + defer cancel() + request := httptest.NewRequest(http.MethodGet, "/readyz", nil).WithContext(requestContext) + recorder := httptest.NewRecorder() + + startedAt := time.Now() + server.ready(recorder, request) + elapsed := time.Since(startedAt) + + assertUnavailableResponse(t, recorder, "POSTGRES_UNAVAILABLE", "postgres unavailable") + if elapsed > 3*time.Second { + t.Fatalf("readiness timeout took %s, want no more than 3s", elapsed) + } +} + +func TestLoginReturnsAuthStoreUnavailableWithinFiveSeconds(t *testing.T) { + db := newExhaustedPostgresStore(t) + var logs bytes.Buffer + server := &Server{store: db, logger: slog.New(slog.NewJSONHandler(&logs, nil))} + requestContext, cancel := context.WithTimeout(context.Background(), 7*time.Second) + defer cancel() + request := httptest.NewRequest(http.MethodPost, "/api/v1/auth/login", strings.NewReader(`{"account":"timeout-test-account","password":"timeout-test-password"}`)).WithContext(requestContext) + recorder := httptest.NewRecorder() + + startedAt := time.Now() + server.login(recorder, request) + elapsed := time.Since(startedAt) + + assertUnavailableResponse(t, recorder, "AUTH_STORE_UNAVAILABLE", "authentication service temporarily unavailable") + if elapsed > 6*time.Second { + t.Fatalf("login timeout took %s, want no more than 6s", elapsed) + } + logOutput := logs.String() + for _, field := range []string{"postgres_pool_max_connections", "postgres_pool_acquired_connections", "postgres_pool_idle_connections", "postgres_pool_empty_acquire_count", "postgres_pool_canceled_acquire_count"} { + if !strings.Contains(logOutput, field) { + t.Fatalf("login failure log did not include %q: %s", field, logOutput) + } + } + if strings.Contains(logOutput, "timeout-test-account") || strings.Contains(logOutput, "timeout-test-password") { + t.Fatalf("login failure log exposed credentials: %s", logOutput) + } +} + +func newExhaustedPostgresStore(t *testing.T) *store.Store { + t.Helper() + databaseURL := strings.TrimSpace(os.Getenv("AI_GATEWAY_TEST_DATABASE_URL")) + if databaseURL == "" { + t.Skip("set AI_GATEWAY_TEST_DATABASE_URL to run PostgreSQL availability timeout tests") + } + parsed, err := url.Parse(databaseURL) + if err != nil { + t.Fatalf("parse test database URL: %v", err) + } + query := parsed.Query() + query.Set("pool_max_conns", "1") + query.Set("pool_min_conns", "0") + parsed.RawQuery = query.Encode() + + ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second) + defer cancel() + db, err := store.Connect(ctx, parsed.String()) + if err != nil { + t.Fatalf("connect timeout test store: %v", err) + } + t.Cleanup(db.Close) + connection, err := db.Pool().Acquire(ctx) + if err != nil { + t.Fatalf("exhaust timeout test pool: %v", err) + } + t.Cleanup(connection.Release) + return db +} + +func assertUnavailableResponse(t *testing.T, recorder *httptest.ResponseRecorder, expectedCode, expectedMessage string) { + t.Helper() + if recorder.Code != http.StatusServiceUnavailable { + t.Fatalf("status = %d, want 503; body=%s", recorder.Code, recorder.Body.String()) + } + var envelope ErrorEnvelope + if err := json.Unmarshal(recorder.Body.Bytes(), &envelope); err != nil { + t.Fatalf("decode unavailable response: %v", err) + } + if envelope.Error.Code != expectedCode { + t.Fatalf("error code = %q, want %q; body=%s", envelope.Error.Code, expectedCode, recorder.Body.String()) + } + if envelope.Error.Message != expectedMessage { + t.Fatalf("error message = %q, want %q; body=%s", envelope.Error.Message, expectedMessage, recorder.Body.String()) + } +} diff --git a/apps/api/internal/httpapi/handlers.go b/apps/api/internal/httpapi/handlers.go index 4d510f3..0e29b08 100644 --- a/apps/api/internal/httpapi/handlers.go +++ b/apps/api/internal/httpapi/handlers.go @@ -17,6 +17,14 @@ import ( "github.com/easyai/easyai-ai-gateway/apps/api/internal/store" ) +const ( + postgresReadinessTimeout = 2 * time.Second + localLoginStoreTimeout = 5 * time.Second + errorCodePostgresDown = "POSTGRES_UNAVAILABLE" + errorCodeAuthStoreDown = "AUTH_STORE_UNAVAILABLE" + authStoreUnavailableError = "authentication service temporarily unavailable" +) + // health godoc // @Summary 健康检查 // @Description 返回服务进程、运行环境和身份模式,供负载均衡或人工排障使用。 @@ -42,8 +50,11 @@ func (s *Server) health(w http.ResponseWriter, r *http.Request) { // @Failure 503 {object} ErrorEnvelope // @Router /readyz [get] func (s *Server) ready(w http.ResponseWriter, r *http.Request) { - if err := s.store.Ping(r.Context()); err != nil { - writeError(w, http.StatusServiceUnavailable, "postgres unavailable") + ctx, cancel := context.WithTimeout(r.Context(), postgresReadinessTimeout) + defer cancel() + if err := s.store.Ping(ctx); err != nil { + s.logPostgresUnavailable("postgres readiness check failed") + writeError(w, http.StatusServiceUnavailable, "postgres unavailable", errorCodePostgresDown) return } writeJSON(w, http.StatusOK, map[string]any{"ok": true}) @@ -121,6 +132,7 @@ func (s *Server) register(w http.ResponseWriter, r *http.Request) { // @Failure 401 {object} ErrorEnvelope // @Failure 403 {object} ErrorEnvelope // @Failure 500 {object} ErrorEnvelope +// @Failure 503 {object} ErrorEnvelope // @Router /api/v1/auth/login [post] func (s *Server) login(w http.ResponseWriter, r *http.Request) { var input store.LocalLoginInput @@ -128,12 +140,19 @@ func (s *Server) login(w http.ResponseWriter, r *http.Request) { writeError(w, http.StatusBadRequest, "invalid json body") return } - user, err := s.store.AuthenticateLocalUser(r.Context(), input) + ctx, cancel := context.WithTimeout(r.Context(), localLoginStoreTimeout) + defer cancel() + user, err := s.store.AuthenticateLocalUser(ctx, input) if err != nil { if errors.Is(err, store.ErrInvalidCredentials) { writeError(w, http.StatusUnauthorized, "invalid account or password") return } + if store.IsPostgresUnavailable(err) { + s.logPostgresUnavailable("login authentication store unavailable") + writeError(w, http.StatusServiceUnavailable, authStoreUnavailableError, errorCodeAuthStoreDown) + return + } s.logger.Error("login local user failed", "error", err) writeError(w, http.StatusInternalServerError, "login failed") return @@ -145,6 +164,22 @@ func (s *Server) login(w http.ResponseWriter, r *http.Request) { s.writeAuthResponse(w, http.StatusOK, user) } +func (s *Server) logPostgresUnavailable(message string) { + if s.logger == nil || s.store == nil || s.store.Pool() == nil { + return + } + statistics := s.store.Pool().Stat() + s.logger.Error(message, + "error_category", "postgres_unavailable", + "postgres_pool_max_connections", statistics.MaxConns(), + "postgres_pool_total_connections", statistics.TotalConns(), + "postgres_pool_acquired_connections", statistics.AcquiredConns(), + "postgres_pool_idle_connections", statistics.IdleConns(), + "postgres_pool_empty_acquire_count", statistics.EmptyAcquireCount(), + "postgres_pool_canceled_acquire_count", statistics.CanceledAcquireCount(), + ) +} + func (s *Server) localIdentityEnabled() bool { mode := strings.ToLower(strings.TrimSpace(s.cfg.IdentityMode)) return mode == "" || mode == "standalone" || mode == "hybrid" diff --git a/apps/api/internal/store/api_key_auth_integration_test.go b/apps/api/internal/store/api_key_auth_integration_test.go new file mode 100644 index 0000000..8247c5e --- /dev/null +++ b/apps/api/internal/store/api_key_auth_integration_test.go @@ -0,0 +1,99 @@ +package store + +import ( + "context" + "sync" + "testing" + "time" + + "github.com/easyai/easyai-ai-gateway/apps/api/internal/auth" + "github.com/jackc/pgx/v5/pgxpool" +) + +func TestVerifyLocalAPIKeyWorksWithSingleConnectionPool(t *testing.T) { + db, verificationStore, created, user := newLocalAPIKeyVerificationFixture(t, 1) + ctx := context.Background() + + verifyCtx, cancel := context.WithTimeout(ctx, 2*time.Second) + defer cancel() + verified, err := verificationStore.VerifyLocalAPIKey(verifyCtx, created.Secret) + if err != nil { + t.Fatalf("verify local API key with one connection: %v", err) + } + if verified.APIKeyID != created.APIKey.ID || verified.GatewayUserID != user.ID { + t.Fatalf("verified identity = %+v, want API key %q and user %q", verified, created.APIKey.ID, user.ID) + } + + var lastUsedAt *time.Time + if err := db.pool.QueryRow(ctx, `SELECT last_used_at FROM gateway_api_keys WHERE id=$1::uuid`, created.APIKey.ID).Scan(&lastUsedAt); err != nil { + t.Fatalf("read API key last_used_at: %v", err) + } + if lastUsedAt == nil { + t.Fatal("successful API key verification did not update last_used_at") + } +} + +func TestVerifyLocalAPIKeyHandlesEightConcurrentRequestsWithFourConnections(t *testing.T) { + _, verificationStore, created, _ := newLocalAPIKeyVerificationFixture(t, 4) + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + + const requestCount = 8 + start := make(chan struct{}) + errorsByRequest := make(chan error, requestCount) + var requests sync.WaitGroup + requests.Add(requestCount) + for range requestCount { + go func() { + defer requests.Done() + <-start + _, err := verificationStore.VerifyLocalAPIKey(ctx, created.Secret) + errorsByRequest <- err + }() + } + close(start) + requests.Wait() + close(errorsByRequest) + + for err := range errorsByRequest { + if err != nil { + t.Fatalf("concurrent API key verification failed: %v", err) + } + } + if acquired := verificationStore.pool.Stat().AcquiredConns(); acquired != 0 { + t.Fatalf("API key verification left %d connections acquired", acquired) + } +} + +func newLocalAPIKeyVerificationFixture(t *testing.T, maxConnections int32) (*Store, *Store, CreatedAPIKey, GatewayUser) { + t.Helper() + db := newIdentityPairingPostgresTestStore(t) + ctx := context.Background() + user, err := db.RegisterLocalUser(ctx, LocalRegisterInput{ + Username: "api-key-verification-user", + Password: "api-key-verification-password", + }) + if err != nil { + t.Fatalf("register API key verification user: %v", err) + } + created, err := db.CreateAPIKey(ctx, CreateAPIKeyInput{Name: "API key verification fixture"}, &auth.User{ + ID: user.ID, + GatewayUserID: user.ID, + GatewayTenantID: user.GatewayTenantID, + TenantID: user.TenantID, + TenantKey: user.TenantKey, + }) + if err != nil { + t.Fatalf("create API key verification fixture: %v", err) + } + + config := db.pool.Config() + config.MaxConns = maxConnections + config.MinConns = 0 + pool, err := pgxpool.NewWithConfig(ctx, config) + if err != nil { + t.Fatalf("create verification pool with %d connections: %v", maxConnections, err) + } + t.Cleanup(pool.Close) + return db, &Store{pool: pool}, created, user +} diff --git a/apps/api/internal/store/api_key_auth_test.go b/apps/api/internal/store/api_key_auth_test.go new file mode 100644 index 0000000..5b8e299 --- /dev/null +++ b/apps/api/internal/store/api_key_auth_test.go @@ -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 +} diff --git a/apps/api/internal/store/identity_pairing_integration_test.go b/apps/api/internal/store/identity_pairing_integration_test.go index 65c8b16..6837ee2 100644 --- a/apps/api/internal/store/identity_pairing_integration_test.go +++ b/apps/api/internal/store/identity_pairing_integration_test.go @@ -685,11 +685,13 @@ func newIdentityPairingPostgresTestStore(t *testing.T) *Store { migrationDirectory := filepath.Join(filepath.Dir(filename), "..", "..", "migrations") for _, migrationName := range []string{ "0001_init.sql", + "0017_task_record_enrichment.sql", "0061_oidc_server_sessions.sql", "0065_identity_configuration_revisions.sql", "0066_identity_onboarding_exchanges.sql", "0067_identity_secret_cleanup_queue.sql", "0068_identity_pairing_start_reservation.sql", + "0069_billing_correctness_v2.sql", } { migration, err := os.ReadFile(filepath.Join(migrationDirectory, migrationName)) if err != nil { diff --git a/apps/api/internal/store/postgres.go b/apps/api/internal/store/postgres.go index 8de239e..ee17719 100644 --- a/apps/api/internal/store/postgres.go +++ b/apps/api/internal/store/postgres.go @@ -7,6 +7,7 @@ import ( "encoding/base64" "encoding/json" "errors" + "net" "strings" "time" "unicode" @@ -22,6 +23,11 @@ type Store struct { pool *pgxpool.Pool } +const ( + postgresApplicationName = "easyai-ai-gateway" + postgresConnectTimeout = 5 * time.Second +) + func defaultAPIKeyScopes() []string { return []string{"chat", "embedding", "rerank", "image", "video", "music", "audio", "voice_clone"} } @@ -67,7 +73,11 @@ var ( ) func Connect(ctx context.Context, databaseURL string) (*Store, error) { - pool, err := pgxpool.New(ctx, databaseURL) + config, err := postgresPoolConfig(databaseURL) + if err != nil { + return nil, err + } + pool, err := pgxpool.NewWithConfig(ctx, config) if err != nil { return nil, err } @@ -78,6 +88,38 @@ func Connect(ctx context.Context, databaseURL string) (*Store, error) { return &Store{pool: pool}, nil } +func postgresPoolConfig(databaseURL string) (*pgxpool.Config, error) { + config, err := pgxpool.ParseConfig(databaseURL) + if err != nil { + return nil, err + } + config.ConnConfig.ConnectTimeout = postgresConnectTimeout + config.ConnConfig.RuntimeParams["application_name"] = postgresApplicationName + return config, nil +} + +func IsPostgresUnavailable(err error) bool { + if err == nil { + return false + } + if errors.Is(err, context.DeadlineExceeded) { + return true + } + var connectError *pgconn.ConnectError + if errors.As(err, &connectError) { + return true + } + var networkError net.Error + if errors.As(err, &networkError) { + return true + } + var postgresError *pgconn.PgError + if errors.As(err, &postgresError) { + return strings.HasPrefix(postgresError.Code, "08") || postgresError.Code == "53300" || strings.HasPrefix(postgresError.Code, "57P0") + } + return pgconn.SafeToRetry(err) +} + func (s *Store) Close() { s.pool.Close() } @@ -1519,11 +1561,35 @@ WHERE subject_type = 'api_key' AND subject_id = $1::uuid`, apiKeyID); err != nil } func (s *Store) VerifyLocalAPIKey(ctx context.Context, secret string) (*auth.User, error) { + return verifyLocalAPIKey(ctx, s.pool, secret) +} + +type localAPIKeyDatabase interface { + Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error) + Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error) +} + +type localAPIKeyCandidate struct { + apiKeyID string + hash string + keyPrefix string + keyName string + scopesBytes []byte + userGroupID string + gatewayUserID string + username string + rolesBytes []byte + gatewayTenantID string + tenantID string + tenantKey string +} + +func verifyLocalAPIKey(ctx context.Context, database localAPIKeyDatabase, secret string) (*auth.User, error) { prefix := apiKeyPrefix(secret) if prefix == "" { return nil, auth.ErrUnauthorized } - rows, err := s.pool.Query(ctx, ` + rows, err := database.Query(ctx, ` SELECT k.id::text, k.key_hash, k.key_prefix, k.name, k.scopes, COALESCE(k.user_group_id::text, u.default_user_group_id::text, ''), u.id::text, u.username, u.roles, COALESCE(u.gateway_tenant_id::text, ''), COALESCE(u.tenant_id, ''), COALESCE(u.tenant_key, '') @@ -1538,49 +1604,51 @@ WHERE k.key_prefix = $1 if err != nil { return nil, err } - defer rows.Close() + candidates, err := pgx.CollectRows(rows, func(row pgx.CollectableRow) (localAPIKeyCandidate, error) { + var candidate localAPIKeyCandidate + err := row.Scan( + &candidate.apiKeyID, + &candidate.hash, + &candidate.keyPrefix, + &candidate.keyName, + &candidate.scopesBytes, + &candidate.userGroupID, + &candidate.gatewayUserID, + &candidate.username, + &candidate.rolesBytes, + &candidate.gatewayTenantID, + &candidate.tenantID, + &candidate.tenantKey, + ) + return candidate, err + }) + if err != nil { + return nil, err + } - for rows.Next() { - var apiKeyID string - var hash string - var keyPrefix string - var keyName string - var scopesBytes []byte - var userGroupID string - var gatewayUserID string - var username string - var rolesBytes []byte - var gatewayTenantID string - var tenantID string - var tenantKey string - if err := rows.Scan(&apiKeyID, &hash, &keyPrefix, &keyName, &scopesBytes, &userGroupID, &gatewayUserID, &username, &rolesBytes, &gatewayTenantID, &tenantID, &tenantKey); err != nil { - return nil, err - } - if bcrypt.CompareHashAndPassword([]byte(hash), []byte(secret)) != nil { + for _, candidate := range candidates { + if bcrypt.CompareHashAndPassword([]byte(candidate.hash), []byte(secret)) != nil { continue } - if _, err := s.pool.Exec(ctx, `UPDATE gateway_api_keys SET last_used_at = now(), updated_at = now() WHERE id = $1::uuid`, apiKeyID); err != nil { + if _, err := database.Exec(ctx, `UPDATE gateway_api_keys SET last_used_at = now(), updated_at = now() WHERE id = $1::uuid`, candidate.apiKeyID); err != nil { return nil, err } return &auth.User{ - ID: gatewayUserID, - Username: username, - Roles: decodeStringArray(rolesBytes), - TenantID: tenantID, - GatewayTenantID: gatewayTenantID, - TenantKey: tenantKey, + ID: candidate.gatewayUserID, + Username: candidate.username, + Roles: decodeStringArray(candidate.rolesBytes), + TenantID: candidate.tenantID, + GatewayTenantID: candidate.gatewayTenantID, + TenantKey: candidate.tenantKey, Source: "gateway", - GatewayUserID: gatewayUserID, - UserGroupID: userGroupID, - APIKeyID: apiKeyID, - APIKeyName: keyName, - APIKeyPrefix: keyPrefix, - APIKeyScopes: decodeStringArray(scopesBytes), + GatewayUserID: candidate.gatewayUserID, + UserGroupID: candidate.userGroupID, + APIKeyID: candidate.apiKeyID, + APIKeyName: candidate.keyName, + APIKeyPrefix: candidate.keyPrefix, + APIKeyScopes: decodeStringArray(candidate.scopesBytes), }, nil } - if err := rows.Err(); err != nil { - return nil, err - } return nil, auth.ErrUnauthorized } diff --git a/apps/api/internal/store/postgres_config_test.go b/apps/api/internal/store/postgres_config_test.go new file mode 100644 index 0000000..622e055 --- /dev/null +++ b/apps/api/internal/store/postgres_config_test.go @@ -0,0 +1,49 @@ +package store + +import ( + "context" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgconn" +) + +func TestPostgresPoolConfigSetsDiagnosticAndConnectTimeout(t *testing.T) { + config, err := postgresPoolConfig("postgresql://gateway:password@localhost:5432/gateway?sslmode=disable") + if err != nil { + t.Fatalf("parse PostgreSQL pool config: %v", err) + } + if config.ConnConfig.ConnectTimeout != 5*time.Second { + t.Fatalf("connect timeout = %s, want 5s", config.ConnConfig.ConnectTimeout) + } + if applicationName := config.ConnConfig.RuntimeParams["application_name"]; applicationName != "easyai-ai-gateway" { + t.Fatalf("application_name = %q, want easyai-ai-gateway", applicationName) + } +} + +func TestPostgresPoolConfigRejectsMalformedURL(t *testing.T) { + if _, err := postgresPoolConfig("://malformed"); err == nil { + t.Fatal("expected malformed PostgreSQL URL to fail") + } +} + +func TestIsPostgresUnavailableClassifiesConnectivityFailures(t *testing.T) { + for _, testCase := range []struct { + name string + err error + }{ + {name: "deadline", err: context.DeadlineExceeded}, + {name: "connection exception", err: &pgconn.PgError{Code: "08006"}}, + {name: "too many connections", err: &pgconn.PgError{Code: "53300"}}, + {name: "cannot connect now", err: &pgconn.PgError{Code: "57P03"}}, + } { + t.Run(testCase.name, func(t *testing.T) { + if !IsPostgresUnavailable(testCase.err) { + t.Fatalf("error %v was not classified as PostgreSQL unavailable", testCase.err) + } + }) + } + if IsPostgresUnavailable(&pgconn.PgError{Code: "42601"}) { + t.Fatal("SQL syntax error was incorrectly classified as PostgreSQL unavailable") + } +} diff --git a/apps/web/src/api.test.ts b/apps/web/src/api.test.ts index bdf8210..470cc3b 100644 --- a/apps/web/src/api.test.ts +++ b/apps/web/src/api.test.ts @@ -9,12 +9,39 @@ import { getAPITask, getCurrentUser, getOpsManagementSkillMetadata, + loginLocalAccount, OIDC_BROWSER_SESSION_CREDENTIAL, startIdentityPairing, retireIdentityPairingSecurityEventConflict, validateIdentityRevision, } from './api'; +describe('local login transport', () => { + afterEach(() => { + vi.useRealTimers(); + vi.unstubAllGlobals(); + }); + + it('aborts after ten seconds and returns a stable login timeout message', async () => { + vi.useFakeTimers(); + const fetchMock = vi.fn((_url: string, init?: RequestInit) => new Promise((_resolve, reject) => { + init?.signal?.addEventListener('abort', () => reject(new DOMException('aborted', 'AbortError'))); + })); + vi.stubGlobal('fetch', fetchMock); + + const login = loginLocalAccount({ account: 'timeout-test-account', password: 'timeout-test-password' }); + const rejection = expect(login).rejects.toMatchObject({ + message: '登录请求超时,请稍后重试', + }); + await vi.advanceTimersByTimeAsync(10_000); + await rejection; + + const [, init] = fetchMock.mock.calls[0] as [string, RequestInit]; + expect(init.signal).toBeInstanceOf(AbortSignal); + expect(init.signal?.aborted).toBe(true); + }); +}); + describe('Gateway provisioning errors', () => { const cases = [ ['GATEWAY_USER_NOT_PROVISIONED', '该账号尚未开通 EasyAI Gateway'], diff --git a/apps/web/src/api.ts b/apps/web/src/api.ts index 834b372..aa79f36 100644 --- a/apps/web/src/api.ts +++ b/apps/web/src/api.ts @@ -121,6 +121,7 @@ export async function loginLocalAccount(input: { account: string; password: stri auth: false, body: input, method: 'POST', + timeoutMs: 10_000, }); } @@ -1206,7 +1207,7 @@ export async function deleteFileStorageChannel(token: string, channelId: string) async function request( path: string, - options: { token?: string; auth?: boolean; method?: string; body?: unknown; headers?: Record; signal?: AbortSignal } = {}, + options: { token?: string; auth?: boolean; method?: string; body?: unknown; headers?: Record; signal?: AbortSignal; timeoutMs?: number } = {}, ): Promise { const headers: Record = { ...(options.headers ?? {}) }; if (options.auth !== false && options.token && options.token !== OIDC_BROWSER_SESSION_CREDENTIAL) { @@ -1215,21 +1216,35 @@ async function request( if (options.body !== undefined) { headers['Content-Type'] = 'application/json'; } - const response = await fetch(`${API_BASE}${path}`, { - method: options.method ?? 'GET', - headers, - body: options.body === undefined ? undefined : JSON.stringify(options.body), - credentials: 'include', - signal: options.signal, - }); - if (!response.ok) { - const body = await response.text(); - throw new GatewayApiError(parseErrorDetails(body, response.status, `Request failed: ${response.status}`)); + const controller = options.timeoutMs ? new AbortController() : undefined; + const timeout = controller ? globalThis.setTimeout(() => controller.abort(), options.timeoutMs) : undefined; + const signal = controller && options.signal + ? AbortSignal.any([controller.signal, options.signal]) + : controller?.signal ?? options.signal; + try { + const response = await fetch(`${API_BASE}${path}`, { + method: options.method ?? 'GET', + headers, + body: options.body === undefined ? undefined : JSON.stringify(options.body), + credentials: 'include', + signal, + }); + if (!response.ok) { + const body = await response.text(); + throw new GatewayApiError(parseErrorDetails(body, response.status, `Request failed: ${response.status}`)); + } + if (response.status === 204) { + return undefined as T; + } + return response.json() as Promise; + } catch (error) { + if (controller?.signal.aborted) { + throw new GatewayApiError('登录请求超时,请稍后重试'); + } + throw error; + } finally { + if (timeout !== undefined) globalThis.clearTimeout(timeout); } - if (response.status === 204) { - return undefined as T; - } - return response.json() as Promise; } function authorizationHeader(token: string): Record { diff --git a/docker/nginx.conf b/docker/nginx.conf index 510bd73..a7ee56a 100644 --- a/docker/nginx.conf +++ b/docker/nginx.conf @@ -33,6 +33,20 @@ server { return 404; } + location = /gateway-api/api/v1/auth/login { + proxy_http_version 1.1; + proxy_set_header Host $host; + proxy_set_header X-Real-IP $remote_addr; + proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for; + proxy_set_header X-Forwarded-Proto $scheme; + proxy_set_header Connection ""; + proxy_connect_timeout 3s; + proxy_read_timeout 15s; + proxy_send_timeout 15s; + proxy_redirect off; + proxy_pass http://api:8088/api/v1/auth/login; + } + location /gateway-api/ { proxy_http_version 1.1; proxy_set_header Host $host;