Go语言学习: 实现轻量线程池
Goroutine 池的核心思想是对 Goroutine 的重用,也就是把 M 个计算任务调度到 N 个 Goroutine 上,而不是为每个计算任务分配一个独享的 Goroutine,从而提高计算资源的利用率。
这里简单地采用channel+select的实现方案,主要可分成3部分:
-
pool的创建与销毁
-
pool中的worker(Goroutine)的管理
-
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()
}