diff --git a/api/pub.proto b/api/pub.proto index cbcb434..83bf525 100644 --- a/api/pub.proto +++ b/api/pub.proto @@ -170,6 +170,7 @@ message Paging { int32 page = 1; int32 size = 2; bool asc = 3; // 升序排序 + string sortBy = 4; // 排序字段 } // 指标元数据 diff --git a/api/trading.proto b/api/trading.proto index c0602e9..de2a320 100644 --- a/api/trading.proto +++ b/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; diff --git a/internal/gateway/fast_gateway.go b/internal/gateway/fast_gateway.go index 1406d66..e381006 100644 --- a/internal/gateway/fast_gateway.go +++ b/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 diff --git a/internal/trading/backtest/trading_plan_backtester.go b/internal/trading/backtest/trading_plan_backtester.go index fab64fe..57fa0ae 100644 --- a/internal/trading/backtest/trading_plan_backtester.go +++ b/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 diff --git a/internal/trading/sig/indicator_context.go b/internal/trading/sig/indicator_context.go index 3be6988..8ca7206 100644 --- a/internal/trading/sig/indicator_context.go +++ b/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: diff --git a/internal/trading/trading_grpc_server.go b/internal/trading/trading_grpc_server.go index 10f0c56..36a515a 100644 --- a/internal/trading/trading_grpc_server.go +++ b/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) } diff --git a/internal/trading/trading_service.go b/internal/trading/trading_service.go index efb4201..b399470 100644 --- a/internal/trading/trading_service.go +++ b/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 } diff --git a/pkg/data/common.go b/pkg/data/common.go index be5a3b9..bb388aa 100644 --- a/pkg/data/common.go +++ b/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, } diff --git a/pkg/strategy/mean_reversion_v1.go b/pkg/strategy/mean_reversion_v1.go index 06a5ca2..14f5ea6 100644 --- a/pkg/strategy/mean_reversion_v1.go +++ b/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 diff --git a/pkg/trade/types.go b/pkg/trade/types.go index 4e1ae43..aab2ee0 100644 --- a/pkg/trade/types.go +++ b/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 { diff --git a/pkg/utils/progress/progress.go b/pkg/utils/progress/progress.go new file mode 100644 index 0000000..e618d3a --- /dev/null +++ b/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 +} diff --git a/pkg/utils/progress/progress_test.go b/pkg/utils/progress/progress_test.go new file mode 100644 index 0000000..9e8e11d --- /dev/null +++ b/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 +}