Go语言学习: 实现轻量线程池


Goroutine 池的核心思想是对 Goroutine 的重用,也就是把 M 个计算任务调度到 N 个 Goroutine 上,而不是为每个计算任务分配一个独享的 Goroutine,从而提高计算资源的利用率。

这里简单地采用channel+select的实现方案,主要可分成3部分:

  1. pool的创建与销毁

  2. pool中的worker(Goroutine)的管理

  3. task的提交与调度

定义一个结构体Pool,它应当具有一些属性:

capacity 是 pool 的一个属性,代表整个 pool 中 worker 的最大容量。使用一个带缓冲的 channel:active,作为 worker 的“计数器”,这种 channel 使用模式就是计数信号量,那么其对应的数据类型就是struct{}。

当 active channel 可写时,就创建一个 worker,用于处理用户通过 Schedule 函数提交的待处理的请求。当 active channel 满了的时候,pool 就会停止 worker 的创建,直到某个 worker 因故退出,active channel 又空出一个位置时,pool 才会创建新的 worker 填补那个空位。

把用户要提交给 workerpool 执行的请求抽象为一个 Task。Task 的提交与调度也很简单:Task 通过 Schedule 函数提交到一个 task channel 中,已经创建的 worker 将从这个 task channel 中读取 task 并执行。

定义了如下结构体:

type Pool struct {
    capacity int         // workerpool大小

    active chan struct{} // 对应上图中的active channel
    tasks  chan Task     // 对应上图中的task channel

    wg   sync.WaitGroup  // 用于在pool销毁时等待所有worker退出
    quit chan struct{}   // 用于通知各个worker退出的信号channel
}

下面实现一个New函数,用于创建一个pool类型实例,并将pool池的worker管理机制运行起来。

func New(capacity int) *Pool {
    if capacity <= 0 {
        capacity = defaultCapacity
    }
    if capacity > maxCapacity {
        capacity = maxCapacity
    }

    p := &Pool{
        capacity: capacity,
        tasks:    make(chan Task),
        quit:     make(chan struct{}),
        active:   make(chan struct{}, capacity),
    }
    fmt.Printf("workpool start\n")
    go p.run()
    return p
}

New函数接收一个参数capacity用于指定workerpool池的容量,这个参数用于控制wokerpool最多只能有capacity个worker,共同处理用户提交的任务请求。函数开始处会检查传参是否合理。

Pool类型实例变量p完成初始化后,创建一个新的Goroutine,用于workerpool进行管理,这个Goroutine用于对workerpool进行管理,这个goroutine执行的是pool类型的run方法。

func (p *Pool) run() {
    index := 0
    for {
        select {
        case <-p.quit:
            return
        case p.active <- struct{}{}:
            index++
            p.newWorker(index)
        }
    }
}

run方法内是一个无限循环,循环体中使用select监视pool类型实例的两个channel:quit和active。这种在for循环中使用select监视多个channel的实现,在Go代码中十分常见。

当接收到来自quit channel的退出"信号"时,这个Goroutine就会结束运行。而当active channel可写时,run方法就会创建一个新的worker Goroutine。此外,为了方便在程序中区分各个worker输出的日志,这里将一个从1开始的变量index作为worker的编号,并将它以参数形式传给创建worker的方法。

将创建新的 worker goroutine 的职责,封装到一个名为 newWorker 的方法中:

func (p *Pool) newWorker(i int) {
    p.wg.Add(1)
    go func() {
        defer func() {
            if err := recover(); err != nil {
                fmt.Printf("worker[%03d]: recover panic[%s] and exit\n", i, err)
                <-p.active
            }
            p.wg.Done()
        }()

        fmt.Printf("worker[%03d]: start\n", i)

        for {
            select {
            case <-p.quit:
                fmt.Printf("worker[%03d]: exit\n", i)
                <-p.active
                return
            case t := <-p.tasks:
                fmt.Printf("worker[%03d]: receive a task\n", i)
                t()
            }
        }
    }()
}

在创建一个新的 worker goroutine 之前,newWorker 方法会先调用 p.wg.Add 方法将 WaitGroup 的等待计数加一。由于每个 worker 运行于一个独立的 Goroutine 中,newWorker 方法通过 go 关键字创建了一个新的 Goroutine 作为 worker。

新 worker 的核心,依然是一个基于 for-select 模式的循环语句,在循环体中,新 worker 通过 select 监视 quit 和 tasks 两个 channel。和前面的 run 方法一样,当接收到来自 quit channel 的退出“信号”时,这个 worker 就会结束运行。tasks channel 中放置的是用户通过 Schedule 方法提交的请求,新 worker 会从这个 channel 中获取最新的 Task 并运行这个 Task。

Task 是一个对用户提交的请求的抽象,它的本质就是一个函数类型:

type Task func()

在新 worker 中,为了防止用户提交的 task 抛出 panic,进而导致整个 workerpool 受到影响,在 worker 代码的开始处,使用了 defer+recover 对 panic 进行捕捉,捕捉后 worker 也是要退出的,于是还通过<-p.active更新了 worker 计数器。并且一旦 worker goroutine 退出,p.wg.Done 也需要被调用,这样可以减少 WaitGroup 的 Goroutine 等待数量。

workerpool 提供给用户提交请求的导出方法 Schedule:

var ErrWorkerPoolFreed = errors.New("workerpool freed")       // workerpool已终止运行

func (p *Pool) Schedule(t Task) error {
    select {
    case <-p.quit:
        return ErrWorkerPoolFreed
    case p.tasks <- t:
        return nil
    }
}

这里要注意的是,这里的 Pool 结构体中的 tasks 是一个无缓冲的 channel,如果 pool 中 worker 数量已达上限,而且 worker 都在处理 task 的状态,那么 Schedule 方法就会阻塞,直到有 worker 变为 idle 状态来读取 tasks channel,schedule 的调用阻塞才会解除。

完整代码:

package main

import (
	"errors"
	"fmt"
	"sync"
	"time"
)

type Task func()

const (
	maxCapacity     = 20
	defaultCapacity = 10
)

type Pool struct {
	capacity int
	wg       sync.WaitGroup
	active   chan struct{}
	quit     chan struct{}
	tasks    chan Task
}

func New(capacity int) *Pool {
	if capacity <= 0 {
		capacity = defaultCapacity
	}
	if capacity > maxCapacity {
		capacity = maxCapacity
	}
	p := &Pool{
		capacity: capacity,
		active:   make(chan struct{}, capacity),
		quit:     make(chan struct{}),
		tasks:    make(chan Task),
	}
	go p.run()
	return p
}

func (p *Pool) run() {
	index := 0
	for {
		select {
		case <-p.quit:
			return
		case p.active <- struct{}{}:
			index++
			p.newWorker(index)
		}
	}
}

func (p *Pool) newWorker(index int) {
	p.wg.Add(1)
	go func() {
		defer func() {
			if err := recover(); err != nil {
				fmt.Printf("worker[%d] panic[%s] and exit\n", index, err)
				<-p.active
			}
			p.wg.Done()
		}()

		fmt.Printf("worker[%03d] start\n", index)
		for {
			select {
			case <-p.quit:
				fmt.Printf("worker[%03d] exit\n", index)
				<-p.active
				return
			case t := <-p.tasks:
				fmt.Printf("worker[%03d] received a task\n", index)
				t()
			}
		}
	}()
}

var ErrWorkerPoolFreed = errors.New("workerpool freed")

func (p *Pool) Schedule(t Task) error {
	select {
	case <-p.quit:
		return ErrWorkerPoolFreed
	case p.tasks <- t:
		return nil
	}
}

func (p *Pool) Free() {
	close(p.quit)
	p.wg.Wait()
	fmt.Println("workpool freed")
}

func main() {
	p := New(5)
	for i := 0; i < 10; i++ {
		err := p.Schedule(func() {
			time.Sleep(time.Second * 3)
		})
		if err != nil {
			println("task: ", i, "err:", err)
		}
	}
	p.Free()
}