diff --git a/bannet/client.go b/bannet/client.go deleted file mode 100644 index 87081cd..0000000 --- a/bannet/client.go +++ /dev/null @@ -1,176 +0,0 @@ -package bannet - -import ( - "encoding/binary" - "fmt" - "io" - "net" - "time" - - "github.com/NeverENG/BanDB/pkg/proto" - "github.com/NeverENG/BanDB/pkg/utils" -) - -// Client 是可复用的 BanNet TCP 客户端:按二进制 TLV 协议对一个 BanNet 服务端发起 -// PUT/GET/DELETE。命令行交互客户端在 cmd `client/`(package main,不可被库导入), -// 故这里在 banNet 包内提供一份库版本,供跨节点转发(分片集群)等复用。 -// -// 协议(见 pkg/proto/codes.go):请求/响应帧为 [dataLen u32 LE][msgIDLen u16 LE][msgID][data]; -// PUT/DEL 响应 data = [statusLen u8][status];GET 响应追加 [valueLen u32 LE][value]。 -type Client struct { - addr string - conn net.Conn - timeout time.Duration -} - -// NewClient 创建一个连往 addr 的客户端(尚未拨号)。timeout<=0 取 5s。 -func NewClient(addr string, timeout time.Duration) *Client { - if timeout <= 0 { - timeout = 5 * time.Second - } - return &Client{addr: addr, timeout: timeout} -} - -// Connect 建立 TCP 连接。 -func (c *Client) Connect() error { - conn, err := net.DialTimeout("tcp", c.addr, c.timeout) - if err != nil { - return fmt.Errorf("banNet client dial %s: %w", c.addr, err) - } - c.conn = conn - return nil -} - -// Close 关闭连接。 -func (c *Client) Close() error { - if c.conn != nil { - err := c.conn.Close() - c.conn = nil - return err - } - return nil -} - -// Put 发送 PUT 并等待成功响应。 -func (c *Client) Put(key, value []byte) error { - msg := utils.NewMessage(proto.MsgPut, key, value) - payload, err := c.roundTrip(msg) - if err != nil { - return err - } - status, _, err := parseStatus(payload) - if err != nil { - return err - } - if status != proto.StatusOK { - return fmt.Errorf("banNet client PUT: server status %q", status) - } - return nil -} - -// Get 发送 GET,返回 value 与是否命中。未命中返回 (nil, false, nil)。 -func (c *Client) Get(key []byte) ([]byte, bool, error) { - keyLen := make([]byte, 4) - binary.LittleEndian.PutUint32(keyLen, uint32(len(key))) - msg := utils.NewMessage2(proto.MsgGet, utils.ByteBuilder(keyLen, key)) - payload, err := c.roundTrip(msg) - if err != nil { - return nil, false, err - } - status, rest, err := parseStatus(payload) - if err != nil { - return nil, false, err - } - if status != proto.StatusOK { - return nil, false, nil // 未命中/远端错误:按未命中处理 - } - if len(rest) < 4 { - return nil, false, fmt.Errorf("banNet client GET: truncated response") - } - valueLen := binary.LittleEndian.Uint32(rest[:4]) - if len(rest) < 4+int(valueLen) { - return nil, false, fmt.Errorf("banNet client GET: incomplete value") - } - return rest[4 : 4+valueLen], true, nil -} - -// Delete 发送 DELETE 并等待成功响应。 -func (c *Client) Delete(key []byte) error { - keyLen := make([]byte, 4) - binary.LittleEndian.PutUint32(keyLen, uint32(len(key))) - msg := utils.NewMessage2(proto.MsgDelete, utils.ByteBuilder(keyLen, key)) - payload, err := c.roundTrip(msg) - if err != nil { - return err - } - status, _, err := parseStatus(payload) - if err != nil { - return err - } - if status != proto.StatusOK { - return fmt.Errorf("banNet client DELETE: server status %q", status) - } - return nil -} - -// roundTrip 打包发送一条消息并读回一条响应的 data 负载。 -func (c *Client) roundTrip(msg *utils.Message) ([]byte, error) { - if c.conn == nil { - return nil, fmt.Errorf("banNet client: not connected") - } - dp := NewDataPack() - packet, err := dp.Pack(msg) - if err != nil { - return nil, fmt.Errorf("banNet client pack: %w", err) - } - c.conn.SetWriteDeadline(time.Now().Add(c.timeout)) - if _, err := c.conn.Write(packet); err != nil { - return nil, err - } - c.conn.SetReadDeadline(time.Now().Add(c.timeout)) - return c.readResponse() -} - -// readResponse 读取一条响应帧,返回其 data 负载。 -func (c *Client) readResponse() ([]byte, error) { - dp := NewDataPack() - header := make([]byte, dp.HeadLen()) - if _, err := io.ReadFull(c.conn, header); err != nil { - return nil, fmt.Errorf("banNet client read header: %w", err) - } - tempMsg, err := dp.UnPack(header) - if err != nil { - return nil, fmt.Errorf("banNet client unpack: %w", err) - } - mImpl, ok := tempMsg.(*Message) - if !ok { - return nil, fmt.Errorf("banNet client: unexpected message type") - } - if mImpl.IDLen > 0 { - idBuf := make([]byte, mImpl.IDLen) - if _, err := io.ReadFull(c.conn, idBuf); err != nil { - return nil, fmt.Errorf("banNet client read msgID: %w", err) - } - } - dataLen := tempMsg.MsgLen() - if dataLen == 0 { - return nil, nil - } - data := make([]byte, dataLen) - if _, err := io.ReadFull(c.conn, data); err != nil { - return nil, fmt.Errorf("banNet client read data: %w", err) - } - return data, nil -} - -// parseStatus 从响应 data 头部解析 [statusLen u8][status],返回 status 与剩余字节。 -func parseStatus(payload []byte) (string, []byte, error) { - if len(payload) < 1 { - return "", nil, fmt.Errorf("banNet client: empty payload") - } - statusLen := int(payload[0]) - if len(payload) < 1+statusLen { - return "", nil, fmt.Errorf("banNet client: truncated status") - } - return string(payload[1 : 1+statusLen]), payload[1+statusLen:], nil -} diff --git a/bannet/connection.go b/bannet/connection.go index 58527ca..9767405 100644 --- a/bannet/connection.go +++ b/bannet/connection.go @@ -77,9 +77,8 @@ func (c *Connection) StartReader() { } // 头部之后, 先按 IDLen 读取 msgID 字符串 - mImpl := msg.(*Message) - if mImpl.IDLen > 0 { - idBuf := make([]byte, mImpl.IDLen) + if msg.IDLen > 0 { + idBuf := make([]byte, msg.IDLen) if _, err := io.ReadFull(reader, idBuf); err != nil { slog.Error("conn read msgID failed", "connID", c.ConnID, "error", err) return diff --git a/bannet/datapack.go b/bannet/datapack.go index c0ee886..8e4625b 100644 --- a/bannet/datapack.go +++ b/bannet/datapack.go @@ -25,7 +25,7 @@ func (dp *DataPack) HeadLen() uint32 { // Pack 编码一帧。按最终长度一次性分配并直接写入定长头部,不经 bytes.Buffer 的增量扩容, // 也不经 binary.Write 的反射路径——Pack 位于每个响应的必经路径上。 -func (dp *DataPack) Pack(msg Frame) ([]byte, error) { +func (dp *DataPack) Pack(msg *Message) ([]byte, error) { id := msg.MsgID() if len(id) > 0xFFFF { return nil, fmt.Errorf("msgID too long: %d", len(id)) @@ -43,7 +43,7 @@ func (dp *DataPack) Pack(msg Frame) ([]byte, error) { // UnPack 只解析定长头部 (6 字节), 返回带 DataLen 与 IDLen 的占位 Message; // 调用方拿到 IDLen 后, 还需要从连接读取 IDLen+DataLen 字节填充 Id 与 Data。 -func (dp *DataPack) UnPack(data []byte) (Frame, error) { +func (dp *DataPack) UnPack(data []byte) (*Message, error) { if len(data) < int(dp.HeadLen()) { return nil, errors.New("head too short") } diff --git a/bannet/interfaces.go b/bannet/interfaces.go index be8f826..3008d6c 100644 --- a/bannet/interfaces.go +++ b/bannet/interfaces.go @@ -15,8 +15,8 @@ type ConnRegistry interface { type Codec interface { HeadLen() uint32 - Pack(msg Frame) ([]byte, error) - UnPack([]byte) (Frame, error) + Pack(msg *Message) ([]byte, error) + UnPack([]byte) (*Message, error) } type Dispatcher interface { @@ -63,12 +63,3 @@ type Conn interface { Property(key string) any RemoveProperty(key string) } - -type Frame interface { - MsgID() string - Payload() []byte - MsgLen() uint32 - SetMsgLen(uint32) - SetData([]byte) - SetMsgID(string) -} diff --git a/bannet/message.go b/bannet/message.go index 8f88b0c..06befb8 100644 --- a/bannet/message.go +++ b/bannet/message.go @@ -10,8 +10,6 @@ type Message struct { Data []byte } -var _ Frame = &Message{} - func NewMessage(id string, data []byte) *Message { return &Message{ ID: id, diff --git a/bannet/request.go b/bannet/request.go index e5ab1d4..98fb1c9 100644 --- a/bannet/request.go +++ b/bannet/request.go @@ -1,13 +1,13 @@ package bannet type request struct { - msg Frame + msg *Message conn Conn } var _ Request = &request{} -func newRequest(msg Frame, conn Conn) *request { +func newRequest(msg *Message, conn Conn) *request { return &request{ msg: msg, conn: conn, diff --git a/bannet/wire_scan_test.go b/bannet/wire_scan_test.go index c751c3b..3474558 100644 --- a/bannet/wire_scan_test.go +++ b/bannet/wire_scan_test.go @@ -41,7 +41,7 @@ func TestScanResponseSurvivesWire(t *testing.T) { if err != nil { t.Fatalf("UnPack 拒绝了大响应(MaxPackageSize 太小?): %v", err) } - m := tempMsg.(*bannet.Message) + m := tempMsg off := headLen + int(m.IDLen) data := packet[off : off+int(tempMsg.MsgLen())] diff --git a/cluster/peerpool.go b/cluster/peerpool.go index df9572f..9383445 100644 --- a/cluster/peerpool.go +++ b/cluster/peerpool.go @@ -1,98 +1,111 @@ package cluster import ( + "context" + "errors" "sync" "time" - "github.com/NeverENG/BanDB/bannet" + bandb "github.com/NeverENG/BanDB/client" ) -// PeerPool 维护到各节点的 BanNet 客户端连接(懒建、缓存),供分片转发复用。 +// peerMaxRetries 是转发到属主节点时的重试次数,取 0(不重试)。 // -// 并发正确性:单条 BanNet 连接是请求-响应式,不能被多个 goroutine 交错读写(否则帧 -// 错位)。故每个 peer 用一把锁串行化其连接上的调用;出错则丢弃连接、下次重连。 -// 需要更高并发时可扩为每 peer 一个连接池,此处先保正确。 +// 转发发生在「客户端 → 入口节点 → 属主节点」链路的第二跳,而入口侧的客户端 SDK 自身 +// 已带重试。若此处再重试,过载时两级会相乘放大请求量,正好在最不该加压的时刻加压。 +// 网络瞬时故障由客户端那一层的重试覆盖即可。 +const peerMaxRetries = -1 // client.Options 中负值表示不重试 + +// PeerPool 维护到各 peer 节点的客户端,供分片转发复用。 +// +// 每个 peer 一个 client.Client:BanNet 是请求-响应协议,一条连接必须收到响应才能发下 +// 一帧,故对同一 peer 的并发转发由 SDK 内部的连接池以多条连接承担,而非串行排队。 type PeerPool struct { mu sync.Mutex timeout time.Duration - peers map[string]*peerConn + peers map[string]*bandb.Client } -type peerConn struct { - mu sync.Mutex - addr string - client *bannet.Client -} - -// NewPeerPool 创建一个转发连接池。timeout 为每次拨号/读写的超时。 +// NewPeerPool 创建一个转发连接池。timeout 为每次拨号/请求的超时。 func NewPeerPool(timeout time.Duration) *PeerPool { - return &PeerPool{timeout: timeout, peers: map[string]*peerConn{}} + return &PeerPool{timeout: timeout, peers: map[string]*bandb.Client{}} } -// conn 取(或创建)某 peer 的连接槽。 -func (p *PeerPool) conn(addr string) *peerConn { +// client 取(或惰性创建)到该 peer 的客户端。 +func (p *PeerPool) client(addr string) (*bandb.Client, error) { p.mu.Lock() defer p.mu.Unlock() - pc, ok := p.peers[addr] - if !ok { - pc = &peerConn{addr: addr} - p.peers[addr] = pc + if c, ok := p.peers[addr]; ok { + return c, nil + } + c, err := bandb.New(bandb.Options{ + Addrs: []string{addr}, + DialTimeout: p.timeout, + RequestTimeout: p.timeout, + MaxRetries: peerMaxRetries, + }) + if err != nil { + return nil, err } - return pc + p.peers[addr] = c + return c, nil } -// withClient 在 peer 连接上串行执行 fn;懒建连接,fn 出错则丢弃连接以便下次重连。 -func (p *PeerPool) withClient(addr string, fn func(*bannet.Client) error) error { - pc := p.conn(addr) - pc.mu.Lock() - defer pc.mu.Unlock() - if pc.client == nil { - c := bannet.NewClient(addr, p.timeout) - if err := c.Connect(); err != nil { - return err - } - pc.client = c - } - if err := fn(pc.client); err != nil { - _ = pc.client.Close() - pc.client = nil // 下次重连 - return err - } - return nil +// ctx 返回一个受 timeout 约束的上下文。 +func (p *PeerPool) ctx() (context.Context, context.CancelFunc) { + return context.WithTimeout(context.Background(), p.timeout) } // Put 转发 PUT 到 addr 节点。 func (p *PeerPool) Put(addr string, key, value []byte) error { - return p.withClient(addr, func(c *bannet.Client) error { return c.Put(key, value) }) + c, err := p.client(addr) + if err != nil { + return err + } + ctx, cancel := p.ctx() + defer cancel() + return c.Put(ctx, key, value) } // Get 转发 GET 到 addr 节点,返回 value 与是否命中。 +// +// 「未命中」与「属主节点出错」严格区分:仅 ErrKeyNotFound 记为未命中,其余错误原样 +// 上抛。二者混同会让上游把远端故障当成「这个 key 不存在」,从而返回错误的空结果。 func (p *PeerPool) Get(addr string, key []byte) ([]byte, bool, error) { - var value []byte - var found bool - err := p.withClient(addr, func(c *bannet.Client) error { - v, ok, err := c.Get(key) - value, found = v, ok - return err - }) - return value, found, err + c, err := p.client(addr) + if err != nil { + return nil, false, err + } + ctx, cancel := p.ctx() + defer cancel() + + value, err := c.Get(ctx, key) + switch { + case errors.Is(err, bandb.ErrKeyNotFound): + return nil, false, nil + case err != nil: + return nil, false, err + } + return value, true, nil } // Delete 转发 DELETE 到 addr 节点。 func (p *PeerPool) Delete(addr string, key []byte) error { - return p.withClient(addr, func(c *bannet.Client) error { return c.Delete(key) }) + c, err := p.client(addr) + if err != nil { + return err + } + ctx, cancel := p.ctx() + defer cancel() + return c.Delete(ctx, key) } -// Close 关闭所有缓存连接。 +// Close 关闭所有缓存的客户端。 func (p *PeerPool) Close() { p.mu.Lock() defer p.mu.Unlock() - for _, pc := range p.peers { - pc.mu.Lock() - if pc.client != nil { - _ = pc.client.Close() - pc.client = nil - } - pc.mu.Unlock() + for addr, c := range p.peers { + _ = c.Close() + delete(p.peers, addr) } } diff --git a/cmd/ban-bench/runner.go b/cmd/ban-bench/runner.go index 732b1ff..18500ac 100644 --- a/cmd/ban-bench/runner.go +++ b/cmd/ban-bench/runner.go @@ -243,7 +243,7 @@ func dial(addr string) (*net.TCPConn, error) { } func send(conn *net.TCPConn, msgID string, data []byte) error { - msg := utils.NewMessage2(msgID, data) + msg := bannet.NewMessage(msgID, data) dp := bannet.NewDataPack() packed, err := dp.Pack(msg) if err != nil { @@ -267,9 +267,8 @@ func recv(conn *net.TCPConn) ([]byte, error) { return nil, err } - mImpl := tempMsg.(*bannet.Message) - if mImpl.IDLen > 0 { - idBuf := make([]byte, mImpl.IDLen) + if tempMsg.IDLen > 0 { + idBuf := make([]byte, tempMsg.IDLen) if _, err := io.ReadFull(conn, idBuf); err != nil { return nil, err } diff --git a/pkg/utils/message.go b/pkg/utils/message.go deleted file mode 100644 index dbea09a..0000000 --- a/pkg/utils/message.go +++ /dev/null @@ -1,59 +0,0 @@ -package utils - -import "encoding/binary" - -type Message struct { - Id string - - DataLen uint32 - Data []byte -} - -func NewKVData(key []byte, value []byte) []byte { - Keylen := make([]byte, 4) - valuelen := make([]byte, 4) - - binary.LittleEndian.PutUint32(Keylen, uint32(len(key))) - binary.LittleEndian.PutUint32(valuelen, uint32(len(value))) - - return ByteBuilder(Keylen, valuelen, key, value) -} - -func NewMessage(id string, key []byte, value []byte) *Message { - data := NewKVData(key, value) - return &Message{ - Id: id, - DataLen: uint32(len(data)), - Data: data, - } -} - -func NewMessage2(id string, data []byte) *Message { - return &Message{ - Id: id, - DataLen: uint32(len(data)), - Data: data, - } -} - -func (m *Message) MsgID() string { - return m.Id -} -func (m *Message) MsgLen() uint32 { - return m.DataLen -} -func (m *Message) Payload() []byte { - return m.Data -} - -func (m *Message) SetMsgID(id string) { - m.Id = id -} - -func (m *Message) SetData(data []byte) { - m.Data = data -} - -func (m *Message) SetMsgLen(id uint32) { - m.DataLen = id -} diff --git a/service/shard_routing_integration_test.go b/service/shard_routing_integration_test.go index 396041f..03509a8 100644 --- a/service/shard_routing_integration_test.go +++ b/service/shard_routing_integration_test.go @@ -1,6 +1,7 @@ package service import ( + "context" "errors" "fmt" "net" @@ -10,6 +11,7 @@ import ( "time" "github.com/NeverENG/BanDB/bannet" + bandb "github.com/NeverENG/BanDB/client" "github.com/NeverENG/BanDB/cluster" "github.com/NeverENG/BanDB/config" "github.com/NeverENG/BanDB/pkg/predicate" @@ -108,12 +110,9 @@ func TestShardRouting_MultiNode(t *testing.T) { value := []byte("payload-v1") // 经入口节点写入 → 应被转发到属主。 - c := bannet.NewClient(entry, 2*time.Second) - if err := c.Connect(); err != nil { - t.Fatal(err) - } - defer c.Close() - if err := c.Put(key, value); err != nil { + ctx := context.Background() + c := newTestClient(t, entry) + if err := c.Put(ctx, key, value); err != nil { t.Fatalf("put via entry failed: %v", err) } @@ -126,9 +125,9 @@ func TestShardRouting_MultiNode(t *testing.T) { } // 从入口读(转发读)→ 命中。 - got, found, err := c.Get(key) - if err != nil || !found { - t.Fatalf("get via entry: found=%v err=%v", found, err) + got, err := c.Get(ctx, key) + if err != nil { + t.Fatalf("get via entry: %v", err) } if string(got) != string(value) { t.Fatalf("get via entry = %q, want %q", got, value) @@ -136,18 +135,14 @@ func TestShardRouting_MultiNode(t *testing.T) { // 从第三个节点(既非入口也非属主)读 → 仍应转发到属主命中。 third := indexOfOther(peers, 0, ownerIdx) - c3 := bannet.NewClient(peers[third], 2*time.Second) - if err := c3.Connect(); err != nil { - t.Fatal(err) - } - defer c3.Close() - got3, found3, err := c3.Get(key) - if err != nil || !found3 || string(got3) != string(value) { - t.Fatalf("get via third node: got=%q found=%v err=%v", got3, found3, err) + c3 := newTestClient(t, peers[third]) + got3, err := c3.Get(ctx, key) + if err != nil || string(got3) != string(value) { + t.Fatalf("get via third node: got=%q err=%v", got3, err) } // 删除经入口转发 → 属主 store 不再有该 key。 - if err := c.Delete(key); err != nil { + if err := c.Delete(ctx, key); err != nil { t.Fatalf("delete via entry failed: %v", err) } if stores[ownerIdx].has(string(key)) { @@ -173,3 +168,18 @@ func indexOfOther(ss []string, a, b int) int { } return -1 } + +// newTestClient 建一个指向 addr 的 SDK 客户端,随测试结束关闭。 +func newTestClient(t *testing.T, addr string) *bandb.Client { + t.Helper() + c, err := bandb.New(bandb.Options{ + Addrs: []string{addr}, + DialTimeout: 2 * time.Second, + RequestTimeout: 2 * time.Second, + }) + if err != nil { + t.Fatalf("new client: %v", err) + } + t.Cleanup(func() { c.Close() }) + return c +}