Files

296 lines
6.3 KiB
Go
Raw Permalink Normal View History

package dispatch
import (
"context"
"fmt"
"sync"
"time"
)
// Task represents a unit of work that can be executed
type Task interface {
ID() string
Execute(ctx context.Context) (interface{}, error)
}
// TaskResult holds the result of a task execution
type TaskResult struct {
TaskID string
Result interface{}
Error error
Duration time.Duration
StartTime time.Time
EndTime time.Time
}
// Dispatcher manages parallel task execution
type Dispatcher struct {
mu sync.RWMutex
maxConcurrency int
results map[string]*TaskResult
inProgress map[string]bool
completed map[string]bool
semaphore chan struct{}
taskOrder []string
}
// NewDispatcher creates a new task dispatcher
func NewDispatcher(maxConcurrency int) *Dispatcher {
if maxConcurrency <= 0 {
maxConcurrency = 10
}
return &Dispatcher{
maxConcurrency: maxConcurrency,
results: make(map[string]*TaskResult),
inProgress: make(map[string]bool),
completed: make(map[string]bool),
semaphore: make(chan struct{}, maxConcurrency),
taskOrder: make([]string, 0),
}
}
// DispatchAll dispatches all tasks concurrently and waits for completion
func (d *Dispatcher) DispatchAll(ctx context.Context, tasks []Task) (map[string]*TaskResult, error) {
if len(tasks) == 0 {
return make(map[string]*TaskResult), nil
}
d.mu.Lock()
d.taskOrder = make([]string, len(tasks))
for i, task := range tasks {
d.taskOrder[i] = task.ID()
}
d.mu.Unlock()
var wg sync.WaitGroup
errChan := make(chan error, len(tasks))
// Launch all tasks concurrently with concurrency limit
for _, task := range tasks {
wg.Add(1)
go func(t Task) {
defer wg.Done()
// Acquire semaphore slot
select {
case d.semaphore <- struct{}{}:
defer func() { <-d.semaphore }()
case <-ctx.Done():
errChan <- ctx.Err()
return
}
err := d.executeTask(ctx, t)
if err != nil {
errChan <- err
}
}(task)
}
// Wait for all tasks to complete
wg.Wait()
close(errChan)
// Collect errors
var errors []error
for err := range errChan {
if err != nil {
errors = append(errors, err)
}
}
d.mu.RLock()
resultsCopy := make(map[string]*TaskResult)
for id, result := range d.results {
resultsCopy[id] = result
}
d.mu.RUnlock()
if len(errors) > 0 {
return resultsCopy, fmt.Errorf("tasks completed with %d errors", len(errors))
}
return resultsCopy, nil
}
// executeTask executes a single task and stores the result
func (d *Dispatcher) executeTask(ctx context.Context, task Task) error {
taskID := task.ID()
d.mu.Lock()
d.inProgress[taskID] = true
d.mu.Unlock()
result := &TaskResult{
TaskID: taskID,
StartTime: time.Now(),
}
// Execute task with context timeout
taskCtx, cancel := context.WithCancel(ctx)
defer cancel()
taskResult, err := task.Execute(taskCtx)
result.EndTime = time.Now()
result.Duration = result.EndTime.Sub(result.StartTime)
result.Result = taskResult
result.Error = err
d.mu.Lock()
d.results[taskID] = result
d.inProgress[taskID] = false
d.completed[taskID] = true
d.mu.Unlock()
return nil
}
// GetResult retrieves the result of a task
func (d *Dispatcher) GetResult(taskID string) (*TaskResult, bool) {
d.mu.RLock()
defer d.mu.RUnlock()
result, exists := d.results[taskID]
return result, exists
}
// GetResults retrieves all results
func (d *Dispatcher) GetResults() map[string]*TaskResult {
d.mu.RLock()
defer d.mu.RUnlock()
resultsCopy := make(map[string]*TaskResult)
for id, result := range d.results {
resultsCopy[id] = result
}
return resultsCopy
}
// GetStats returns dispatcher statistics
func (d *Dispatcher) GetStats() map[string]interface{} {
d.mu.RLock()
defer d.mu.RUnlock()
completed := len(d.completed)
totalDuration := time.Duration(0)
maxDuration := time.Duration(0)
minDuration := time.Duration(0)
for _, result := range d.results {
totalDuration += result.Duration
if result.Duration > maxDuration {
maxDuration = result.Duration
}
if minDuration == 0 || result.Duration < minDuration {
minDuration = result.Duration
}
}
avgDuration := time.Duration(0)
if completed > 0 {
avgDuration = totalDuration / time.Duration(completed)
}
return map[string]interface{}{
"total_tasks": len(d.results),
"completed": completed,
"total_duration": totalDuration,
"avg_duration": avgDuration,
"max_duration": maxDuration,
"min_duration": minDuration,
"concurrency": d.maxConcurrency,
}
}
// GetExecutionTime returns the total execution time (wallclock)
func (d *Dispatcher) GetExecutionTime() time.Duration {
d.mu.RLock()
defer d.mu.RUnlock()
if len(d.results) == 0 {
return 0
}
var minStart time.Time
var maxEnd time.Time
for _, result := range d.results {
if minStart.IsZero() || result.StartTime.Before(minStart) {
minStart = result.StartTime
}
if result.EndTime.After(maxEnd) {
maxEnd = result.EndTime
}
}
return maxEnd.Sub(minStart)
}
// GetTotalTaskDuration returns the sum of all task durations
func (d *Dispatcher) GetTotalTaskDuration() time.Duration {
d.mu.RLock()
defer d.mu.RUnlock()
total := time.Duration(0)
for _, result := range d.results {
total += result.Duration
}
return total
}
// GetSpeedup returns the speedup factor (sum of task durations / wallclock time)
func (d *Dispatcher) GetSpeedup() float64 {
totalDuration := d.GetTotalTaskDuration()
executionTime := d.GetExecutionTime()
if executionTime == 0 {
return 0
}
return float64(totalDuration) / float64(executionTime)
}
// IsComplete checks if a task is complete
func (d *Dispatcher) IsComplete(taskID string) bool {
d.mu.RLock()
defer d.mu.RUnlock()
return d.completed[taskID]
}
// AreAllComplete checks if all tasks are complete
func (d *Dispatcher) AreAllComplete() bool {
d.mu.RLock()
defer d.mu.RUnlock()
return len(d.completed) == len(d.results)
}
// GetCompletedCount returns the number of completed tasks
func (d *Dispatcher) GetCompletedCount() int {
d.mu.RLock()
defer d.mu.RUnlock()
return len(d.completed)
}
// WaitForCompletion waits for all tasks to complete or context to be cancelled
func (d *Dispatcher) WaitForCompletion(ctx context.Context) error {
ticker := time.NewTicker(10 * time.Millisecond)
defer ticker.Stop()
for {
select {
case <-ctx.Done():
return ctx.Err()
case <-ticker.C:
if d.AreAllComplete() {
return nil
}
}
}
}