Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
33 changes: 33 additions & 0 deletions payment/audit.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,33 @@
// SPDX-License-Identifier: Apache-2.0
// Copyright (c) 2026 Lanka Software Foundation

package payment

import "github.com/OpenNSW/core/shared/audit"

var _ audit.Details = AuditDetails{}

// AuditDetails is the payment-owned payload on an audit.Event.
type AuditDetails struct {
GatewayID string
Reference string
Status string // domain PaymentStatus, not audit.Status
Error string
}

func (d AuditDetails) Metadata() map[string]any {
m := make(map[string]any, 4)
if d.GatewayID != "" {
m["gateway_id"] = d.GatewayID
}
if d.Reference != "" {
m["reference"] = d.Reference
}
if d.Status != "" {
m["status"] = d.Status
}
if d.Error != "" {
m["error"] = d.Error
}
return m
}
3 changes: 3 additions & 0 deletions payment/go.mod
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@ go 1.26

require (
github.com/DATA-DOG/go-sqlmock v1.5.2
github.com/OpenNSW/core/shared v0.3.0
github.com/google/uuid v1.6.0
github.com/shopspring/decimal v1.4.0
github.com/stretchr/testify v1.12.1
Expand All @@ -23,3 +24,5 @@ require (
golang.org/x/sync v0.17.0 // indirect
golang.org/x/text v0.29.0 // indirect
)

replace github.com/OpenNSW/core/shared => ../shared
Comment thread
ginaxu1 marked this conversation as resolved.
2 changes: 2 additions & 0 deletions payment/handler_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ import (
"net/http/httptest"
"testing"

"github.com/OpenNSW/core/shared/audit"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/mock"
"github.com/stretchr/testify/require"
Expand All @@ -37,6 +38,7 @@ func (m *mockService) ProcessWebhook(context.Context, string, []byte, map[string
return m.webhookResp, m.webhookErr
}
func (m *mockService) SetTaskCompleter(TaskCompleter) {}
func (m *mockService) WithAuditor(audit.Auditor) {}

// serve routes a webhook POST through a mux so PathValue("gatewayId") resolves.
func serveWebhook(svc PaymentService) *httptest.ResponseRecorder {
Expand Down
168 changes: 159 additions & 9 deletions payment/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@ import (
"strings"
"time"

"github.com/OpenNSW/core/shared/audit"
"github.com/google/uuid"
)

Expand Down Expand Up @@ -83,12 +84,16 @@ type PaymentService interface {
// SetTaskCompleter injects the dependency used to advance the workflow when
// a payment settles. Wired post-construction to avoid an import cycle with taskv2.
SetTaskCompleter(completer TaskCompleter)

// WithAuditor injects an optional auditor for recording payment audit events.
WithAuditor(auditor audit.Auditor)
}

type paymentService struct {
repo PaymentRepository
registry GatewayRegistry
taskCompleter TaskCompleter
auditor audit.Auditor
}

// NewPaymentService initializes a new payment service.
Expand All @@ -103,6 +108,30 @@ func (s *paymentService) SetTaskCompleter(completer TaskCompleter) {
s.taskCompleter = completer
}

// WithAuditor injects an optional auditor for recording payment audit events.
// Passing nil disables auditing.
func (s *paymentService) WithAuditor(auditor audit.Auditor) {
s.auditor = auditor
}

const eventTypePayment = "PAYMENT"

// auditPayment is a nil-safe helper that emits a payment audit event.
func (s *paymentService) auditPayment(ctx context.Context, action audit.Action, status audit.Status, d AuditDetails) {
if s.auditor == nil {
return
}
s.auditor.Audit(ctx, audit.Event{
Timestamp: time.Now().UTC(),
EventType: eventTypePayment,
Action: action,
Status: status,
TargetType: "RESOURCE",
TargetID: d.Reference,
Details: d,
})
}

func (s *paymentService) ListAvailableMethods(ctx context.Context) ([]GatewayInfo, error) {
return s.registry.ListInfo(), nil
}
Expand Down Expand Up @@ -196,7 +225,14 @@ func (s *paymentService) CreateCheckoutSession(ctx context.Context, req CreateCh
slog.ErrorContext(ctx, "payment: failed to mark transaction failed after gateway error",
"reference", tx.ReferenceNumber, "error", uerr)
}
return nil, fmt.Errorf("gateway failed to create session: %w", err)
err = fmt.Errorf("gateway failed to create session: %w", err)
s.auditPayment(ctx, audit.ActionCreate, audit.StatusFailure, AuditDetails{
GatewayID: req.GatewayID,
Reference: tx.ReferenceNumber,
Status: string(PaymentStatusFailed),
Error: err.Error(),
})
return nil, err
}

// 4. Persist the gateway-assigned session id.
Expand All @@ -211,6 +247,11 @@ func (s *paymentService) CreateCheckoutSession(ctx context.Context, req CreateCh
return nil, errors.New("gateway returned nil session response")
}

s.auditPayment(ctx, audit.ActionCreate, audit.StatusSuccess, AuditDetails{
GatewayID: req.GatewayID,
Reference: generatedRef,
Status: string(PaymentStatusPending),
})
return &CreateCheckoutResponse{
ReferenceNumber: generatedRef,
SessionID: sessionResp.SessionID,
Expand All @@ -227,26 +268,46 @@ func (s *paymentService) ValidateReference(ctx context.Context, gatewayID string
// 1. Get the gateway from the registry using the ID from the URL
gateway, err := s.registry.Get(gatewayID)
if err != nil {
return nil, fmt.Errorf("gateway %s not found: %w", gatewayID, err)
err = fmt.Errorf("gateway %s not found: %w", gatewayID, err)
s.auditPayment(ctx, audit.ActionRead, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Error: err.Error(),
})
return nil, err
}

// 2. Verify the caller before any gateway-specific parsing runs. No
// reference lookup, and no presentment info, may be disclosed to an
// unverified caller.
if err := verifyCaller(ctx, gateway, gatewayID, rawBody, headers); err != nil {
s.auditPayment(ctx, audit.ActionRead, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Error: err.Error(),
})
return nil, err
}

// 3. Extract reference number from raw body
refNo, err := gateway.ExtractReferenceNumber(ctx, rawBody)
if err != nil {
return nil, fmt.Errorf("failed to extract reference number: %w", err)
err = fmt.Errorf("failed to extract reference number: %w", err)
s.auditPayment(ctx, audit.ActionRead, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Error: err.Error(),
})
return nil, err
}

// 4. Look up the transaction metadata from the DB
tx, err := s.repo.GetByReferenceNumber(ctx, refNo)
if err != nil {
return nil, fmt.Errorf("failed to retrieve payment reference: %w", err)
err = fmt.Errorf("failed to retrieve payment reference: %w", err)
s.auditPayment(ctx, audit.ActionRead, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Reference: refNo,
Error: err.Error(),
})
return nil, err
}

// 5. Map the internal record to the gateway DTO and decide payability.
Expand All @@ -272,36 +333,92 @@ func (s *paymentService) ValidateReference(ctx context.Context, gatewayID string
}
}

// Domain status is known once the transaction has been mapped, including on
// subsequent gateway formatting failures. Leave it empty only when there is
// no usable domain transaction (unknown reference or gateway mismatch).
status := ""
if validationTx != nil {
status = validationTx.Status
}

// 5. Delegate the protocol-specific response formatting to the gateway.
return gateway.HandleValidateReference(ctx, validationTx, isPayable, rawBody)
resp, err := gateway.HandleValidateReference(ctx, validationTx, isPayable, rawBody)
if err != nil {
s.auditPayment(ctx, audit.ActionRead, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Reference: refNo,
Status: status,
Error: err.Error(),
})
return nil, err
}
if resp == nil {
err = errors.New("gateway returned nil validation response")
s.auditPayment(ctx, audit.ActionRead, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Reference: refNo,
Status: status,
Error: err.Error(),
})
return nil, err
}
s.auditPayment(ctx, audit.ActionRead, audit.StatusSuccess, AuditDetails{
GatewayID: gatewayID, Reference: refNo, Status: status,
})
return resp, nil
}

func (s *paymentService) ProcessWebhook(ctx context.Context, gatewayID string, body []byte, headers map[string][]string) (*WebhookResponse, error) {
gateway, err := s.registry.Get(gatewayID)
if err != nil {
return nil, fmt.Errorf("failed to get gateway %s: %w", gatewayID, err)
err = fmt.Errorf("failed to get gateway %s: %w", gatewayID, err)
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Error: err.Error(),
})
return nil, err
}

// Verify the caller before any gateway-specific parsing runs. No
// transaction may be settled on the strength of an unverified caller.
if err := verifyCaller(ctx, gateway, gatewayID, body, headers); err != nil {
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Error: err.Error(),
})
return nil, err
}

gwPayload, webhookResp, err := gateway.ParseWebhook(ctx, body, headers)
if err != nil {
return nil, fmt.Errorf("gateway failed to parse webhook: %w", err)
err = fmt.Errorf("gateway failed to parse webhook: %w", err)
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Error: err.Error(),
})
return nil, err
}

if gwPayload == nil {
return nil, fmt.Errorf("gateway returned nil webhook payload")
err = fmt.Errorf("gateway returned nil webhook payload")
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Error: err.Error(),
})
return nil, err
}

// Translate the canonical gateway status into our domain status, rejecting
// anything unrecognized (defense-in-depth against a misbehaving gateway).
newStatus, err := toDomainStatus(gwPayload.Status)
if err != nil {
return nil, fmt.Errorf("webhook for %s: %w", gwPayload.ReferenceNumber, err)
err = fmt.Errorf("webhook for %s: %w", gwPayload.ReferenceNumber, err)
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Reference: gwPayload.ReferenceNumber,
Error: err.Error(),
})
return nil, err
}

// Claim and apply the status transition atomically. A row-level lock makes
Expand Down Expand Up @@ -329,6 +446,7 @@ func (s *paymentService) ProcessWebhook(ctx context.Context, gatewayID string, b
// concurrent delivery that committed first) — nothing more to do.
if tx.Status == PaymentStatusSuccess || tx.Status == PaymentStatusFailed {
slog.InfoContext(ctx, "webhook ignored (idempotent)", "reference", tx.ReferenceNumber, "current_status", tx.Status)
finalStatus = tx.Status
return nil
}

Expand Down Expand Up @@ -363,6 +481,11 @@ func (s *paymentService) ProcessWebhook(ctx context.Context, gatewayID string, b
return nil
})
if err != nil {
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Reference: gwPayload.ReferenceNumber,
Error: err.Error(),
})
return nil, err
}

Expand All @@ -371,6 +494,11 @@ func (s *paymentService) ProcessWebhook(ctx context.Context, gatewayID string, b

// Already terminal / nothing claimed — don't advance again.
if !advance {
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusSuccess, AuditDetails{
GatewayID: gatewayID,
Reference: gwPayload.ReferenceNumber,
Status: string(finalStatus),
})
return webhookResp, nil
}

Expand All @@ -381,6 +509,11 @@ func (s *paymentService) ProcessWebhook(ctx context.Context, gatewayID string, b
// task signal; any other status leaves the task untouched so a non-terminal or
// unrecognized gateway status can't be misread as paid.
if s.taskCompleter == nil {
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusSuccess, AuditDetails{
GatewayID: gatewayID,
Reference: gwPayload.ReferenceNumber,
Status: string(finalStatus),
})
return webhookResp, nil
}

Expand All @@ -394,6 +527,11 @@ func (s *paymentService) ProcessWebhook(ctx context.Context, gatewayID string, b
if statusStr == "" {
slog.WarnContext(ctx, "payment: non-terminal webhook status, not advancing task",
"reference", gwPayload.ReferenceNumber, "task_id", advanceTask, "status", finalStatus)
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusSuccess, AuditDetails{
GatewayID: gatewayID,
Reference: gwPayload.ReferenceNumber,
Status: string(finalStatus),
})
return webhookResp, nil
}

Expand All @@ -410,8 +548,20 @@ func (s *paymentService) ProcessWebhook(ctx context.Context, gatewayID string, b
// The transaction is already persisted; log and let the gateway retry
// drive a re-attempt rather than masking the failure as success.
slog.ErrorContext(ctx, "payment: failed to advance task step", "task_id", advanceTask, "error", err)
s.auditPayment(ctx, audit.ActionUpdate, audit.StatusFailure, AuditDetails{
GatewayID: gatewayID,
Reference: gwPayload.ReferenceNumber,
Status: string(finalStatus),
Error: err.Error(),
})
return nil, fmt.Errorf("failed to advance task step for %s: %w", advanceTask, err)
}

s.auditPayment(ctx, audit.ActionUpdate, audit.StatusSuccess, AuditDetails{
GatewayID: gatewayID,
Reference: gwPayload.ReferenceNumber,
Status: string(finalStatus),
})

return webhookResp, nil
}
Loading
Loading