mirror of
https://github.com/lwch/natpass.git
synced 2024-04-21 12:41:54 +00:00
131 lines
3.0 KiB
Go
131 lines
3.0 KiB
Go
package handler
|
|
|
|
import (
|
|
"context"
|
|
"natpass/code/client/global"
|
|
"natpass/code/client/tunnel"
|
|
"natpass/code/network"
|
|
"net"
|
|
"sync"
|
|
|
|
"github.com/lwch/logging"
|
|
"github.com/lwch/runtime"
|
|
)
|
|
|
|
type Handler struct {
|
|
sync.RWMutex
|
|
ctx context.Context
|
|
ctl network.Natpass_ControlClient
|
|
fwd network.Natpass_ForwardClient
|
|
mgr *tunnel.Mgr
|
|
|
|
// runtime
|
|
conns map[string]net.Conn
|
|
}
|
|
|
|
func New(ctx context.Context,
|
|
ctl network.Natpass_ControlClient, fwd network.Natpass_ForwardClient) *Handler {
|
|
return &Handler{
|
|
ctx: ctx,
|
|
ctl: ctl,
|
|
fwd: fwd,
|
|
mgr: tunnel.NewMgr(),
|
|
}
|
|
}
|
|
|
|
func (h *Handler) CreateTunnels(id string, tunnels []global.Tunnel) {
|
|
cfgs := make(map[string]global.Tunnel, len(tunnels)) // name => config
|
|
h.conns = make(map[string]net.Conn) // local_cid => connection
|
|
for _, t := range tunnels {
|
|
l, err := tunnel.NewListener(t.Type, t.Name, t.LocalAddr, t.LocalPort)
|
|
runtime.Assert(err)
|
|
cfgs[t.Name] = t
|
|
go l.Loop(func(c net.Conn, name string) {
|
|
id := createTunnel(h.ctx, cfgs[name], h.ctl, id)
|
|
h.Lock()
|
|
h.conns[id] = c
|
|
h.Unlock()
|
|
})
|
|
}
|
|
}
|
|
|
|
func (h *Handler) HandleControl() {
|
|
for {
|
|
data, err := h.ctl.Recv()
|
|
if err != nil {
|
|
logging.Error("handler: %v", err)
|
|
return
|
|
}
|
|
switch data.GetXType() {
|
|
case network.ControlData_connect_req:
|
|
h.handleConnectReq(data)
|
|
case network.ControlData_connect_rep:
|
|
h.handleConnectRep(data)
|
|
}
|
|
}
|
|
}
|
|
|
|
func (h *Handler) handleConnectReq(data *network.ControlData) {
|
|
req := data.GetCreq()
|
|
cid, err := runtime.UUID(cidLength)
|
|
rep := &network.ConnectResponse{
|
|
Name: req.GetName(),
|
|
RemoteCid: req.GetCid(),
|
|
}
|
|
if err != nil {
|
|
logging.Error("generate channel_id failed, err=%v", err)
|
|
rep.Ok = false
|
|
rep.Msg = err.Error()
|
|
connectResponse(h.ctl, rep, data.GetTo(), data.GetFrom())
|
|
return
|
|
}
|
|
t := "tcp"
|
|
if req.GetXType() == network.ConnectRequest_udp {
|
|
t = "udp"
|
|
}
|
|
tn, err := tunnel.NewConnect(req.GetName(), data.GetTo(), cid,
|
|
data.GetFrom(), req.GetCid(),
|
|
t, req.GetAddr(), req.GetPort())
|
|
if err != nil {
|
|
logging.Error("new connect failed: %v", err)
|
|
rep.Ok = false
|
|
rep.Msg = err.Error()
|
|
connectResponse(h.ctl, rep, data.GetTo(), data.GetFrom())
|
|
return
|
|
}
|
|
h.mgr.Add(tn)
|
|
rep.Ok = true
|
|
rep.LocalCid = cid
|
|
connectResponse(h.ctl, rep, data.GetTo(), data.GetFrom())
|
|
go tn.ForwardLocal(h.fwd)
|
|
}
|
|
|
|
func (h *Handler) handleConnectRep(data *network.ControlData) {
|
|
rep := data.GetCrep()
|
|
if !rep.GetOk() {
|
|
logging.Error("connect %s failed, err=%s", rep.GetName(), rep.GetMsg())
|
|
return
|
|
}
|
|
tn, err := tunnel.NewListen(rep.GetName(), data.GetTo(), rep.GetRemoteCid(),
|
|
data.GetFrom(), rep.GetLocalCid(), h.conns[rep.GetRemoteCid()])
|
|
runtime.Assert(err)
|
|
h.mgr.Add(tn)
|
|
go tn.ForwardLocal(h.fwd)
|
|
}
|
|
|
|
func (h *Handler) HandleForward() {
|
|
for {
|
|
data, err := h.fwd.Recv()
|
|
if err != nil {
|
|
logging.Error("read data from remote failed, err=%v", err)
|
|
continue
|
|
}
|
|
tn := h.mgr.Find(data.GetCid())
|
|
if tn == nil {
|
|
logging.Error("tunnel not found, from=%s, to=%s", data.GetFrom(), data.GetTo())
|
|
continue
|
|
}
|
|
tn.WriteLocal(data.GetData())
|
|
}
|
|
}
|