1. 增加代码注释

2. hook了disconnect消息
This commit is contained in:
lwch
2021-08-12 17:39:10 +08:00
parent fe1718433c
commit e643d0cd7a
4 changed files with 131 additions and 100 deletions
+7 -93
View File
@@ -20,6 +20,7 @@ type Client struct {
tunnels map[string]*tunnel
}
// New create client
func New(cfg *global.Configure, conn *network.Conn) *Client {
return &Client{
cfg: cfg,
@@ -28,6 +29,7 @@ func New(cfg *global.Configure, conn *network.Conn) *Client {
}
}
// Run main loop
func (c *Client) Run() {
err := c.writeHandshake()
runtime.Assert(err)
@@ -61,19 +63,7 @@ func (c *Client) Run() {
}
}
func (c *Client) writeHandshake() error {
var msg network.Msg
msg.XType = network.Msg_handshake
msg.From = c.cfg.ID
msg.To = "server"
msg.Payload = &network.Msg_Hsp{
Hsp: &network.HandshakePayload{
Enc: c.cfg.Enc[:],
},
}
return c.conn.WriteMessage(&msg, 5*time.Second)
}
// handleTcpTunnel local listen to tcp tunnel
func (c *Client) handleTcpTunnel(t global.Tunnel) {
defer func() {
if err := recover(); err != nil {
@@ -111,31 +101,12 @@ func (c *Client) handleTcpTunnel(t global.Tunnel) {
}
}
// handleUdpTunnel local listen to udp tunnel
func (c *Client) handleUdpTunnel(t global.Tunnel) {
// TODO
}
func (c *Client) sendConnect(id string, t global.Tunnel) {
tp := network.ConnectRequest_tcp
if t.Type != "tcp" {
tp = network.ConnectRequest_udp
}
var msg network.Msg
msg.From = c.cfg.ID
msg.To = t.Target
msg.XType = network.Msg_connect_req
msg.Payload = &network.Msg_Creq{
Creq: &network.ConnectRequest{
Id: id,
Name: t.Name,
XType: tp,
Addr: t.RemoteAddr,
Port: uint32(t.RemotePort),
},
}
c.conn.WriteMessage(&msg, 5*time.Second)
}
// handleConnect handle connect request message from remote, local dial to remomte addr
func (c *Client) handleConnect(from, to string, req *network.ConnectRequest) {
dial := "tcp"
if req.GetXType() == network.ConnectRequest_udp {
@@ -154,49 +125,7 @@ func (c *Client) handleConnect(from, to string, req *network.ConnectRequest) {
go tn.loop()
}
func (c *Client) connectError(to, id, m string) {
var msg network.Msg
msg.From = c.cfg.ID
msg.To = to
msg.XType = network.Msg_connect_rep
msg.Payload = &network.Msg_Crep{
Crep: &network.ConnectResponse{
Id: id,
Ok: false,
Msg: m,
},
}
c.conn.WriteMessage(&msg, time.Second)
}
func (c *Client) connectOK(to, id string) {
var msg network.Msg
msg.From = c.cfg.ID
msg.To = to
msg.XType = network.Msg_connect_rep
msg.Payload = &network.Msg_Crep{
Crep: &network.ConnectResponse{
Id: id,
Ok: true,
},
}
c.conn.WriteMessage(&msg, time.Second)
}
func (c *Client) send(id, target string, data []byte) {
var msg network.Msg
msg.From = c.cfg.ID
msg.To = target
msg.XType = network.Msg_forward
msg.Payload = &network.Msg_XData{
XData: &network.Data{
Cid: id,
Data: data,
},
}
c.conn.WriteMessage(&msg, time.Second)
}
// handleDisconnect handle disconnect message from remote, this means remote connection is closed
func (c *Client) handleDisconnect(data *network.Disconnect) {
id := data.GetId()
@@ -213,6 +142,7 @@ func (c *Client) handleDisconnect(data *network.Disconnect) {
}
}
// handleData handle forward data message, write data to local connection
func (c *Client) handleData(data *network.Data) {
id := data.GetCid()
c.RLock()
@@ -224,19 +154,3 @@ func (c *Client) handleData(data *network.Data) {
}
tn.write(data.GetData())
}
func (c *Client) disconnect(id, to string) {
var msg network.Msg
msg.From = c.cfg.ID
msg.To = to
msg.XType = network.Msg_disconnect
msg.Payload = &network.Msg_XDisconnect{
XDisconnect: &network.Disconnect{
Id: id,
},
}
c.conn.WriteMessage(&msg, time.Second)
c.Lock()
delete(c.tunnels, id)
c.Unlock()
}
+100
View File
@@ -0,0 +1,100 @@
package client
import (
"natpass/code/client/global"
"natpass/code/network"
"time"
)
func (c *Client) writeHandshake() error {
var msg network.Msg
msg.XType = network.Msg_handshake
msg.From = c.cfg.ID
msg.To = "server"
msg.Payload = &network.Msg_Hsp{
Hsp: &network.HandshakePayload{
Enc: c.cfg.Enc[:],
},
}
return c.conn.WriteMessage(&msg, 5*time.Second)
}
func (c *Client) sendConnect(id string, t global.Tunnel) {
tp := network.ConnectRequest_tcp
if t.Type != "tcp" {
tp = network.ConnectRequest_udp
}
var msg network.Msg
msg.From = c.cfg.ID
msg.To = t.Target
msg.XType = network.Msg_connect_req
msg.Payload = &network.Msg_Creq{
Creq: &network.ConnectRequest{
Id: id,
Name: t.Name,
XType: tp,
Addr: t.RemoteAddr,
Port: uint32(t.RemotePort),
},
}
c.conn.WriteMessage(&msg, 5*time.Second)
}
func (c *Client) connectError(to, id, m string) {
var msg network.Msg
msg.From = c.cfg.ID
msg.To = to
msg.XType = network.Msg_connect_rep
msg.Payload = &network.Msg_Crep{
Crep: &network.ConnectResponse{
Id: id,
Ok: false,
Msg: m,
},
}
c.conn.WriteMessage(&msg, time.Second)
}
func (c *Client) connectOK(to, id string) {
var msg network.Msg
msg.From = c.cfg.ID
msg.To = to
msg.XType = network.Msg_connect_rep
msg.Payload = &network.Msg_Crep{
Crep: &network.ConnectResponse{
Id: id,
Ok: true,
},
}
c.conn.WriteMessage(&msg, time.Second)
}
func (c *Client) send(id, target string, data []byte) {
var msg network.Msg
msg.From = c.cfg.ID
msg.To = target
msg.XType = network.Msg_forward
msg.Payload = &network.Msg_XData{
XData: &network.Data{
Cid: id,
Data: data,
},
}
c.conn.WriteMessage(&msg, time.Second)
}
func (c *Client) disconnect(id, to string) {
var msg network.Msg
msg.From = c.cfg.ID
msg.To = to
msg.XType = network.Msg_disconnect
msg.Payload = &network.Msg_XDisconnect{
XDisconnect: &network.Disconnect{
Id: id,
},
}
c.conn.WriteMessage(&msg, time.Second)
c.Lock()
delete(c.tunnels, id)
c.Unlock()
}
+6
View File
@@ -50,6 +50,12 @@ func (c *client) addTunnel(id string) {
c.Unlock()
}
func (c *client) removeTunnel(id string) {
c.Lock()
delete(c.tunnels, id)
c.Unlock()
}
func (c *client) getTunnels() []string {
ret := make([]string, 0, len(c.tunnels))
c.RLock()
+18 -7
View File
@@ -18,6 +18,7 @@ type Handler struct {
tunnels map[string][2]*client // tunnel id => endpoints
}
// New create handler
func New(cfg *global.Configure) *Handler {
return &Handler{
cfg: cfg,
@@ -26,6 +27,7 @@ func New(cfg *global.Configure) *Handler {
}
}
// Handle main loop
func (h *Handler) Handle(conn net.Conn) {
c := network.NewConn(conn)
var id string
@@ -63,6 +65,7 @@ func (h *Handler) Handle(conn net.Conn) {
cli.run()
}
// readHandshake read handshake message and compare secret encoded from md5
func (h *Handler) readHandshake(c *network.Conn) (string, error) {
msg, err := c.ReadMessage(5 * time.Second)
if err != nil {
@@ -78,6 +81,7 @@ func (h *Handler) readHandshake(c *network.Conn) (string, error) {
return msg.GetFrom(), nil
}
// onMessage forward message
func (h *Handler) onMessage(msg *network.Msg) {
to := msg.GetTo()
h.RLock()
@@ -87,20 +91,21 @@ func (h *Handler) onMessage(msg *network.Msg) {
logging.Error("client %s not found", to)
return
}
h.msgFilter(msg)
h.msgHook(msg)
cli.writeMessage(msg)
}
func (h *Handler) msgFilter(msg *network.Msg) {
// msgHook hook from on message
func (h *Handler) msgHook(msg *network.Msg) {
from := msg.GetFrom()
to := msg.GetTo()
h.RLock()
fromCli := h.clients[from]
toCli := h.clients[to]
h.RUnlock()
switch msg.GetXType() {
case network.Msg_connect_rep:
if msg.GetCrep().GetOk() {
h.RLock()
fromCli := h.clients[from]
toCli := h.clients[to]
h.RUnlock()
id := msg.GetCrep().GetId()
var pair [2]*client
if fromCli != nil {
@@ -116,10 +121,16 @@ func (h *Handler) msgFilter(msg *network.Msg) {
h.Unlock()
}
case network.Msg_disconnect:
if fromCli != nil {
fromCli.removeTunnel(msg.GetXDisconnect().GetId())
}
if toCli != nil {
toCli.removeTunnel(msg.GetXDisconnect().GetId())
}
}
}
// closeAll close all tunnel from client
func (h *Handler) closeAll(cli *client) {
tunnels := cli.getTunnels()
for _, t := range tunnels {