Files
Meshray-Manager/core/connect/fake_tcp.go
T
2026-06-30 15:14:37 +08:00

220 lines
4.9 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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
}