红魔咖啡馆

头发越掉越多,头发越掉越少

0%

【Golang】控制并发的三种方式

参考视频:

【用 10 分鐘了解 Go 語言 context package 使用場景及介紹】

【Go语言中对Context的一些见解,不了解Context的同学可以进来看一看。】

控制并发

控制并发常用的有三种方式

WaitGroup

go中的WaitGroup位于sync库内,用来等待一组并发任务完成

可以理解为一个计数器+阻塞等待器,主要方法有三种:

  • Add(delta int):用于增加/减少计数器的值,delta可正可负
  • Done():等价于Add(-1),表示一个任务已经完成
  • Wait():阻塞直到计数器变为0

注意:

  • 计数器不能小于0,否则会panic
  • Wait要放在所有Goroutine启动后,否则会死锁
  • Add一定要在启动Goroutine的Goroutine中调用
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
package main

import (
	"fmt"
	"sync"
	"time"
)
func main() {
	var wg sync.WaitGroup

	wg.Add(2)
	go func() {
		defer wg.Done()
		time.Sleep(2*time.Second)
		println("Hello1")
	}()
	go func() {
		defer wg.Done()
		time.Sleep(1*time.Second)
		println("Hello2")
	}()
	wg.Wait()
	fmt.Println("All goroutines finished executing")
}

原理:

  • WaitGroup内部主要由两个字段构成:
    • 表示当前计数器的state(原子操作)
    • 信号量sema用于唤醒Wait阻塞的Goroutine
  • 首先Add使用原子操作修改计数器
  • Done使Add(-1),若修改后计数器为0,则通过sema唤醒所有被Wait阻塞的Goroutine
  • Wait内部判断计数器若大于0,就通过sema睡眠,被唤醒后返回

Channel

如果需要主动通知停止,可以使用Channel+Select

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
package main

import (
	"fmt"
	"time"
)

func main() {
	stop := make(chan struct{})
	go func() {
		for {
			select {
			case <-stop:
				fmt.Println("stop channel")
				return
			default:
				fmt.Println("working")
				time.Sleep(1*time.Second)
			}
		}
	}()

	time.Sleep(5*time.Second)
	close(stop)
	time.Sleep(1*time.Second)
}

context

若有多个Goroutine或Groutine中又有Goroutine,这时就适合用context

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
package main

import (
	"context"
	"time"
)

func main() {
	ctx, cancel := context.WithCancel(context.Background())
	go worker(ctx, "worker1")
	go worker(ctx, "worker2")
	go worker(ctx, "worker3")	
	time.Sleep(5*time.Second)
	cancel()
	time.Sleep(1*time.Second)


}
func worker(ctx context.Context, name string) {
	go func() {
		for {
			select {
			case <-ctx.Done():
				println(name, "done")
				return
			default:
				println(name, "working")
				time.Sleep(1*time.Second)
			}
		}
	}()
}

实例

下面创建了一个客户端和一个服务端,并从客户端向服务端发送了一条消息

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
package main

import (
	"log"
	"net"
	"time"
)

func main() {
	// server
	go func() {
		ln, err := net.Listen("tcp", ":1234")
		if err != nil {
			panic(err)
		}
		for {
			conn, err := ln.Accept()
			if err != nil {
				panic(err)
			}

			go func() {
				b := make([]byte, 1234)
				n, err := conn.Read(b)
				if err != nil {
					log.Println(err)
					return 
				}
				log.Println(string(b[:n]))
			}()
		}
	}()
	time.Sleep(100*time.Millisecond)
	// client
	cc, err := net.Dial("tcp", "localhost:1234")
	if err != nil {
		panic(err)
	}
	cc.Write([]byte("hello world"))
	time.Sleep(100*time.Second)
}

可以将客户端和服务端的连接看作一个管道,永远不会关闭

若我们想向服务端写入1g数据,若在1s内没有写完就返回

这个要求需要我们在写入数据的时候同时监测时间,但是有一个问题:我们无法通过在函数外部检测函数内的某个状态来决定何时跳出函数,只能在函数内主动跳出,因此我们可以开一个channel和线程,用AfterFunc计时,1s后关闭channel,当写入数据的线程检测channel关闭,就返回超时

这里检测channel是否关闭就是在函数内设置的一个关键点位,用于检测是否跳出函数

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
package main

import (
	"fmt"
	"log"
	"net"
	"strings"
	"time"
)

func main() {
	// server
	go func() {
		ln, err := net.Listen("tcp", ":1234")
		if err != nil {
			panic(err)
		}
		for {
			conn, err := ln.Accept()
			if err != nil {
				panic(err)
			}

			go func() {
				b := make([]byte, 1234)
				n, err := conn.Read(b)
				if err != nil {
					log.Println(err)
					return 
				}
				log.Println(string(b[:n]))
			}()
		}
	}()
	time.Sleep(100*time.Millisecond)
	// client
	cc, err := net.Dial("tcp", "localhost:1234")
	if err != nil {
		panic(err)
	}

	// 记录超时线程
	ch := make(chan struct{})
	go func() {
		time.AfterFunc(1*time.Second, func ()  {
			close(ch)
		})
	}()

	// 写入数据线程
	for i := 0; i < 1024; i++ {
		time.Sleep(10*1*time.Millisecond)
		cc.Write([]byte(strings.Repeat("a", 1024*1024)))
			select {
			case <-ch:
				fmt.Println("time exceeded", i)
				return 
			default:
			}
			
		}
	time.Sleep(100*time.Second)
}

这个channel的思路可以用context实现

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
// 使用上下文
ctx, cancel := context.WithTimeout(context.Background(), 1*time.Second)
defer cancel()
// // 记录超时线程
// ch := make(chan struct{})
// go func() {
// 	time.AfterFunc(1*time.Second, func ()  {
// 		close(ch)
// 	})
// }()

// 写入数据线程
for i := 0; i < 1024; i++ {
    time.Sleep(10*1*time.Millisecond)
    cc.Write([]byte(strings.Repeat("a", 1024*1024)))
        select {
        case <-ctx.Done():
            fmt.Println("time exceeded", i)
            return 
        default:
        }
}

或者使用ctx, cancel := context.WithDeadline(context.Background(), time.Now().Add(1*time.Second))或者ctx, cancel := context.WithCancel(context.Background()),但后者不要直接defer,而是需要自己实现cancel,即再开一个routine,休眠一秒后cancel()