diff --git a/cmd/server/main.go b/cmd/server/main.go index ee78690..8aa452a 100644 --- a/cmd/server/main.go +++ b/cmd/server/main.go @@ -10,6 +10,7 @@ import ( "fmt" "io" "log" + "log/slog" "net/http" "os" "os/signal" @@ -411,6 +412,60 @@ func main() { apiKeyService := aggregators.NewAPIKeyService(aggregatorRepo, oauthClient.ClientApp) log.Println("✅ API key service initialized") + // Start aggregator token refresh background job + // Timing rationale: + // - Runs every 30 minutes to catch tokens before they expire + // - 1-hour expiry buffer ensures we refresh well before expiration + // - This gives us 2 attempts (at 60min and 30min before expiry) to refresh + // - Note: APIKeyService.TokenRefreshBuffer (5min) is for on-demand refresh during API calls, + // while this background job provides proactive refresh for idle aggregators + tokenRefreshCtx, tokenRefreshCancel := context.WithCancel(context.Background()) + go func() { + defer func() { + if r := recover(); r != nil { + slog.Error("[TOKEN-REFRESH] CRITICAL: Background job panicked", + "panic", r, + ) + } + }() + + ticker := time.NewTicker(30 * time.Minute) + defer ticker.Stop() + + // Heartbeat counter for periodic health logging + cycleCount := 0 + + for { + select { + case <-tokenRefreshCtx.Done(): + slog.Info("[TOKEN-REFRESH] Aggregator token refresh job stopped") + return + case <-ticker.C: + cycleCount++ + refreshed, errs := apiKeyService.RefreshExpiringTokens(tokenRefreshCtx, 1*time.Hour) + if len(errs) > 0 { + slog.Warn("[TOKEN-REFRESH] Aggregator refresh completed with errors", + "refreshed", refreshed, + "failed", len(errs), + ) + for _, err := range errs { + slog.Error("[TOKEN-REFRESH] Refresh error", "error", err) + } + } else if refreshed > 0 { + slog.Info("[TOKEN-REFRESH] Aggregator refresh completed", + "refreshed", refreshed, + ) + } else if cycleCount%6 == 0 { + // Log heartbeat every 6 cycles (3 hours) when no work is done + slog.Info("[TOKEN-REFRESH] Heartbeat: background job running, no tokens needed refresh", + "cycles_completed", cycleCount, + ) + } + } + } + }() + log.Println("Started aggregator token refresh background job (runs every 30 minutes)") + // Get instance DID for service auth validator audience serviceDID := instanceDID // Use instance DID as the service audience @@ -770,8 +825,9 @@ func main() { ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second) defer cancel() - // Stop OAuth cleanup background job + // Stop background jobs cleanupCancel() + tokenRefreshCancel() if err := server.Shutdown(ctx); err != nil { log.Fatalf("Server shutdown error: %v", err) diff --git a/internal/api/middleware/apikey_adapter_test.go b/internal/api/middleware/apikey_adapter_test.go index 9819bff..fcf0d86 100644 --- a/internal/api/middleware/apikey_adapter_test.go +++ b/internal/api/middleware/apikey_adapter_test.go @@ -197,6 +197,10 @@ func (m *mockAPIKeyServiceRepository) GetCredentialsByAPIKeyHash(ctx context.Con return nil, aggregators.ErrAggregatorNotFound } +func (m *mockAPIKeyServiceRepository) ListAggregatorsNeedingTokenRefresh(ctx context.Context, expiryBuffer time.Duration) ([]*aggregators.AggregatorCredentials, error) { + return nil, nil +} + // ============================================================================= // ValidateKey Delegation Tests // ============================================================================= diff --git a/internal/core/aggregators/apikey_service.go b/internal/core/aggregators/apikey_service.go index f329586..d10cd85 100644 --- a/internal/core/aggregators/apikey_service.go +++ b/internal/core/aggregators/apikey_service.go @@ -371,3 +371,64 @@ func (s *APIKeyService) GetFailedLastUsedUpdates() int64 { func (s *APIKeyService) GetFailedNonceUpdates() int64 { return s.failedNonceUpdates.Load() } + +// perAggregatorRefreshTimeout is the maximum time allowed for refreshing +// a single aggregator's tokens. This prevents a slow OAuth server from +// blocking the entire refresh job. +const perAggregatorRefreshTimeout = 30 * time.Second + +// RefreshExpiringTokens proactively refreshes tokens for all aggregators +// whose tokens will expire within the given buffer period. +// Returns count of successful refreshes and any errors encountered. +// Each aggregator refresh has a 30-second timeout to prevent slow OAuth servers +// from blocking the entire job. +func (s *APIKeyService) RefreshExpiringTokens(ctx context.Context, expiryBuffer time.Duration) (refreshed int, errors []error) { + // Get all aggregators with tokens expiring within the buffer period + aggregators, err := s.repo.ListAggregatorsNeedingTokenRefresh(ctx, expiryBuffer) + if err != nil { + slog.Error("[TOKEN-REFRESH] Failed to list aggregators needing token refresh", + "error", err, + "expiry_buffer", expiryBuffer, + ) + return 0, []error{fmt.Errorf("failed to list aggregators needing refresh: %w", err)} + } + + if len(aggregators) == 0 { + return 0, nil + } + + slog.Info("[TOKEN-REFRESH] Starting proactive token refresh", + "aggregator_count", len(aggregators), + "expiry_buffer", expiryBuffer, + ) + + // Refresh tokens for each aggregator with per-aggregator timeout + for _, creds := range aggregators { + slog.Info("[TOKEN-REFRESH] Attempting token refresh for aggregator", + "did", creds.DID, + "token_expires_at", creds.OAuthTokenExpiresAt, + ) + + // Create per-aggregator timeout context to prevent slow OAuth servers + // from blocking the entire refresh cycle + refreshCtx, cancel := context.WithTimeout(ctx, perAggregatorRefreshTimeout) + err := s.RefreshTokensIfNeeded(refreshCtx, creds) + cancel() + + if err != nil { + slog.Error("[TOKEN-REFRESH] Failed to refresh tokens for aggregator", + "did", creds.DID, + "error", err, + ) + errors = append(errors, fmt.Errorf("aggregator %s: %w", creds.DID, err)) + } else { + slog.Info("[TOKEN-REFRESH] Successfully refreshed tokens for aggregator", + "did", creds.DID, + "new_expires_at", creds.OAuthTokenExpiresAt, + ) + refreshed++ + } + } + + return refreshed, errors +} diff --git a/internal/core/aggregators/apikey_service_test.go b/internal/core/aggregators/apikey_service_test.go index 9408d6c..d4d9284 100644 --- a/internal/core/aggregators/apikey_service_test.go +++ b/internal/core/aggregators/apikey_service_test.go @@ -34,15 +34,16 @@ func newTestAPIKeyService(repo Repository) *APIKeyService { // mockRepository implements Repository interface for testing type mockRepository struct { - getAggregatorFunc func(ctx context.Context, did string) (*Aggregator, error) - getByAPIKeyHashFunc func(ctx context.Context, keyHash string) (*Aggregator, error) - getCredentialsByAPIKeyHashFunc func(ctx context.Context, keyHash string) (*AggregatorCredentials, error) - getAggregatorCredentialsFunc func(ctx context.Context, did string) (*AggregatorCredentials, error) - setAPIKeyFunc func(ctx context.Context, did, keyPrefix, keyHash string, oauthCreds *OAuthCredentials) error - updateOAuthTokensFunc func(ctx context.Context, did, accessToken, refreshToken string, expiresAt time.Time) error - updateOAuthNoncesFunc func(ctx context.Context, did, authServerNonce, pdsNonce string) error - updateAPIKeyLastUsedFunc func(ctx context.Context, did string) error - revokeAPIKeyFunc func(ctx context.Context, did string) error + getAggregatorFunc func(ctx context.Context, did string) (*Aggregator, error) + getByAPIKeyHashFunc func(ctx context.Context, keyHash string) (*Aggregator, error) + getCredentialsByAPIKeyHashFunc func(ctx context.Context, keyHash string) (*AggregatorCredentials, error) + getAggregatorCredentialsFunc func(ctx context.Context, did string) (*AggregatorCredentials, error) + setAPIKeyFunc func(ctx context.Context, did, keyPrefix, keyHash string, oauthCreds *OAuthCredentials) error + updateOAuthTokensFunc func(ctx context.Context, did, accessToken, refreshToken string, expiresAt time.Time) error + updateOAuthNoncesFunc func(ctx context.Context, did, authServerNonce, pdsNonce string) error + updateAPIKeyLastUsedFunc func(ctx context.Context, did string) error + revokeAPIKeyFunc func(ctx context.Context, did string) error + listAggregatorsNeedingTokenRefreshFunc func(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) } func (m *mockRepository) GetAggregator(ctx context.Context, did string) (*Aggregator, error) { @@ -181,6 +182,13 @@ func (m *mockRepository) GetCredentialsByAPIKeyHash(ctx context.Context, keyHash return nil, ErrAggregatorNotFound } +func (m *mockRepository) ListAggregatorsNeedingTokenRefresh(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) { + if m.listAggregatorsNeedingTokenRefreshFunc != nil { + return m.listAggregatorsNeedingTokenRefreshFunc(ctx, expiryBuffer) + } + return nil, nil +} + func TestHashAPIKey(t *testing.T) { plainKey := "ckapi_abcdef1234567890abcdef1234567890" @@ -1141,3 +1149,214 @@ func TestAPIKeyService_FailedLastUsedUpdates_IncrementsOnError(t *testing.T) { t.Errorf("GetFailedLastUsedUpdates() after failure = %d, want 1", got) } } + +// ============================================================================= +// RefreshExpiringTokens Tests +// ============================================================================= + +func TestAPIKeyService_RefreshExpiringTokens_DatabaseError(t *testing.T) { + expectedError := errors.New("database connection failed") + repo := &mockRepository{ + listAggregatorsNeedingTokenRefreshFunc: func(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) { + return nil, expectedError + }, + } + service := newTestAPIKeyService(repo) + + refreshed, errs := service.RefreshExpiringTokens(context.Background(), 1*time.Hour) + + if refreshed != 0 { + t.Errorf("RefreshExpiringTokens() refreshed = %d, want 0", refreshed) + } + if len(errs) != 1 { + t.Fatalf("RefreshExpiringTokens() errors count = %d, want 1", len(errs)) + } + if !errors.Is(errs[0], expectedError) { + t.Errorf("RefreshExpiringTokens() error = %v, want %v", errs[0], expectedError) + } +} + +func TestAPIKeyService_RefreshExpiringTokens_EmptyList(t *testing.T) { + repo := &mockRepository{ + listAggregatorsNeedingTokenRefreshFunc: func(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) { + return []*AggregatorCredentials{}, nil + }, + } + service := newTestAPIKeyService(repo) + + refreshed, errs := service.RefreshExpiringTokens(context.Background(), 1*time.Hour) + + if refreshed != 0 { + t.Errorf("RefreshExpiringTokens() refreshed = %d, want 0", refreshed) + } + if len(errs) != 0 { + t.Errorf("RefreshExpiringTokens() errors count = %d, want 0", len(errs)) + } +} + +func TestAPIKeyService_RefreshExpiringTokens_NilList(t *testing.T) { + repo := &mockRepository{ + listAggregatorsNeedingTokenRefreshFunc: func(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) { + return nil, nil + }, + } + service := newTestAPIKeyService(repo) + + refreshed, errs := service.RefreshExpiringTokens(context.Background(), 1*time.Hour) + + if refreshed != 0 { + t.Errorf("RefreshExpiringTokens() refreshed = %d, want 0", refreshed) + } + if len(errs) != 0 { + t.Errorf("RefreshExpiringTokens() errors count = %d, want 0", len(errs)) + } +} + +func TestAPIKeyService_RefreshExpiringTokens_PassesCorrectExpiryBuffer(t *testing.T) { + expectedBuffer := 2 * time.Hour + var capturedBuffer time.Duration + + repo := &mockRepository{ + listAggregatorsNeedingTokenRefreshFunc: func(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) { + capturedBuffer = expiryBuffer + return nil, nil + }, + } + service := newTestAPIKeyService(repo) + + service.RefreshExpiringTokens(context.Background(), expectedBuffer) + + if capturedBuffer != expectedBuffer { + t.Errorf("RefreshExpiringTokens() passed expiryBuffer = %v, want %v", capturedBuffer, expectedBuffer) + } +} + +func TestAPIKeyService_RefreshExpiringTokens_TokensStillValid(t *testing.T) { + // When tokens are still valid (not within refresh buffer), no refresh should happen + // This tests the case where RefreshTokensIfNeeded returns early because tokens are valid + expiresAt := time.Now().Add(1 * time.Hour) // Well beyond the 5 minute buffer + + repo := &mockRepository{ + listAggregatorsNeedingTokenRefreshFunc: func(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) { + return []*AggregatorCredentials{ + { + DID: "did:plc:aggregator1", + OAuthTokenExpiresAt: &expiresAt, + OAuthAccessToken: "valid_token", + OAuthRefreshToken: "refresh_token", + }, + }, nil + }, + } + service := newTestAPIKeyService(repo) + + refreshed, errs := service.RefreshExpiringTokens(context.Background(), 1*time.Hour) + + // Tokens are valid, so RefreshTokensIfNeeded returns early without error + // and counts as "refreshed" (even though no actual refresh was needed) + if refreshed != 1 { + t.Errorf("RefreshExpiringTokens() refreshed = %d, want 1", refreshed) + } + if len(errs) != 0 { + t.Errorf("RefreshExpiringTokens() errors count = %d, want 0", len(errs)) + } +} + +func TestAPIKeyService_RefreshExpiringTokens_TokensExpired_RefreshFails(t *testing.T) { + // When tokens are expired and refresh fails (no OAuth app configured) + expiresAt := time.Now().Add(-1 * time.Hour) // Already expired + + repo := &mockRepository{ + listAggregatorsNeedingTokenRefreshFunc: func(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) { + return []*AggregatorCredentials{ + { + DID: "did:plc:aggregator1", + OAuthTokenExpiresAt: &expiresAt, + OAuthAccessToken: "expired_token", + OAuthRefreshToken: "refresh_token", + }, + }, nil + }, + } + service := newTestAPIKeyService(repo) + + refreshed, errs := service.RefreshExpiringTokens(context.Background(), 1*time.Hour) + + if refreshed != 0 { + t.Errorf("RefreshExpiringTokens() refreshed = %d, want 0", refreshed) + } + if len(errs) != 1 { + t.Errorf("RefreshExpiringTokens() errors count = %d, want 1", len(errs)) + } +} + +func TestAPIKeyService_RefreshExpiringTokens_MixedResults(t *testing.T) { + // Multiple aggregators: some with valid tokens, some with expired tokens + validExpiry := time.Now().Add(1 * time.Hour) + expiredExpiry := time.Now().Add(-1 * time.Hour) + + repo := &mockRepository{ + listAggregatorsNeedingTokenRefreshFunc: func(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) { + return []*AggregatorCredentials{ + { + DID: "did:plc:valid1", + OAuthTokenExpiresAt: &validExpiry, + OAuthAccessToken: "valid_token", + }, + { + DID: "did:plc:expired1", + OAuthTokenExpiresAt: &expiredExpiry, + OAuthAccessToken: "expired_token", + OAuthRefreshToken: "refresh_token", + }, + { + DID: "did:plc:valid2", + OAuthTokenExpiresAt: &validExpiry, + OAuthAccessToken: "valid_token2", + }, + }, nil + }, + } + service := newTestAPIKeyService(repo) + + refreshed, errs := service.RefreshExpiringTokens(context.Background(), 1*time.Hour) + + // 2 valid tokens should count as refreshed, 1 expired token should fail + if refreshed != 2 { + t.Errorf("RefreshExpiringTokens() refreshed = %d, want 2", refreshed) + } + if len(errs) != 1 { + t.Errorf("RefreshExpiringTokens() errors count = %d, want 1", len(errs)) + } +} + +func TestAPIKeyService_RefreshExpiringTokens_ContextCancellation(t *testing.T) { + // Test that context cancellation is respected + ctx, cancel := context.WithCancel(context.Background()) + cancel() // Cancel immediately + + repo := &mockRepository{ + listAggregatorsNeedingTokenRefreshFunc: func(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) { + // Check if context is already cancelled + select { + case <-ctx.Done(): + return nil, ctx.Err() + default: + return []*AggregatorCredentials{ + {DID: "did:plc:test"}, + }, nil + } + }, + } + service := newTestAPIKeyService(repo) + + refreshed, errs := service.RefreshExpiringTokens(ctx, 1*time.Hour) + + // Should fail due to context cancellation + if refreshed != 0 { + t.Errorf("RefreshExpiringTokens() refreshed = %d, want 0", refreshed) + } + if len(errs) != 1 { + t.Errorf("RefreshExpiringTokens() errors count = %d, want 1", len(errs)) + } +} diff --git a/internal/core/aggregators/interfaces.go b/internal/core/aggregators/interfaces.go index a32c387..60d7f3b 100644 --- a/internal/core/aggregators/interfaces.go +++ b/internal/core/aggregators/interfaces.go @@ -57,6 +57,10 @@ type Repository interface { UpdateAPIKeyLastUsed(ctx context.Context, did string) error // RevokeAPIKey marks an API key as revoked (sets api_key_revoked_at) RevokeAPIKey(ctx context.Context, did string) error + + // ListAggregatorsNeedingTokenRefresh returns aggregators with active API keys + // whose OAuth tokens expire within the given buffer period + ListAggregatorsNeedingTokenRefresh(ctx context.Context, expiryBuffer time.Duration) ([]*AggregatorCredentials, error) } // Service defines the interface for aggregator business logic diff --git a/internal/db/postgres/aggregator_repo.go b/internal/db/postgres/aggregator_repo.go index ba088ce..28908f2 100644 --- a/internal/db/postgres/aggregator_repo.go +++ b/internal/db/postgres/aggregator_repo.go @@ -1164,6 +1164,114 @@ func (r *postgresAggregatorRepo) GetCredentialsByAPIKeyHash(ctx context.Context, return creds, nil } +// ListAggregatorsNeedingTokenRefresh returns aggregators with active API keys +// whose OAuth tokens expire within the given buffer period. +// Used by background job to proactively refresh tokens before they expire. +func (r *postgresAggregatorRepo) ListAggregatorsNeedingTokenRefresh(ctx context.Context, expiryBuffer time.Duration) ([]*aggregators.AggregatorCredentials, error) { + query := ` + SELECT + did, + api_key_prefix, api_key_hash, api_key_created_at, api_key_revoked_at, api_key_last_used_at, + CASE + WHEN oauth_access_token_encrypted IS NOT NULL + THEN pgp_sym_decrypt(oauth_access_token_encrypted, (SELECT encode(key_data, 'hex') FROM encryption_keys WHERE id = 1)) + ELSE NULL + END as oauth_access_token, + CASE + WHEN oauth_refresh_token_encrypted IS NOT NULL + THEN pgp_sym_decrypt(oauth_refresh_token_encrypted, (SELECT encode(key_data, 'hex') FROM encryption_keys WHERE id = 1)) + ELSE NULL + END as oauth_refresh_token, + oauth_token_expires_at, + oauth_pds_url, oauth_auth_server_iss, oauth_auth_server_token_endpoint, + CASE + WHEN oauth_dpop_private_key_encrypted IS NOT NULL + THEN pgp_sym_decrypt(oauth_dpop_private_key_encrypted, (SELECT encode(key_data, 'hex') FROM encryption_keys WHERE id = 1)) + ELSE NULL + END as oauth_dpop_private_key_multibase, + oauth_dpop_authserver_nonce, oauth_dpop_pds_nonce + FROM aggregators + WHERE api_key_hash IS NOT NULL + AND api_key_revoked_at IS NULL + AND oauth_token_expires_at IS NOT NULL + AND oauth_token_expires_at <= NOW() + $1` + + rows, err := r.db.QueryContext(ctx, query, expiryBuffer) + if err != nil { + return nil, fmt.Errorf("failed to list aggregators needing token refresh: %w", err) + } + defer func() { _ = rows.Close() }() + + var results []*aggregators.AggregatorCredentials + for rows.Next() { + creds := &aggregators.AggregatorCredentials{} + var apiKeyPrefix, apiKeyHash sql.NullString + var oauthAccessToken, oauthRefreshToken sql.NullString + var oauthPDSURL, oauthAuthServerIss, oauthAuthServerTokenEndpoint sql.NullString + var oauthDPoPPrivateKey, oauthDPoPAuthServerNonce, oauthDPoPPDSNonce sql.NullString + var apiKeyCreatedAt, apiKeyRevokedAt, apiKeyLastUsed, oauthTokenExpiresAt sql.NullTime + + err := rows.Scan( + &creds.DID, + &apiKeyPrefix, + &apiKeyHash, + &apiKeyCreatedAt, + &apiKeyRevokedAt, + &apiKeyLastUsed, + &oauthAccessToken, + &oauthRefreshToken, + &oauthTokenExpiresAt, + &oauthPDSURL, + &oauthAuthServerIss, + &oauthAuthServerTokenEndpoint, + &oauthDPoPPrivateKey, + &oauthDPoPAuthServerNonce, + &oauthDPoPPDSNonce, + ) + if err != nil { + return nil, fmt.Errorf("failed to scan aggregator credentials: %w", err) + } + + // Map nullable string fields + creds.APIKeyPrefix = apiKeyPrefix.String + creds.APIKeyHash = apiKeyHash.String + creds.OAuthAccessToken = oauthAccessToken.String + creds.OAuthRefreshToken = oauthRefreshToken.String + creds.OAuthPDSURL = oauthPDSURL.String + creds.OAuthAuthServerIss = oauthAuthServerIss.String + creds.OAuthAuthServerTokenEndpoint = oauthAuthServerTokenEndpoint.String + creds.OAuthDPoPPrivateKeyMultibase = oauthDPoPPrivateKey.String + creds.OAuthDPoPAuthServerNonce = oauthDPoPAuthServerNonce.String + creds.OAuthDPoPPDSNonce = oauthDPoPPDSNonce.String + + // Map nullable time fields + if apiKeyCreatedAt.Valid { + t := apiKeyCreatedAt.Time + creds.APIKeyCreatedAt = &t + } + if apiKeyRevokedAt.Valid { + t := apiKeyRevokedAt.Time + creds.APIKeyRevokedAt = &t + } + if apiKeyLastUsed.Valid { + t := apiKeyLastUsed.Time + creds.APIKeyLastUsed = &t + } + if oauthTokenExpiresAt.Valid { + t := oauthTokenExpiresAt.Time + creds.OAuthTokenExpiresAt = &t + } + + results = append(results, creds) + } + + if err = rows.Err(); err != nil { + return nil, fmt.Errorf("error iterating aggregators needing token refresh: %w", err) + } + + return results, nil +} + // ===== Helper Functions ===== // scanAuthorizations is a helper to scan multiple authorization rows