fix(purchase): correct gift amount deduction logic and enhance payment processing comments
This commit is contained in:
+298
-51
@@ -1,90 +1,302 @@
|
||||
// Package deduction provides functionality for calculating remaining amounts
|
||||
// in subscription billing systems, supporting various time units and traffic-based calculations.
|
||||
package deduction
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"math"
|
||||
"time"
|
||||
|
||||
"github.com/perfect-panel/server/pkg/tool"
|
||||
)
|
||||
|
||||
const (
|
||||
UnitTimeNoLimit = "NoLimit"
|
||||
UnitTimeYear = "Year"
|
||||
UnitTimeMonth = "Month"
|
||||
UnitTimeDay = "Day"
|
||||
UintTimeHour = "Hour"
|
||||
UintTimeMinute = "Minute"
|
||||
// Time unit constants for subscription billing
|
||||
UnitTimeNoLimit = "NoLimit" // Unlimited time subscription
|
||||
UnitTimeYear = "Year" // Annual subscription
|
||||
UnitTimeMonth = "Month" // Monthly subscription
|
||||
UnitTimeDay = "Day" // Daily subscription
|
||||
UnitTimeHour = "Hour" // Hourly subscription
|
||||
UnitTimeMinute = "Minute" // Per-minute subscription
|
||||
|
||||
ResetCycleNone = 0
|
||||
ResetCycle1st = 1
|
||||
ResetCycleMonthly = 2
|
||||
ResetCycleYear = 3
|
||||
// Reset cycle constants for traffic resets
|
||||
ResetCycleNone = 0 // No reset cycle
|
||||
ResetCycle1st = 1 // Reset on 1st of each month
|
||||
ResetCycleMonthly = 2 // Reset monthly based on start date
|
||||
ResetCycleYear = 3 // Reset yearly based on start date
|
||||
|
||||
// Safety limits for overflow protection
|
||||
maxInt64 = math.MaxInt64
|
||||
minInt64 = math.MinInt64
|
||||
)
|
||||
|
||||
// Error definitions for validation and calculation failures
|
||||
var (
|
||||
ErrInvalidQuantity = errors.New("order quantity cannot be zero or negative")
|
||||
ErrInvalidAmount = errors.New("order amount cannot be negative")
|
||||
ErrInvalidTraffic = errors.New("traffic values cannot be negative")
|
||||
ErrInvalidTimeRange = errors.New("expire time must be after start time")
|
||||
ErrInvalidUnitTime = errors.New("invalid unit time")
|
||||
ErrInvalidDeductionRatio = errors.New("deduction ratio must be between 0 and 100")
|
||||
ErrOverflow = errors.New("calculation overflow")
|
||||
)
|
||||
|
||||
// Subscribe represents a subscription with time and traffic limits
|
||||
type Subscribe struct {
|
||||
StartTime time.Time
|
||||
ExpireTime time.Time
|
||||
Traffic int64
|
||||
Download int64
|
||||
Upload int64
|
||||
UnitTime string
|
||||
UnitPrice int64
|
||||
ResetCycle int64
|
||||
DeductionRatio int64
|
||||
StartTime time.Time // Subscription start time
|
||||
ExpireTime time.Time // Subscription expiration time
|
||||
Traffic int64 // Total traffic allowance in bytes
|
||||
Download int64 // Downloaded traffic in bytes
|
||||
Upload int64 // Uploaded traffic in bytes
|
||||
UnitTime string // Time unit for billing (Year, Month, Day, etc.)
|
||||
UnitPrice int64 // Price per unit time
|
||||
ResetCycle int64 // Traffic reset cycle
|
||||
DeductionRatio int64 // Deduction ratio for weighted calculations (0-100)
|
||||
}
|
||||
|
||||
// Order represents a purchase order for subscription calculation
|
||||
type Order struct {
|
||||
Amount int64
|
||||
Quantity int64
|
||||
Amount int64 // Total order amount
|
||||
Quantity int64 // Order quantity
|
||||
}
|
||||
|
||||
func CalculateRemainingAmount(sub Subscribe, order Order) int64 {
|
||||
if sub.UnitTime == UnitTimeNoLimit && sub.ResetCycle != 0 {
|
||||
return 0
|
||||
// Validate checks if the Subscribe struct contains valid data
|
||||
func (s *Subscribe) Validate() error {
|
||||
if s.Traffic < 0 || s.Download < 0 || s.Upload < 0 {
|
||||
return ErrInvalidTraffic
|
||||
}
|
||||
// 实际单价
|
||||
sub.UnitPrice = order.Amount / order.Quantity
|
||||
now := time.Now()
|
||||
|
||||
if s.Download+s.Upload > s.Traffic {
|
||||
return fmt.Errorf("download + upload (%d) cannot exceed total traffic (%d)", s.Download+s.Upload, s.Traffic)
|
||||
}
|
||||
|
||||
if !s.ExpireTime.After(s.StartTime) {
|
||||
return ErrInvalidTimeRange
|
||||
}
|
||||
|
||||
if s.DeductionRatio < 0 || s.DeductionRatio > 100 {
|
||||
return ErrInvalidDeductionRatio
|
||||
}
|
||||
|
||||
validUnitTimes := []string{UnitTimeNoLimit, UnitTimeYear, UnitTimeMonth, UnitTimeDay, UnitTimeHour, UnitTimeMinute}
|
||||
valid := false
|
||||
for _, ut := range validUnitTimes {
|
||||
if s.UnitTime == ut {
|
||||
valid = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !valid {
|
||||
return ErrInvalidUnitTime
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Validate checks if the Order struct contains valid data
|
||||
func (o *Order) Validate() error {
|
||||
if o.Quantity <= 0 {
|
||||
return ErrInvalidQuantity
|
||||
}
|
||||
if o.Amount < 0 {
|
||||
return ErrInvalidAmount
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// safeMultiply performs multiplication with overflow protection
|
||||
func safeMultiply(a, b int64) (int64, error) {
|
||||
if a == 0 || b == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
if a > 0 && b > 0 {
|
||||
if a > maxInt64/b {
|
||||
return 0, ErrOverflow
|
||||
}
|
||||
} else if a < 0 && b < 0 {
|
||||
if a < maxInt64/b {
|
||||
return 0, ErrOverflow
|
||||
}
|
||||
} else {
|
||||
if (a > 0 && b < minInt64/a) || (a < 0 && b > minInt64/a) {
|
||||
return 0, ErrOverflow
|
||||
}
|
||||
}
|
||||
|
||||
return a * b, nil
|
||||
}
|
||||
|
||||
// safeAdd performs addition with overflow protection
|
||||
func safeAdd(a, b int64) (int64, error) {
|
||||
if (b > 0 && a > maxInt64-b) || (b < 0 && a < minInt64-b) {
|
||||
return 0, ErrOverflow
|
||||
}
|
||||
return a + b, nil
|
||||
}
|
||||
|
||||
// safeDivide performs division with zero-division protection
|
||||
func safeDivide(a, b int64) (int64, error) {
|
||||
if b == 0 {
|
||||
return 0, errors.New("division by zero")
|
||||
}
|
||||
return a / b, nil
|
||||
}
|
||||
|
||||
// CalculateRemainingAmount calculates the remaining refund amount for a subscription
|
||||
// based on unused time and traffic. Returns the amount and any calculation errors.
|
||||
func CalculateRemainingAmount(sub Subscribe, order Order) (int64, error) {
|
||||
if err := sub.Validate(); err != nil {
|
||||
return 0, fmt.Errorf("invalid subscription: %w", err)
|
||||
}
|
||||
|
||||
if err := order.Validate(); err != nil {
|
||||
return 0, fmt.Errorf("invalid order: %w", err)
|
||||
}
|
||||
|
||||
if sub.UnitTime == UnitTimeNoLimit && sub.ResetCycle != 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
unitPrice, err := safeDivide(order.Amount, order.Quantity)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to calculate unit price: %w", err)
|
||||
}
|
||||
sub.UnitPrice = unitPrice
|
||||
|
||||
loc, err := time.LoadLocation(sub.StartTime.Location().String())
|
||||
if err != nil {
|
||||
loc = time.UTC
|
||||
}
|
||||
now := time.Now().In(loc)
|
||||
|
||||
switch sub.UnitTime {
|
||||
case UnitTimeNoLimit:
|
||||
usedTraffic := sub.Traffic - sub.Download - sub.Upload
|
||||
unitPrice := float64(order.Amount) / float64(sub.Traffic)
|
||||
return int64(float64(usedTraffic) * unitPrice)
|
||||
return calculateNoLimitAmount(sub, order)
|
||||
|
||||
case UnitTimeYear:
|
||||
remainingYears := tool.YearDiff(now, sub.ExpireTime)
|
||||
remainingUnitTimeAmount := calculateRemainingUnitTimeAmount(sub)
|
||||
return int64(remainingYears)*sub.UnitPrice + remainingUnitTimeAmount
|
||||
remainingUnitTimeAmount, err := calculateRemainingUnitTimeAmount(sub)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
yearAmount, err := safeMultiply(int64(remainingYears), sub.UnitPrice)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("year calculation overflow: %w", err)
|
||||
}
|
||||
|
||||
total, err := safeAdd(yearAmount, remainingUnitTimeAmount)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("total calculation overflow: %w", err)
|
||||
}
|
||||
|
||||
return total, nil
|
||||
|
||||
case UnitTimeMonth:
|
||||
remainingMonths := tool.MonthDiff(now, sub.ExpireTime)
|
||||
remainingUnitTimeAmount := calculateRemainingUnitTimeAmount(sub)
|
||||
return int64(remainingMonths)*sub.UnitPrice + remainingUnitTimeAmount
|
||||
remainingUnitTimeAmount, err := calculateRemainingUnitTimeAmount(sub)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
monthAmount, err := safeMultiply(int64(remainingMonths), sub.UnitPrice)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("month calculation overflow: %w", err)
|
||||
}
|
||||
|
||||
total, err := safeAdd(monthAmount, remainingUnitTimeAmount)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("total calculation overflow: %w", err)
|
||||
}
|
||||
|
||||
return total, nil
|
||||
|
||||
case UnitTimeDay:
|
||||
remainingDays := tool.DayDiff(now, sub.ExpireTime)
|
||||
remainingUnitTimeAmount := calculateRemainingUnitTimeAmount(sub)
|
||||
return remainingDays*sub.UnitPrice + remainingUnitTimeAmount
|
||||
remainingUnitTimeAmount, err := calculateRemainingUnitTimeAmount(sub)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
|
||||
dayAmount, err := safeMultiply(remainingDays, sub.UnitPrice)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("day calculation overflow: %w", err)
|
||||
}
|
||||
|
||||
total, err := safeAdd(dayAmount, remainingUnitTimeAmount)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("total calculation overflow: %w", err)
|
||||
}
|
||||
|
||||
return total, nil
|
||||
}
|
||||
|
||||
return 0
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
func calculateRemainingUnitTimeAmount(sub Subscribe) int64 {
|
||||
// calculateNoLimitAmount calculates refund amount for unlimited time subscriptions
|
||||
// based on unused traffic only
|
||||
func calculateNoLimitAmount(sub Subscribe, order Order) (int64, error) {
|
||||
if sub.Traffic == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
usedTraffic := sub.Traffic - sub.Download - sub.Upload
|
||||
if usedTraffic < 0 {
|
||||
usedTraffic = 0
|
||||
}
|
||||
|
||||
unitPrice := float64(order.Amount) / float64(sub.Traffic)
|
||||
result := float64(usedTraffic) * unitPrice
|
||||
|
||||
if result > float64(maxInt64) || result < float64(minInt64) {
|
||||
return 0, ErrOverflow
|
||||
}
|
||||
|
||||
return int64(result), nil
|
||||
}
|
||||
|
||||
// calculateRemainingUnitTimeAmount calculates the remaining amount based on
|
||||
// both time and traffic usage, applying deduction ratios when specified
|
||||
func calculateRemainingUnitTimeAmount(sub Subscribe) (int64, error) {
|
||||
now := time.Now()
|
||||
trafficWeight, timeWeight := calculateWeights(sub.DeductionRatio)
|
||||
remainingDays, totalDays := getRemainingAndTotalDays(sub, now)
|
||||
remainingTraffic := sub.Traffic - sub.Download - sub.Upload
|
||||
remainingTimeAmount := calculateProportionalAmount(sub.UnitPrice, remainingDays, totalDays)
|
||||
remainingTrafficAmount := calculateProportionalAmount(sub.UnitPrice, remainingTraffic, sub.Traffic)
|
||||
if sub.Traffic == 0 {
|
||||
return remainingTimeAmount
|
||||
|
||||
if totalDays == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
remainingTraffic := sub.Traffic - sub.Download - sub.Upload
|
||||
if remainingTraffic < 0 {
|
||||
remainingTraffic = 0
|
||||
}
|
||||
|
||||
remainingTimeAmount, err := calculateProportionalAmount(sub.UnitPrice, remainingDays, totalDays)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("time amount calculation failed: %w", err)
|
||||
}
|
||||
|
||||
if sub.Traffic == 0 {
|
||||
return remainingTimeAmount, nil
|
||||
}
|
||||
|
||||
remainingTrafficAmount, err := calculateProportionalAmount(sub.UnitPrice, remainingTraffic, sub.Traffic)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("traffic amount calculation failed: %w", err)
|
||||
}
|
||||
|
||||
if sub.DeductionRatio != 0 {
|
||||
return calculateWeightedAmount(sub.UnitPrice, remainingTraffic, sub.Traffic, remainingDays, totalDays, trafficWeight, timeWeight)
|
||||
}
|
||||
|
||||
return min(remainingTimeAmount, remainingTrafficAmount)
|
||||
return min(remainingTimeAmount, remainingTrafficAmount), nil
|
||||
}
|
||||
|
||||
// calculateWeights converts deduction ratio to traffic and time weights
|
||||
// for weighted calculations
|
||||
func calculateWeights(deductionRatio int64) (float64, float64) {
|
||||
if deductionRatio == 0 {
|
||||
return 0, 0
|
||||
@@ -94,20 +306,32 @@ func calculateWeights(deductionRatio int64) (float64, float64) {
|
||||
return trafficWeight, timeWeight
|
||||
}
|
||||
|
||||
// getRemainingAndTotalDays calculates remaining and total days based on
|
||||
// the subscription's reset cycle configuration
|
||||
func getRemainingAndTotalDays(sub Subscribe, now time.Time) (int64, int64) {
|
||||
switch sub.ResetCycle {
|
||||
case ResetCycleNone:
|
||||
|
||||
remaining := sub.ExpireTime.Sub(now).Hours() / 24
|
||||
total := sub.ExpireTime.Sub(sub.StartTime).Hours() / 24
|
||||
if remaining < 0 {
|
||||
remaining = 0
|
||||
}
|
||||
if total < 0 {
|
||||
total = 0
|
||||
}
|
||||
return int64(remaining), int64(total)
|
||||
|
||||
case ResetCycle1st:
|
||||
return tool.DaysToNextMonth(now), tool.GetLastDayOfMonth(now)
|
||||
|
||||
case ResetCycleMonthly:
|
||||
// -1 to include the current day
|
||||
return tool.DaysToMonthDay(now, sub.StartTime.Day()) - 1, tool.DaysToMonthDay(now, sub.StartTime.Day())
|
||||
remaining := tool.DaysToMonthDay(now, sub.StartTime.Day()) - 1
|
||||
total := tool.DaysToMonthDay(now, sub.StartTime.Day())
|
||||
if remaining < 0 {
|
||||
remaining = 0
|
||||
}
|
||||
return remaining, total
|
||||
|
||||
case ResetCycleYear:
|
||||
return tool.DaysToYearDay(now, int(sub.StartTime.Month()), sub.StartTime.Day()),
|
||||
tool.GetYearDays(now, int(sub.StartTime.Month()), sub.StartTime.Day())
|
||||
@@ -115,13 +339,36 @@ func getRemainingAndTotalDays(sub Subscribe, now time.Time) (int64, int64) {
|
||||
return 0, 0
|
||||
}
|
||||
|
||||
func calculateWeightedAmount(unitPrice, remainingTraffic, totalTraffic, remainingDays, totalDays int64, trafficWeight, timeWeight float64) int64 {
|
||||
// calculateWeightedAmount applies weighted calculation combining both time and traffic
|
||||
// remaining ratios based on the specified weights
|
||||
func calculateWeightedAmount(unitPrice, remainingTraffic, totalTraffic, remainingDays, totalDays int64, trafficWeight, timeWeight float64) (int64, error) {
|
||||
if totalDays == 0 || totalTraffic == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
remainingTimeRatio := float64(remainingDays) / float64(totalDays)
|
||||
remainingTrafficRatio := float64(remainingTraffic) / float64(totalTraffic)
|
||||
weightedRemainingRatio := (timeWeight * remainingTimeRatio) + (trafficWeight * remainingTrafficRatio)
|
||||
return int64(float64(unitPrice) * weightedRemainingRatio)
|
||||
|
||||
result := float64(unitPrice) * weightedRemainingRatio
|
||||
if result > float64(maxInt64) || result < float64(minInt64) {
|
||||
return 0, ErrOverflow
|
||||
}
|
||||
|
||||
return int64(result), nil
|
||||
}
|
||||
|
||||
func calculateProportionalAmount(unitPrice, remaining, total int64) int64 {
|
||||
return int64(float64(unitPrice) * (float64(remaining) / float64(total)))
|
||||
// calculateProportionalAmount calculates proportional amount based on
|
||||
// remaining vs total ratio with overflow protection
|
||||
func calculateProportionalAmount(unitPrice, remaining, total int64) (int64, error) {
|
||||
if total == 0 {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
result := float64(unitPrice) * (float64(remaining) / float64(total))
|
||||
if result > float64(maxInt64) || result < float64(minInt64) {
|
||||
return 0, ErrOverflow
|
||||
}
|
||||
|
||||
return int64(result), nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,665 @@
|
||||
package deduction
|
||||
|
||||
import (
|
||||
"math"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestSubscribe_Validate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
sub Subscribe
|
||||
wantErr bool
|
||||
errType error
|
||||
}{
|
||||
{
|
||||
name: "valid subscription",
|
||||
sub: Subscribe{
|
||||
StartTime: time.Now(),
|
||||
ExpireTime: time.Now().Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 100,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeMonth,
|
||||
DeductionRatio: 50,
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "negative traffic",
|
||||
sub: Subscribe{
|
||||
StartTime: time.Now(),
|
||||
ExpireTime: time.Now().Add(24 * time.Hour),
|
||||
Traffic: -1000,
|
||||
Download: 100,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeMonth,
|
||||
DeductionRatio: 50,
|
||||
},
|
||||
wantErr: true,
|
||||
errType: ErrInvalidTraffic,
|
||||
},
|
||||
{
|
||||
name: "negative download",
|
||||
sub: Subscribe{
|
||||
StartTime: time.Now(),
|
||||
ExpireTime: time.Now().Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: -100,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeMonth,
|
||||
DeductionRatio: 50,
|
||||
},
|
||||
wantErr: true,
|
||||
errType: ErrInvalidTraffic,
|
||||
},
|
||||
{
|
||||
name: "download + upload exceeds traffic",
|
||||
sub: Subscribe{
|
||||
StartTime: time.Now(),
|
||||
ExpireTime: time.Now().Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 600,
|
||||
Upload: 500,
|
||||
UnitTime: UnitTimeMonth,
|
||||
DeductionRatio: 50,
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "expire time before start time",
|
||||
sub: Subscribe{
|
||||
StartTime: time.Now(),
|
||||
ExpireTime: time.Now().Add(-24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 100,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeMonth,
|
||||
DeductionRatio: 50,
|
||||
},
|
||||
wantErr: true,
|
||||
errType: ErrInvalidTimeRange,
|
||||
},
|
||||
{
|
||||
name: "invalid deduction ratio - negative",
|
||||
sub: Subscribe{
|
||||
StartTime: time.Now(),
|
||||
ExpireTime: time.Now().Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 100,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeMonth,
|
||||
DeductionRatio: -10,
|
||||
},
|
||||
wantErr: true,
|
||||
errType: ErrInvalidDeductionRatio,
|
||||
},
|
||||
{
|
||||
name: "invalid deduction ratio - over 100",
|
||||
sub: Subscribe{
|
||||
StartTime: time.Now(),
|
||||
ExpireTime: time.Now().Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 100,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeMonth,
|
||||
DeductionRatio: 150,
|
||||
},
|
||||
wantErr: true,
|
||||
errType: ErrInvalidDeductionRatio,
|
||||
},
|
||||
{
|
||||
name: "invalid unit time",
|
||||
sub: Subscribe{
|
||||
StartTime: time.Now(),
|
||||
ExpireTime: time.Now().Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 100,
|
||||
Upload: 200,
|
||||
UnitTime: "InvalidUnit",
|
||||
DeductionRatio: 50,
|
||||
},
|
||||
wantErr: true,
|
||||
errType: ErrInvalidUnitTime,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := tt.sub.Validate()
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("Subscribe.Validate() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if tt.errType != nil && err != tt.errType {
|
||||
t.Errorf("Subscribe.Validate() error = %v, want %v", err, tt.errType)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestOrder_Validate(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
order Order
|
||||
wantErr bool
|
||||
errType error
|
||||
}{
|
||||
{
|
||||
name: "valid order",
|
||||
order: Order{Amount: 1000, Quantity: 2},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "zero quantity",
|
||||
order: Order{Amount: 1000, Quantity: 0},
|
||||
wantErr: true,
|
||||
errType: ErrInvalidQuantity,
|
||||
},
|
||||
{
|
||||
name: "negative quantity",
|
||||
order: Order{Amount: 1000, Quantity: -1},
|
||||
wantErr: true,
|
||||
errType: ErrInvalidQuantity,
|
||||
},
|
||||
{
|
||||
name: "negative amount",
|
||||
order: Order{Amount: -1000, Quantity: 2},
|
||||
wantErr: true,
|
||||
errType: ErrInvalidAmount,
|
||||
},
|
||||
{
|
||||
name: "zero amount is valid",
|
||||
order: Order{Amount: 0, Quantity: 1},
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
err := tt.order.Validate()
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("Order.Validate() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if tt.errType != nil && err != tt.errType {
|
||||
t.Errorf("Order.Validate() error = %v, want %v", err, tt.errType)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeMultiply(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
a, b int64
|
||||
want int64
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "normal multiplication",
|
||||
a: 10,
|
||||
b: 20,
|
||||
want: 200,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "zero multiplication",
|
||||
a: 10,
|
||||
b: 0,
|
||||
want: 0,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "negative multiplication",
|
||||
a: -10,
|
||||
b: 20,
|
||||
want: -200,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "overflow case",
|
||||
a: math.MaxInt64,
|
||||
b: 2,
|
||||
want: 0,
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "large numbers no overflow",
|
||||
a: 1000000,
|
||||
b: 1000000,
|
||||
want: 1000000000000,
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := safeMultiply(tt.a, tt.b)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("safeMultiply() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if got != tt.want {
|
||||
t.Errorf("safeMultiply() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeAdd(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
a, b int64
|
||||
want int64
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "normal addition",
|
||||
a: 10,
|
||||
b: 20,
|
||||
want: 30,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "negative addition",
|
||||
a: -10,
|
||||
b: 5,
|
||||
want: -5,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "overflow case",
|
||||
a: math.MaxInt64,
|
||||
b: 1,
|
||||
want: 0,
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "underflow case",
|
||||
a: math.MinInt64,
|
||||
b: -1,
|
||||
want: 0,
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := safeAdd(tt.a, tt.b)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("safeAdd() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if got != tt.want {
|
||||
t.Errorf("safeAdd() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestSafeDivide(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
a, b int64
|
||||
want int64
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "normal division",
|
||||
a: 20,
|
||||
b: 10,
|
||||
want: 2,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "division by zero",
|
||||
a: 20,
|
||||
b: 0,
|
||||
want: 0,
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "negative division",
|
||||
a: -20,
|
||||
b: 10,
|
||||
want: -2,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "zero dividend",
|
||||
a: 0,
|
||||
b: 10,
|
||||
want: 0,
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := safeDivide(tt.a, tt.b)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("safeDivide() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if got != tt.want {
|
||||
t.Errorf("safeDivide() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCalculateWeights(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
deductionRatio int64
|
||||
wantTrafficWeight float64
|
||||
wantTimeWeight float64
|
||||
}{
|
||||
{
|
||||
name: "zero ratio",
|
||||
deductionRatio: 0,
|
||||
wantTrafficWeight: 0,
|
||||
wantTimeWeight: 0,
|
||||
},
|
||||
{
|
||||
name: "50% ratio",
|
||||
deductionRatio: 50,
|
||||
wantTrafficWeight: 0.5,
|
||||
wantTimeWeight: 0.5,
|
||||
},
|
||||
{
|
||||
name: "75% ratio",
|
||||
deductionRatio: 75,
|
||||
wantTrafficWeight: 0.75,
|
||||
wantTimeWeight: 0.25,
|
||||
},
|
||||
{
|
||||
name: "100% ratio",
|
||||
deductionRatio: 100,
|
||||
wantTrafficWeight: 1.0,
|
||||
wantTimeWeight: 0.0,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
gotTrafficWeight, gotTimeWeight := calculateWeights(tt.deductionRatio)
|
||||
if gotTrafficWeight != tt.wantTrafficWeight {
|
||||
t.Errorf("calculateWeights() trafficWeight = %v, want %v", gotTrafficWeight, tt.wantTrafficWeight)
|
||||
}
|
||||
if gotTimeWeight != tt.wantTimeWeight {
|
||||
t.Errorf("calculateWeights() timeWeight = %v, want %v", gotTimeWeight, tt.wantTimeWeight)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCalculateProportionalAmount(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
unitPrice int64
|
||||
remaining int64
|
||||
total int64
|
||||
want int64
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "normal calculation",
|
||||
unitPrice: 100,
|
||||
remaining: 50,
|
||||
total: 100,
|
||||
want: 50,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "zero total",
|
||||
unitPrice: 100,
|
||||
remaining: 50,
|
||||
total: 0,
|
||||
want: 0,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "zero remaining",
|
||||
unitPrice: 100,
|
||||
remaining: 0,
|
||||
total: 100,
|
||||
want: 0,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "quarter remaining",
|
||||
unitPrice: 200,
|
||||
remaining: 25,
|
||||
total: 100,
|
||||
want: 50,
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := calculateProportionalAmount(tt.unitPrice, tt.remaining, tt.total)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("calculateProportionalAmount() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if got != tt.want {
|
||||
t.Errorf("calculateProportionalAmount() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCalculateNoLimitAmount(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
sub Subscribe
|
||||
order Order
|
||||
want int64
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "normal no limit calculation",
|
||||
sub: Subscribe{
|
||||
Traffic: 1000,
|
||||
Download: 300,
|
||||
Upload: 200,
|
||||
},
|
||||
order: Order{
|
||||
Amount: 1000,
|
||||
},
|
||||
want: 500, // (1000 - 300 - 200) / 1000 * 1000 = 500
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "zero traffic",
|
||||
sub: Subscribe{
|
||||
Traffic: 0,
|
||||
Download: 0,
|
||||
Upload: 0,
|
||||
},
|
||||
order: Order{
|
||||
Amount: 1000,
|
||||
},
|
||||
want: 0,
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "overused traffic",
|
||||
sub: Subscribe{
|
||||
Traffic: 1000,
|
||||
Download: 600,
|
||||
Upload: 500,
|
||||
},
|
||||
order: Order{
|
||||
Amount: 1000,
|
||||
},
|
||||
want: 0, // usedTraffic would be negative, clamped to 0
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got, err := calculateNoLimitAmount(tt.sub, tt.order)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("calculateNoLimitAmount() error = %v, wantErr %v", err, tt.wantErr)
|
||||
return
|
||||
}
|
||||
if got != tt.want {
|
||||
t.Errorf("calculateNoLimitAmount() = %v, want %v", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCalculateRemainingAmount(t *testing.T) {
|
||||
now := time.Now()
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
sub Subscribe
|
||||
order Order
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "valid no limit subscription",
|
||||
sub: Subscribe{
|
||||
StartTime: now.Add(-24 * time.Hour),
|
||||
ExpireTime: now.Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 300,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeNoLimit,
|
||||
ResetCycle: ResetCycleNone,
|
||||
DeductionRatio: 0,
|
||||
},
|
||||
order: Order{
|
||||
Amount: 1000,
|
||||
Quantity: 1,
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
{
|
||||
name: "invalid subscription",
|
||||
sub: Subscribe{
|
||||
StartTime: now,
|
||||
ExpireTime: now.Add(-24 * time.Hour), // Invalid: expire before start
|
||||
Traffic: 1000,
|
||||
Download: 300,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeMonth,
|
||||
DeductionRatio: 0,
|
||||
},
|
||||
order: Order{
|
||||
Amount: 1000,
|
||||
Quantity: 1,
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "invalid order",
|
||||
sub: Subscribe{
|
||||
StartTime: now.Add(-24 * time.Hour),
|
||||
ExpireTime: now.Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 300,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeMonth,
|
||||
DeductionRatio: 0,
|
||||
},
|
||||
order: Order{
|
||||
Amount: 1000,
|
||||
Quantity: 0, // Invalid: zero quantity
|
||||
},
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "no limit with reset cycle",
|
||||
sub: Subscribe{
|
||||
StartTime: now.Add(-24 * time.Hour),
|
||||
ExpireTime: now.Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 300,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeNoLimit,
|
||||
ResetCycle: ResetCycleMonthly, // Should return 0
|
||||
DeductionRatio: 0,
|
||||
},
|
||||
order: Order{
|
||||
Amount: 1000,
|
||||
Quantity: 1,
|
||||
},
|
||||
wantErr: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
_, err := CalculateRemainingAmount(tt.sub, tt.order)
|
||||
if (err != nil) != tt.wantErr {
|
||||
t.Errorf("CalculateRemainingAmount() error = %v, wantErr %v", err, tt.wantErr)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestCalculateRemainingAmount_NoLimitWithResetCycle(t *testing.T) {
|
||||
now := time.Now()
|
||||
sub := Subscribe{
|
||||
StartTime: now.Add(-24 * time.Hour),
|
||||
ExpireTime: now.Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 300,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeNoLimit,
|
||||
ResetCycle: ResetCycleMonthly,
|
||||
DeductionRatio: 0,
|
||||
}
|
||||
order := Order{
|
||||
Amount: 1000,
|
||||
Quantity: 1,
|
||||
}
|
||||
|
||||
got, err := CalculateRemainingAmount(sub, order)
|
||||
if err != nil {
|
||||
t.Errorf("CalculateRemainingAmount() error = %v", err)
|
||||
return
|
||||
}
|
||||
if got != 0 {
|
||||
t.Errorf("CalculateRemainingAmount() = %v, want 0", got)
|
||||
}
|
||||
}
|
||||
|
||||
// Benchmark tests
|
||||
func BenchmarkCalculateRemainingAmount(b *testing.B) {
|
||||
now := time.Now()
|
||||
sub := Subscribe{
|
||||
StartTime: now.Add(-24 * time.Hour),
|
||||
ExpireTime: now.Add(24 * time.Hour),
|
||||
Traffic: 1000,
|
||||
Download: 300,
|
||||
Upload: 200,
|
||||
UnitTime: UnitTimeMonth,
|
||||
ResetCycle: ResetCycleNone,
|
||||
DeductionRatio: 50,
|
||||
}
|
||||
order := Order{
|
||||
Amount: 1000,
|
||||
Quantity: 1,
|
||||
}
|
||||
|
||||
b.ResetTimer()
|
||||
for i := 0; i < b.N; i++ {
|
||||
_, _ = CalculateRemainingAmount(sub, order)
|
||||
}
|
||||
}
|
||||
|
||||
func BenchmarkSafeMultiply(b *testing.B) {
|
||||
for i := 0; i < b.N; i++ {
|
||||
_, _ = safeMultiply(12345, 67890)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user