Browse Source

backtest runnnig progress

main
strange 6 months ago
parent
commit
5aac7d3065
  1. 1
      api/pub.proto
  2. 9
      api/trading.proto
  3. 2
      internal/gateway/fast_gateway.go
  4. 8
      internal/trading/backtest/trading_plan_backtester.go
  5. 3
      internal/trading/sig/indicator_context.go
  6. 16
      internal/trading/trading_grpc_server.go
  7. 51
      internal/trading/trading_service.go
  8. 6
      pkg/data/common.go
  9. 105
      pkg/strategy/mean_reversion_v1.go
  10. 55
      pkg/trade/types.go
  11. 169
      pkg/utils/progress/progress.go
  12. 45
      pkg/utils/progress/progress_test.go

1
api/pub.proto

@ -170,6 +170,7 @@ message Paging {
int32 page = 1;
int32 size = 2;
bool asc = 3; //
string sortBy = 4; //
}
//

9
api/trading.proto

@ -101,7 +101,8 @@ message ReqBacktest {
string etime = 3;
}
message RspBacktest {
int64 backtest_id = 1;
string task_id = 1;
int64 backtest_id = 2;
}
message BacktestLog {
@ -110,9 +111,15 @@ message BacktestLog {
int64 stime = 3;
int64 etime = 4;
string interval = 5; //
int32 status = 6; // 1.
int64 progress = 7;
int64 total = 8;
string pct = 9;
string taskId = 10;
}
message ReqBacktestLog {
Paging paging = 1;
bool running = 2; //
}
message RspBacktestLog {
repeated BacktestLog logs = 1;

2
internal/gateway/fast_gateway.go

@ -141,7 +141,7 @@ func (s *FastGatewayServer) reverseProxyGrpcGenericCall(c *fasthttp.RequestCtx)
// put session
ctx := context.Background()
ctx = session.PutSubject(ctx, session.NewRpcSubject("123456"))
ctx, cancel = context.WithTimeout(ctx, time.Second*20)
ctx, cancel = context.WithTimeout(ctx, time.Second*60)
defer cancel()
// todo config call options

8
internal/trading/backtest/trading_plan_backtester.go

@ -59,9 +59,11 @@ func (b *TradingPlanBacktester) Init(cash float64, plan entity.TradePlan, sr *pb
if err = sonic.UnmarshalString(plan.SigStrategyParam, &b.sigStrategyInput); err != nil {
return
}
if err = b.sigStrategy.Init(b.sigStrategyInput); err != nil {
return
}
// 填充参数默认值
// sig.FillDefaultInputs(b.sigStrategyInput, b.sigStrategy.Meta().Input)
// if err = b.sigStrategy.Init(b.sigStrategyInput); err != nil {
// return
// }
// 交易策略参数
var tradeStrategyInput, closeStrategyInput, riskStrategyInput types.Input

3
internal/trading/sig/indicator_context.go

@ -184,6 +184,9 @@ inputLoop:
case types.Input:
input = v
break inputLoop
case map[string]any:
input = v
break inputLoop
}
}
windowLoop:

16
internal/trading/trading_grpc_server.go

@ -4,6 +4,7 @@ import (
"context"
"fmt"
"sig-pub/api/pb"
"sig-pub/pkg/data"
"sig-pub/pkg/types"
"sig-pub/pkg/utils/collect"
"sig-pub/pkg/utils/misc"
@ -180,17 +181,21 @@ func (svr *TradingGrpcServer) Backtest(ctx context.Context, req *pb.ReqBacktest)
if err != nil {
return
}
err = svr.tradingService.Backtest(ctx, req.PlanId, stime.UnixMilli(), etime.UnixMilli())
taskId, err := svr.tradingService.Backtest(ctx, req.PlanId, stime.UnixMilli(), etime.UnixMilli())
if err != nil {
return
}
rsp = new(pb.RspBacktest)
rsp = &pb.RspBacktest{TaskId: taskId}
return
}
// BacktestLog 回测记录查询
func (svr *TradingGrpcServer) BacktestLog(ctx context.Context, req *pb.ReqBacktestLog) (rsp *pb.RspBacktestLog, err error) {
logs, err := svr.tradingService.BacktestLog(ctx, 10001)
if req.Paging == nil {
req.Paging = &pb.Paging{Page: 1, Size: 20, Asc: true, SortBy: "id"}
}
paging := data.PageArgs0(int(req.Paging.Page), int(req.Paging.Size), req.Paging.SortBy, req.Paging.Asc)
logs, err := svr.tradingService.BacktestLog(ctx, 10001, paging, req.Running)
if err != nil {
return
}
@ -201,6 +206,11 @@ func (svr *TradingGrpcServer) BacktestLog(ctx context.Context, req *pb.ReqBackte
Stime: l.SeriesBefore,
Etime: l.SeriesAfter,
BacktestId: l.Id,
Status: l.RunningStatus,
Progress: l.RunningProgress,
Total: l.RunningTotal,
Pct: l.RunningPct,
TaskId: l.RunningTaskId,
}
rsp.Logs = append(rsp.Logs, log)
}

51
internal/trading/trading_service.go

@ -16,6 +16,7 @@ import (
"sig-pub/pkg/types"
"sig-pub/pkg/utils/collect"
"sig-pub/pkg/utils/misc"
"sig-pub/pkg/utils/progress"
"sig-pub/pkg/utils/times"
"sig-pub/pkg/zlog"
"sort"
@ -38,6 +39,7 @@ type TradingService struct {
strategyReg *strategy.SigStrategyRegistry // 注册信号策略
signalPublisher *publish.Publisher[int64, strategy.StrategyType] // planId -> strategyType
tradingPlans *collect.SyncMap[int64, *sig.TradingPlan] // 运行中交易计划
backtestingTasks *progress.ProgressManager // 回测进度管理
}
func NewTradingService(
@ -54,6 +56,7 @@ func NewTradingService(
strategyReg: strategy.NewSigStrategyRegistry(),
signalPublisher: publish.NewPublisher[int64, strategy.StrategyType](16),
tradingPlans: collect.NewSyncMap[int64, *sig.TradingPlan](),
backtestingTasks: progress.NewProgressManager(),
}
}
@ -456,7 +459,7 @@ func (svc *TradingService) StrategySeries(ctx context.Context, req *pb.ReqStrate
}
// Backtest 回测交易计划
func (svc *TradingService) Backtest(ctx context.Context, planId, stime, etime int64) (err error) {
func (svc *TradingService) Backtest(ctx context.Context, planId, stime, etime int64) (taskId string, err error) {
plan, err := svc.tradingDataPersist.GetTradePlanById(planId)
if err != nil {
return
@ -487,22 +490,48 @@ func (svc *TradingService) Backtest(ctx context.Context, planId, stime, etime in
if err = tester.Init(10000, *plan, sr); err != nil {
return
}
w := times.NewWatch()
backtestTradingPlan, err := tester.Backtest(ctx)
if err != nil {
return
}
zlog.Infof("backtest use %s", w.ElapsedFmt(""))
w.Reset()
err = svc.tradingDataPersist.SaveBacktestTradingPlan(ctx, backtestTradingPlan)
zlog.Infof("insert backtest ret use %s", w.ElapsedFmt(""))
// 生成异步任务, 获取任务进度
taskId = svc.backtestingTasks.StartTask(context.Background(), fmt.Sprintf("backtest %d", plan.Id), func(ctx context.Context, progress progress.IProgressReporter) error {
w := times.NewWatch()
backtestTradingPlan, err := tester.Backtest(ctx)
if err != nil {
return err
}
zlog.Infof("backtest use %s", w.ElapsedFmt(""))
w.Reset()
err = svc.tradingDataPersist.SaveBacktestTradingPlan(ctx, backtestTradingPlan)
zlog.Infof("insert backtest ret use %s", w.ElapsedFmt(""))
return nil
})
return
}
// BacktestLog 回测记录查询
func (svc *TradingService) BacktestLog(ctx context.Context, userId int64) (backtestLogs []*trade.BacktestTradingPlan, err error) {
func (svc *TradingService) BacktestLog(ctx context.Context, userId int64, paging data.Page, running bool) (backtestLogs []*trade.BacktestTradingPlan, err error) {
// todo userid from ctx
backtestLogs, err = svc.tradingDataPersist.ListBacktestLogs(userId)
if err != nil {
return
}
// 在首页加载运行中任务
if running && paging.Page == 1 {
runningTasks := svc.backtestingTasks.ListRunningTask()
var runningLogs []*trade.BacktestTradingPlan
for _, task := range runningTasks {
progress, total := task.Progress.Get()
runningLogs = append(runningLogs, &trade.BacktestTradingPlan{
RunningStatus: int32(task.Status.Load()),
RunningTaskId: task.ID,
RunningProgress: progress,
RunningTotal: total,
RunningPct: fmt.Sprintf("%.2f", float64(progress)/float64(total)),
})
}
if len(runningLogs) > 0 {
backtestLogs = append(runningLogs, backtestLogs...)
}
}
return
}

6
pkg/data/common.go

@ -49,6 +49,10 @@ func PageArgs(ctx *gin.Context) Page {
sortBy := ctx.Query("sortBy")
asc, _ := strconv.Atoi(ctx.Query("asc"))
return PageArgs0(page, pageSize, sortBy, asc == 1)
}
func PageArgs0(page, pageSize int, sortBy string, asc bool) Page {
if page <= 0 {
page = 1
}
@ -60,7 +64,7 @@ func PageArgs(ctx *gin.Context) Page {
Page: page,
PageSize: pageSize,
SortBy: sortBy,
Asc: asc == 1,
Asc: asc,
Offset: offset,
Limit: limit,
}

105
pkg/strategy/mean_reversion_v1.go

@ -1,16 +1,21 @@
package strategy
import (
"sort"
"sig-pub/pkg/types"
)
// MeanReversionV1
type MeanReversionV1 struct {
IIntervalSigStrategy
interval types.Interval
period int
threshold float64
buckets int
interval types.Interval
period int
threshold float64
buckets int
dominanceRatio float64 // POC volume dominance ratio (vs 2nd highest)
rsiPeriod int
rsiThreshold float64
}
func (s *MeanReversionV1) New() ISigStrategy {
@ -20,12 +25,15 @@ func (s *MeanReversionV1) New() ISigStrategy {
func (s *MeanReversionV1) Meta() StrategyMeta {
return StrategyMeta{
Name: "MeanReversionV1",
Desc: "VRVP Mean Reversion Strategy",
Desc: "VRVP Mean Reversion Strategy with Volume Dominance and RSI Filter",
Input: []types.InputArg{
{Name: "interval", Type: types.InputTypeString, Desc: "Target Interval (e.g., 1m, 1h)", Default: "1m"},
{Name: "period", Type: types.InputTypeInt, Desc: "VRVP calculation window", Default: 100},
{Name: "threshold", Type: types.InputTypeUFloat, Desc: "Reversion Threshold Ratio (e.g. 0.01)", Default: 0.01},
{Name: "buckets", Type: types.InputTypeInt, Desc: "VRVP Buckets", Default: 24},
{Name: "dominance_ratio", Type: types.InputTypeUFloat, Desc: "POC Volume Dominance Ratio (e.g. 1.2)", Default: 1.2},
{Name: "rsi_period", Type: types.InputTypeInt, Desc: "RSI Period", Default: 14},
{Name: "rsi_threshold", Type: types.InputTypeUFloat, Desc: "RSI Threshold (e.g. 30 for 30/70)", Default: 30},
},
}
}
@ -44,61 +52,102 @@ func (s *MeanReversionV1) Init(input types.Input) (err error) {
if s.buckets <= 0 {
s.buckets = 24
}
s.dominanceRatio = input.Float("dominance_ratio")
if s.dominanceRatio < 1.0 {
s.dominanceRatio = 1.0
}
s.rsiPeriod = input.Int("rsi_period")
if s.rsiPeriod <= 0 {
s.rsiPeriod = 14
}
s.rsiThreshold = input.Float("rsi_threshold")
if s.rsiThreshold <= 0 || s.rsiThreshold >= 50 {
s.rsiThreshold = 30 // Default to standard 30 (implying 70 upper)
}
return
}
func (s *MeanReversionV1) CandlePeriods(ctx IIntervalSigStrategyContext) (iss *types.IntervalState[int16]) {
iss = types.NewIntervalState[int16]()
iss.Set(s.interval, int16(s.period))
// We need enough candles for both VRVP and RSI
// VRVP needs 'period' candles.
// RSI needs 'rsiPeriod' candles (maybe +1).
// To be safe, we take the max.
needed := int16(s.period)
if int16(s.rsiPeriod+5) > needed {
needed = int16(s.rsiPeriod + 5)
}
iss.Set(s.interval, needed)
return
}
func (s *MeanReversionV1) Update(ctx IIntervalSigStrategyContext) (side types.Side) {
// Calculate VRVP period candles
// 1. Get VRVP Summary
summaryObj := ctx.SummaryIndicator(s.interval, "VRVP", map[string]any{"buckets": s.buckets})
// Calculate for the last 'period' candles
summaryAny, ok := summaryObj.Summary(0, int16(s.period))
if !ok {
return
}
vrvpSummary, ok := summaryAny.(*types.VRVPSummary)
if !ok || vrvpSummary == nil || len(vrvpSummary.Buckets) == 0 {
if !ok || vrvpSummary == nil || len(vrvpSummary.Buckets) < 2 {
return
}
// Find POC (Point of Control) - Bucket with max volume
var maxVol float64
var pocPrice float64
found := false
// 2. Find POC (Point of Control) and Second Highest Volume
// Create a slice of buckets to sort
type volBucket struct {
Price float64
Volume float64
}
sortedBuckets := make([]volBucket, len(vrvpSummary.Buckets))
for i, b := range vrvpSummary.Buckets {
sortedBuckets[i] = volBucket{Price: b.Price, Volume: b.Volume}
}
for _, bucket := range vrvpSummary.Buckets {
if bucket.Volume > maxVol {
maxVol = bucket.Volume
pocPrice = bucket.Price
found = true
}
// Sort descending by volume
sort.Slice(sortedBuckets, func(i, j int) bool {
return sortedBuckets[i].Volume > sortedBuckets[j].Volume
})
pocBucket := sortedBuckets[0]
secondBucket := sortedBuckets[1]
// Check Dominance
if pocBucket.Volume < secondBucket.Volume*s.dominanceRatio {
// POC is not dominant enough
return
}
if !found {
pocPrice := pocBucket.Price
if pocPrice <= 0 {
return
}
// Get Current Price
// 3. Get Current Price
k := ctx.Get(s.interval, 0)
currentPrice := k.CloseF64()
// Deviation ratio
if pocPrice <= 0 {
return
}
// 4. Calculate RSI for Confirmation
rsiSeries := ctx.Indicator(s.interval, "RSI", map[string]any{"window": s.rsiPeriod})
currentRSI := rsiSeries.Get(0)
// 5. Generate Signal
deviation := (currentPrice - pocPrice) / pocPrice
if deviation > s.threshold {
// 当前价格高于 POC, 预期回落
side = types.SideShort
// Price is significantly higher than POC, expect reversion (Sell)
// Filter: RSI should be overbought (> 100 - threshold, e.g. > 70)
if currentRSI > (100 - s.rsiThreshold) {
side = types.SideShort
}
} else if deviation < -s.threshold {
// 当前价格低于 POC, 预期回升
side = types.SideLong
// Price is significantly lower than POC, expect reversion (Buy)
// Filter: RSI should be oversold (< threshold, e.g. < 30)
if currentRSI < s.rsiThreshold {
side = types.SideLong
}
}
return

55
pkg/trade/types.go

@ -57,31 +57,36 @@ type Position struct {
// BacktestTradingPlan 交易计划回测结果
type BacktestTradingPlan struct {
Id int64 `json:"id" gorm:"column:id;primaryKey"` // 测试id
UserId int64 `json:"userId" gorm:"column:user_id"` // 用户id
PlanId int64 `json:"planId" gorm:"column:plan_id"` // 交易计划id
InstId string `json:"instId" gorm:"column:inst_id"` // 交易产品id
Exchange pb.ExchangeType `json:"exchange" gorm:"column:exchange"` // 交易所
Interval string `json:"interval" gorm:"column:interval"` // 交易周期
SeriesBefore int64 `json:"seriesBefore" gorm:"column:series_before"` // 回测周期开始时间
SeriesAfter int64 `json:"seriesAfter" gorm:"column:series_after"` // 回测周期结束时间
Ctime int64 `json:"ctime" gorm:"column:ctime"` // 创建时间(测试时间)
Etime int64 `json:"etime" gorm:"column:etime"` // 测试结束时间
Cash float64 `json:"cash" gorm:"column:cash"` // 起始金额
EndCash float64 `json:"endCash" gorm:"column:end_cash"` // 结束金额
Profit float64 `json:"profit" gorm:"column:profit"` // 利润
Singals int `json:"singals" gorm:"column:singals"` // 交易信号数
TotalTrades int `json:"totalTrades" gorm:"column:total_trades"` // 总单数
WinningTrades int `json:"winningTrades" gorm:"column:winning_trades"` // 盈利单数
LosingTrades int `json:"losingTrades" gorm:"column:losing_trades"` // 亏损单数
Fee float64 `json:"fee" gorm:"column:fee"` // 总手续费
MaxDrawdown float64 `json:"maxDrawdown" gorm:"column:max_drawdown"` // 最大回撤
SharpeRatio float64 `json:"sharpeRatio" gorm:"column:sharpe_ratio"` // 夏普比率
SortinoRatio float64 `json:"sortinoRatio" gorm:"column:sortino_ratio"` // 索提诺比率
CalmarRatio float64 `json:"calmarRatio" gorm:"column:calmar_ratio"` // 卡尔玛比率
ProfitFactor float64 `json:"profitFactor" gorm:"column:profit_factor"` // 盈利因子
WinRate float64 `json:"winRate" gorm:"column:win_rate"` // 胜率
Trades []*TradeOrder `json:"-" gorm:"-"` // 回测交易单
Id int64 `json:"id" gorm:"column:id;primaryKey"` // 测试id
UserId int64 `json:"userId" gorm:"column:user_id"` // 用户id
PlanId int64 `json:"planId" gorm:"column:plan_id"` // 交易计划id
InstId string `json:"instId" gorm:"column:inst_id"` // 交易产品id
Exchange pb.ExchangeType `json:"exchange" gorm:"column:exchange"` // 交易所
Interval string `json:"interval" gorm:"column:interval"` // 交易周期
SeriesBefore int64 `json:"seriesBefore" gorm:"column:series_before"` // 回测周期开始时间
SeriesAfter int64 `json:"seriesAfter" gorm:"column:series_after"` // 回测周期结束时间
Ctime int64 `json:"ctime" gorm:"column:ctime"` // 创建时间(测试时间)
Etime int64 `json:"etime" gorm:"column:etime"` // 测试结束时间
Cash float64 `json:"cash" gorm:"column:cash"` // 起始金额
EndCash float64 `json:"endCash" gorm:"column:end_cash"` // 结束金额
Profit float64 `json:"profit" gorm:"column:profit"` // 利润
Singals int `json:"singals" gorm:"column:singals"` // 交易信号数
TotalTrades int `json:"totalTrades" gorm:"column:total_trades"` // 总单数
WinningTrades int `json:"winningTrades" gorm:"column:winning_trades"` // 盈利单数
LosingTrades int `json:"losingTrades" gorm:"column:losing_trades"` // 亏损单数
Fee float64 `json:"fee" gorm:"column:fee"` // 总手续费
MaxDrawdown float64 `json:"maxDrawdown" gorm:"column:max_drawdown"` // 最大回撤
SharpeRatio float64 `json:"sharpeRatio" gorm:"column:sharpe_ratio"` // 夏普比率
SortinoRatio float64 `json:"sortinoRatio" gorm:"column:sortino_ratio"` // 索提诺比率
CalmarRatio float64 `json:"calmarRatio" gorm:"column:calmar_ratio"` // 卡尔玛比率
ProfitFactor float64 `json:"profitFactor" gorm:"column:profit_factor"` // 盈利因子
WinRate float64 `json:"winRate" gorm:"column:win_rate"` // 胜率
Trades []*TradeOrder `json:"-" gorm:"-"` // 回测交易单
RunningStatus int32 `json:"-" gorm:"-"` // 运行中回测状态
RunningTaskId string `json:"-" gorm:"-"` // 运行中回测任务id
RunningProgress int64 `json:"-" gorm:"-"`
RunningTotal int64 `json:"-" gorm:"-"`
RunningPct string `json:"-" gorm:"-"`
}
func (BacktestTradingPlan) TableName() string {

169
pkg/utils/progress/progress.go

@ -0,0 +1,169 @@
package progress
import (
"context"
"fmt"
"sig-pub/pkg/types"
"sig-pub/pkg/zlog"
"sync"
"sync/atomic"
"time"
)
// ProgressReporter 进度上报接口
type IProgressReporter interface {
SetTotal(total int64)
SetProgress(progress int64)
AddProgress(progress int64)
Get() (progress, total int64)
}
type reporter struct {
total atomic.Int64
progress atomic.Int64
}
func (r *reporter) SetTotal(total int64) {
r.total.Store(total)
}
func (r *reporter) SetProgress(progress int64) {
r.progress.Store(progress)
}
func (r *reporter) AddProgress(progress int64) {
r.progress.Add(progress)
}
func (r *reporter) Get() (progress, total int64) {
return r.progress.Load(), r.total.Load()
}
// ==================== 任务状态 ====================
type Status int32
const (
// StatusPending Status = "pending"
StatusRunning Status = 1 //"running"
StatusSuccess Status = 2 //"success"
StatusFailed Status = 3 //"failed"
StatusCancelled Status = 4 //"cancelled"
)
// Task 任务结构体(带锁)
type Task struct {
ID string `json:"id"` // 任务id
Name string `json:"name"`
Status atomic.Int64 `json:"status"`
Error error `json:"error,omitempty"`
Ctime int64 `json:"ctime"`
Etime int64 `json:"etime,omitempty"`
Progress IProgressReporter
cancel context.CancelFunc
}
// ProgressManager 管理器核心
type ProgressManager struct {
nextID atomic.Int64
tasks sync.Map // string -> *Task
completeTasks *types.RingSeries[*Task]
mu sync.RWMutex
}
// NewProgressManager 创建管理器实例
func NewProgressManager() *ProgressManager {
return &ProgressManager{
completeTasks: types.NewRingSeries[*Task](100, 0),
}
}
func (m *ProgressManager) generateID() string {
id := m.nextID.Add(1)
return fmt.Sprintf("task_%d_%d", id, time.Now().UnixNano()%100000)
}
func (m *ProgressManager) StartTask(ctx context.Context, name string, fn func(ctx context.Context, progress IProgressReporter) error) string {
id := m.generateID()
ctx, cancel := context.WithCancel(ctx)
task := &Task{
ID: id,
Name: name,
Status: atomic.Int64{},
Error: nil,
Ctime: time.Now().UnixMilli(),
Etime: 0,
Progress: &reporter{},
cancel: cancel,
}
m.tasks.Store(id, task)
go m.execTask(ctx, task, fn)
return id
}
func (m *ProgressManager) execTask(ctx context.Context, task *Task, fn func(ctx context.Context, progress IProgressReporter) error) {
defer func() {
if r := recover(); r != nil {
if !task.Status.CompareAndSwap(int64(StatusRunning), int64(StatusFailed)) {
zlog.Error("task status not running, can not set panic failed: %s, name=%s, status=%d", task.ID, task.Name, task.Status.Load())
}
task.Error = fmt.Errorf("panic: %v", r)
task.Etime = time.Now().UnixMilli()
}
}()
defer func() {
m.tasks.Delete(task.ID)
m.mu.Lock()
defer m.mu.Unlock()
m.completeTasks.Push(task)
}()
task.Status.Store(int64(StatusRunning))
err := fn(ctx, task.Progress)
if err != nil {
if !task.Status.CompareAndSwap(int64(StatusRunning), int64(StatusFailed)) {
zlog.Error("task status not running, can not set failed: %s, name=%s, status=%d", task.ID, task.Name, task.Status.Load())
}
task.Error = err
task.Etime = time.Now().UnixMilli()
return
}
task.Etime = time.Now().UnixMilli()
if !task.Status.CompareAndSwap(int64(StatusRunning), int64(StatusSuccess)) {
zlog.Error("task status not running, can not set success: %s, name=%s, status=%d", task.ID, task.Name, task.Status.Load())
}
}
func (m *ProgressManager) GetTask(id string) (*Task, bool) {
task, ok := m.tasks.Load(id)
if !ok {
return nil, false
}
return task.(*Task), true
}
func (m *ProgressManager) ListTask() ([]*Task, bool) {
runningTasks := m.ListRunningTask()
m.mu.RLock()
defer m.mu.RUnlock()
tasks, ok := m.completeTasks.Series(0, m.completeTasks.Length())
if !ok {
return nil, false
}
return append(runningTasks, tasks...), true
}
func (m *ProgressManager) ListRunningTask() (tasks []*Task) {
m.tasks.Range(func(key, value any) bool {
task := value.(*Task)
if task.Status.Load() == int64(StatusRunning) {
tasks = append(tasks, task)
}
return true
})
return
}

45
pkg/utils/progress/progress_test.go

@ -0,0 +1,45 @@
package progress
import (
"context"
"fmt"
"testing"
"time"
)
func TestProgressManager(t *testing.T) {
c := make(chan struct{})
m := NewProgressManager()
taskId := m.StartTask(context.Background(), "testing", func(ctx context.Context, progress IProgressReporter) error {
defer close(c)
progress.SetTotal(100)
for range 100 {
progress.AddProgress(1)
time.Sleep(110 * time.Millisecond)
}
return nil
})
go func() {
task, _ := m.GetTask(taskId)
for {
select {
case <-c:
return
default:
}
progress, total := task.Progress.Get()
if total != 0 {
fmt.Printf("task %s progress: %.2f%%\n", task.Name, float64(progress)/float64(total)*100)
if progress == total {
return
}
}
time.Sleep(300 * time.Millisecond)
}
}()
<-c
}
Loading…
Cancel
Save