From e643d0cd7a0a2385e51be0a72a753b7dcdb78748 Mon Sep 17 00:00:00 2001 From: lwch Date: Thu, 12 Aug 2021 17:39:10 +0800 Subject: [PATCH] =?UTF-8?q?1.=20=E5=A2=9E=E5=8A=A0=E4=BB=A3=E7=A0=81?= =?UTF-8?q?=E6=B3=A8=E9=87=8A=202.=20hook=E4=BA=86disconnect=E6=B6=88?= =?UTF-8?q?=E6=81=AF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- code/client/client/client.go | 100 +++------------------------------ code/client/client/send.go | 100 +++++++++++++++++++++++++++++++++ code/server/handler/client.go | 6 ++ code/server/handler/handler.go | 25 ++++++--- 4 files changed, 131 insertions(+), 100 deletions(-) create mode 100644 code/client/client/send.go diff --git a/code/client/client/client.go b/code/client/client/client.go index 2151823..90468fb 100644 --- a/code/client/client/client.go +++ b/code/client/client/client.go @@ -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() -} diff --git a/code/client/client/send.go b/code/client/client/send.go new file mode 100644 index 0000000..865be92 --- /dev/null +++ b/code/client/client/send.go @@ -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() +} diff --git a/code/server/handler/client.go b/code/server/handler/client.go index 1d89773..57ab000 100644 --- a/code/server/handler/client.go +++ b/code/server/handler/client.go @@ -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() diff --git a/code/server/handler/handler.go b/code/server/handler/handler.go index dc092a2..14f41de 100644 --- a/code/server/handler/handler.go +++ b/code/server/handler/handler.go @@ -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 {