220 lines
4.9 KiB
Go
220 lines
4.9 KiB
Go
package connect
|
||
|
||
import (
|
||
"context"
|
||
"encoding/binary"
|
||
"fmt"
|
||
"io"
|
||
"net"
|
||
"sync"
|
||
"time"
|
||
|
||
"go.uber.org/zap"
|
||
)
|
||
|
||
// FakeTCPConn FakeTCP 连接(UDP 包封装为 TCP 流)
|
||
type FakeTCPConn struct {
|
||
conn net.Conn
|
||
mu sync.Mutex
|
||
closed bool
|
||
readBuffer []byte
|
||
}
|
||
|
||
// NewFakeTCPConn 创建 FakeTCP 连接
|
||
func NewFakeTCPConn(conn net.Conn) *FakeTCPConn {
|
||
return &FakeTCPConn{
|
||
conn: conn,
|
||
readBuffer: make([]byte, 0),
|
||
}
|
||
}
|
||
|
||
// Read 读取数据(带长度前缀解析)
|
||
func (c *FakeTCPConn) Read(b []byte) (int, error) {
|
||
c.mu.Lock()
|
||
defer c.mu.Unlock()
|
||
|
||
// 如果缓冲区有数据,直接返回
|
||
if len(c.readBuffer) > 0 {
|
||
n := copy(b, c.readBuffer)
|
||
c.readBuffer = c.readBuffer[n:]
|
||
return n, nil
|
||
}
|
||
|
||
// 读取长度前缀(4 字节)
|
||
var length uint32
|
||
if err := binary.Read(c.conn, binary.BigEndian, &length); err != nil {
|
||
return 0, err
|
||
}
|
||
|
||
// 限制最大长度(防止恶意攻击)
|
||
if length > 65535 {
|
||
return 0, fmt.Errorf("packet too large: %d bytes", length)
|
||
}
|
||
|
||
// 读取实际数据
|
||
data := make([]byte, length)
|
||
if _, err := io.ReadFull(c.conn, data); err != nil {
|
||
return 0, err
|
||
}
|
||
|
||
// 返回请求的数据
|
||
n := copy(b, data)
|
||
if n < len(data) {
|
||
// 剩余数据存入缓冲区
|
||
c.readBuffer = data[n:]
|
||
}
|
||
|
||
return n, nil
|
||
}
|
||
|
||
// Write 写入数据(添加 4 字节长度前缀)
|
||
func (c *FakeTCPConn) Write(b []byte) (int, error) {
|
||
c.mu.Lock()
|
||
defer c.mu.Unlock()
|
||
|
||
if c.closed {
|
||
return 0, fmt.Errorf("connection closed")
|
||
}
|
||
|
||
// 写入长度前缀
|
||
length := uint32(len(b))
|
||
if err := binary.Write(c.conn, binary.BigEndian, length); err != nil {
|
||
return 0, err
|
||
}
|
||
|
||
// 写入实际数据
|
||
n, err := c.conn.Write(b)
|
||
return n, err
|
||
}
|
||
|
||
// Close 关闭连接
|
||
func (c *FakeTCPConn) Close() error {
|
||
c.mu.Lock()
|
||
defer c.mu.Unlock()
|
||
|
||
c.closed = true
|
||
return c.conn.Close()
|
||
}
|
||
|
||
// LocalAddr 本地地址
|
||
func (c *FakeTCPConn) LocalAddr() net.Addr {
|
||
return c.conn.LocalAddr()
|
||
}
|
||
|
||
// RemoteAddr 远程地址
|
||
func (c *FakeTCPConn) RemoteAddr() net.Addr {
|
||
return c.conn.RemoteAddr()
|
||
}
|
||
|
||
// SetDeadline 设置截止时间
|
||
func (c *FakeTCPConn) SetDeadline(t time.Time) error {
|
||
return c.conn.SetDeadline(t)
|
||
}
|
||
|
||
// SetReadDeadline 设置读截止时间
|
||
func (c *FakeTCPConn) SetReadDeadline(t time.Time) error {
|
||
return c.conn.SetReadDeadline(t)
|
||
}
|
||
|
||
// SetWriteDeadline 设置写截止时间
|
||
func (c *FakeTCPConn) SetWriteDeadline(t time.Time) error {
|
||
return c.conn.SetWriteDeadline(t)
|
||
}
|
||
|
||
// DialFakeTCP 拨号 FakeTCP 连接
|
||
func DialFakeTCP(ctx context.Context, network, addr string, logger *zap.Logger) (net.Conn, error) {
|
||
logger.Debug("dialing FakeTCP", zap.String("addr", addr))
|
||
|
||
// 建立 TCP 连接
|
||
conn, err := (&net.Dialer{}).DialContext(ctx, network, addr)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("failed to dial TCP: %w", err)
|
||
}
|
||
|
||
// 包装为 FakeTCP 连接
|
||
return NewFakeTCPConn(conn), nil
|
||
}
|
||
|
||
// ListenFakeTCP 监听 FakeTCP 端口
|
||
func ListenFakeTCP(network, addr string, logger *zap.Logger) (net.Listener, error) {
|
||
logger.Info("listening FakeTCP", zap.String("addr", addr))
|
||
|
||
// 监听 TCP 端口
|
||
listener, err := net.Listen(network, addr)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("failed to listen TCP: %w", err)
|
||
}
|
||
|
||
return &fakeTCPListener{
|
||
Listener: listener,
|
||
logger: logger,
|
||
}, nil
|
||
}
|
||
|
||
// fakeTCPListener FakeTCP 监听器
|
||
type fakeTCPListener struct {
|
||
net.Listener
|
||
logger *zap.Logger
|
||
}
|
||
|
||
// Accept 接受连接并包装为 FakeTCPConn
|
||
func (l *fakeTCPListener) Accept() (net.Conn, error) {
|
||
conn, err := l.Listener.Accept()
|
||
if err != nil {
|
||
return nil, err
|
||
}
|
||
|
||
l.logger.Debug("accepted FakeTCP connection", zap.String("addr", conn.RemoteAddr().String()))
|
||
return NewFakeTCPConn(conn), nil
|
||
}
|
||
|
||
// FakeTCPFactory FakeTCP 传输工厂
|
||
type FakeTCPFactory struct {
|
||
logger *zap.Logger
|
||
}
|
||
|
||
// NewFakeTCPFactory 创建 FakeTCP 工厂
|
||
func NewFakeTCPFactory(logger *zap.Logger) *FakeTCPFactory {
|
||
return &FakeTCPFactory{
|
||
logger: logger,
|
||
}
|
||
}
|
||
|
||
// Layer 返回传输层类型
|
||
func (f *FakeTCPFactory) Layer() Layer {
|
||
return LayerFakeTCP
|
||
}
|
||
|
||
// Name 返回名称
|
||
func (f *FakeTCPFactory) Name() string {
|
||
return "FakeTCP"
|
||
}
|
||
|
||
// Dial 建立 FakeTCP 连接
|
||
func (f *FakeTCPFactory) Dial(ctx context.Context, config *DialConfig) (net.Conn, error) {
|
||
f.logger.Info("开始建立 FakeTCP 连接",
|
||
zap.String("peer_id", config.PeerID))
|
||
|
||
// 1. 解析对端地址(PeerID 格式应为 "ip:port")
|
||
if config.PeerID == "" {
|
||
return nil, fmt.Errorf("PeerID 为空")
|
||
}
|
||
|
||
// 2. 建立 TCP 连接
|
||
dialer := &net.Dialer{Timeout: config.Timeout}
|
||
conn, err := dialer.DialContext(ctx, "tcp", config.PeerID)
|
||
if err != nil {
|
||
return nil, fmt.Errorf("TCP 连接失败:%w", err)
|
||
}
|
||
|
||
// 3. 包装为 FakeTCP 连接(UDP 包封装为 TCP 流)
|
||
fakeConn := NewFakeTCPConn(conn)
|
||
|
||
f.logger.Info("FakeTCP 连接建立成功",
|
||
zap.String("peer_id", config.PeerID),
|
||
zap.String("local_addr", conn.LocalAddr().String()),
|
||
zap.String("remote_addr", conn.RemoteAddr().String()))
|
||
|
||
return fakeConn, nil
|
||
}
|