Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
17 changes: 15 additions & 2 deletions bannet/connection.go
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,10 @@ type Connection struct {

property map[string]any
propertyLock sync.RWMutex

// 构造时从全局配置快照的两项策略,避免在每帧读取路径上访问可变全局状态。
maxPackageSize uint32
useWorkerPool bool
}

func NewConnection(conn *net.TCPConn, connID uint32, handle Dispatcher, server *Server) *Connection {
Expand All @@ -45,6 +49,9 @@ func NewConnection(conn *net.TCPConn, connID uint32, handle Dispatcher, server *
msgChan: make(chan []byte, 10), // 高优通道加小缓冲,避免硬阻塞
msgBuffChan: make(chan []byte, config.G.MaxMsgChanLen),
property: make(map[string]any), // 必须初始化,否则 SetProperty 写 nil map 会 panic

maxPackageSize: config.G.MaxPackageSize,
useWorkerPool: config.G.WorkerPoolSize > 0,
}
c.TCPServer.Conns().Add(c)
return c
Expand Down Expand Up @@ -75,6 +82,12 @@ func (c *Connection) StartReader() {
slog.Error("conn unpack header failed", "connID", c.ConnID, "error", err)
return
}
// 帧长上限在此执行:读取负载之前拒绝超限帧,避免按对端声称的长度分配内存。
if c.maxPackageSize > 0 && msg.MsgLen() > c.maxPackageSize {
slog.Warn("conn frame exceeds max package size",
"connID", c.ConnID, "dataLen", msg.MsgLen(), "max", c.maxPackageSize)
return
}

// 头部之后, 先按 IDLen 读取 msgID 字符串
if msg.IDLen > 0 {
Expand All @@ -97,8 +110,8 @@ func (c *Connection) StartReader() {
}
msg.SetData(data)
req := newRequest(msg, c)
// 根据有没有启动 WorkPool 选择不同的结果
if config.G.WorkerPoolSize > 0 {
// 根据有没有启动 worker 池选择投递方式
if c.useWorkerPool {
c.MsgHandle.SendMsgToTaskQueue(req)
} else {
go c.MsgHandle.DoMsgHandle(req)
Expand Down
15 changes: 7 additions & 8 deletions bannet/datapack.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,6 @@ import (
"encoding/binary"
"errors"
"fmt"
"log/slog"

"github.com/NeverENG/BanDB/config"
)

// 报文格式:
Expand Down Expand Up @@ -41,7 +38,13 @@ func (dp *DataPack) Pack(msg *Message) ([]byte, error) {
return buf, nil
}

// UnPack 只解析定长头部 (6 字节), 返回带 DataLen 与 IDLen 的占位 Message;
// UnPack 只解析定长头部。
//
// 它不校验帧长上限:那是策略而非编解码,由连接侧在读取负载前执行(见 Connection.
// StartReader)。此前该校验在此处每帧读两次全局配置——把策略留在边界,解码器才能保持
// 无状态、不依赖全局。
//
// 原注释:只解析定长头部 (6 字节), 返回带 DataLen 与 IDLen 的占位 Message;
// 调用方拿到 IDLen 后, 还需要从连接读取 IDLen+DataLen 字节填充 Id 与 Data。
func (dp *DataPack) UnPack(data []byte) (*Message, error) {
if len(data) < int(dp.HeadLen()) {
Expand All @@ -53,10 +56,6 @@ func (dp *DataPack) UnPack(data []byte) (*Message, error) {
msg.DataLen = binary.LittleEndian.Uint32(data[0:4])
idLen := binary.LittleEndian.Uint16(data[4:6])

if config.G.MaxPackageSize > 0 && msg.DataLen > config.G.MaxPackageSize {
slog.Warn("banNet frame exceeds max package size", "dataLen", msg.DataLen, "max", config.G.MaxPackageSize)
return nil, errors.New("data too large")
}
// 借用 Id 暂存 IDLen 信息: 调用方先从 MsgID() 拿不到东西, 通过头部之后另读 IDLen 字节填回。
// 这里用 SetMsgLen 仅保留 DataLen 不冲突, IDLen 通过返回的 Message.IDLen 提供。
msg.IDLen = idLen
Expand Down
91 changes: 91 additions & 0 deletions bannet/oversized_frame_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,91 @@
package bannet_test

import (
"encoding/binary"
"net"
"strconv"
"testing"
"time"

"github.com/NeverENG/BanDB/bannet"
"github.com/NeverENG/BanDB/config"
"github.com/NeverENG/BanDB/proto"
)

// countingHandler 记录被分派到的请求数,用于断言超限帧未进入业务处理。
type countingHandler struct {
bannet.BaseRouter
handled chan struct{}
}

func (h *countingHandler) Handle(bannet.Request) {
select {
case h.handled <- struct{}{}:
default:
}
}

// TestOversizedFrameRejectedBeforeReadingPayload 守护帧长上限的执行点。
//
// 该校验此前在 DataPack.UnPack 内,每帧读两次全局配置;现已移到连接侧(读取负载之前),
// 让编解码器保持无状态。移动执行点意味着必须有测试证明它仍在执行——否则一个恶意或损坏的
// 帧头就能让服务端按对端声称的长度分配内存。
//
// 构造方式:只发一个声称负载极大的帧头,不发负载。服务端必须在读取负载前就断开连接,
// 且该帧不得被分派给业务处理。
func TestOversizedFrameRejectedBeforeReadingPayload(t *testing.T) {
ln, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil {
t.Fatalf("listen: %v", err)
}
host, portStr, _ := net.SplitHostPort(ln.Addr().String())
port, _ := strconv.Atoi(portStr)
addr := ln.Addr().String()
ln.Close()

oldMax := config.G.MaxPackageSize
config.G.MaxPackageSize = 1024
t.Cleanup(func() { config.G.MaxPackageSize = oldMax })

h := &countingHandler{handled: make(chan struct{}, 1)}
srv := bannet.NewServer()
srv.IP, srv.Port = host, port
srv.AddRouter(proto.MsgPut, h)
srv.Start()
t.Cleanup(srv.Stop)

var conn net.Conn
for i := 0; i < 100; i++ {
if conn, err = net.DialTimeout("tcp", addr, 200*time.Millisecond); err == nil {
break
}
time.Sleep(20 * time.Millisecond)
}
if conn == nil {
t.Fatalf("服务端未就绪: %v", err)
}
defer conn.Close()

// 帧头: [dataLen u32 LE][idLen u16 LE],声称 64MiB 负载但一个字节都不发。
head := make([]byte, 6)
binary.LittleEndian.PutUint32(head[0:4], 64<<20)
binary.LittleEndian.PutUint16(head[4:6], uint16(len(proto.MsgPut)))
if _, err := conn.Write(append(head, []byte(proto.MsgPut)...)); err != nil {
t.Fatalf("write: %v", err)
}

// 服务端应断开连接:读将返回 EOF/RST,而不是一直等着那 64MiB。
conn.SetReadDeadline(time.Now().Add(2 * time.Second))
buf := make([]byte, 1)
if _, err := conn.Read(buf); err == nil {
t.Fatal("超限帧后连接仍可读,应已被服务端关闭")
} else if ne, ok := err.(net.Error); ok && ne.Timeout() {
t.Fatal("服务端既未拒绝也未关闭连接——很可能正按声称的长度等待/分配 64MiB")
}

select {
case <-h.handled:
t.Fatal("超限帧不应被分派到业务处理")
default:
}
}
2 changes: 1 addition & 1 deletion cmd/ban-ingest/main.go
Original file line number Diff line number Diff line change
Expand Up @@ -99,7 +99,7 @@ func setupEngine(memTableSize int) (*storage.Engine, func()) {
config.G.WALPath = filepath.Join(tmp, "wal.log")
config.G.MaxMemTableSize = memTableSize

memTable := storage.NewEngine()
memTable := storage.NewEngine(storage.DefaultOptions())
cleanup := func() {
_ = memTable.Close()
os.RemoveAll(tmp)
Expand Down
2 changes: 1 addition & 1 deletion service/fsm.go
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@ type KVServer struct {
func NewKVServer() *KVServer {
// 初始化存储
kv := &KVServer{
storage: storage.NewEngine(),
storage: storage.NewEngine(storage.DefaultOptions()),
}

if config.G.Mode == config.ModeStandalone {
Expand Down
19 changes: 6 additions & 13 deletions storage/bench_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,15 +2,12 @@ package storage_test

import (
"fmt"
"testing"

"github.com/NeverENG/BanDB/config"
"github.com/NeverENG/BanDB/storage"
"testing"
)

func benchmarkEnginePut(b *testing.B, valueSize int) {
config.G.MaxMemTableSize = 1000000 // prevent flush during bench
memTable := storage.NewEngine()
memTable := storage.NewEngine(storage.Options{Dir: b.TempDir(), MaxMemTableSize: 1000000}) // 阈值取大,基准期间不触发 flush

value := make([]byte, valueSize)
for i := range value {
Expand All @@ -30,8 +27,7 @@ func BenchmarkEngine_Put_1KB(b *testing.B) { benchmarkEnginePut(b, 1024) }
func BenchmarkEngine_Put_4KB(b *testing.B) { benchmarkEnginePut(b, 4096) }

func BenchmarkEngine_Get(b *testing.B) {
config.G.MaxMemTableSize = 1000000
memTable := storage.NewEngine()
memTable := storage.NewEngine(storage.Options{Dir: b.TempDir(), MaxMemTableSize: 1000000}) // 阈值取大,基准期间不触发 flush

value := make([]byte, 256)
for i := 0; i < 10000; i++ {
Expand All @@ -47,8 +43,7 @@ func BenchmarkEngine_Get(b *testing.B) {
}

func BenchmarkEngine_Delete(b *testing.B) {
config.G.MaxMemTableSize = 1000000
memTable := storage.NewEngine()
memTable := storage.NewEngine(storage.Options{Dir: b.TempDir(), MaxMemTableSize: 1000000}) // 阈值取大,基准期间不触发 flush

value := make([]byte, 256)
keys := make([][]byte, b.N)
Expand All @@ -64,8 +59,7 @@ func BenchmarkEngine_Delete(b *testing.B) {
}

func BenchmarkMemTable_Put(b *testing.B) {
config.G.MaxMemTableSize = 1000000
mt := storage.NewEngine()
mt := storage.NewEngine(storage.Options{Dir: b.TempDir(), MaxMemTableSize: 1000000}) // 阈值取大,基准期间不触发 flush
value := []byte("benchmark-value-data-256-bytes-padding-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx")

b.ResetTimer()
Expand All @@ -76,8 +70,7 @@ func BenchmarkMemTable_Put(b *testing.B) {
}

func BenchmarkMemTable_Get(b *testing.B) {
config.G.MaxMemTableSize = 1000000
mt := storage.NewEngine()
mt := storage.NewEngine(storage.Options{Dir: b.TempDir(), MaxMemTableSize: 1000000}) // 阈值取大,基准期间不触发 flush
value := []byte("benchmark-value")
for i := 0; i < 100000; i++ {
key := []byte(fmt.Sprintf("key-%08d", i))
Expand Down
13 changes: 4 additions & 9 deletions storage/compaction_bench_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,8 +5,6 @@ import (
"runtime"
"testing"
"time"

"github.com/NeverENG/BanDB/config"
)

// TestCompactionBench 是 compaction 压测台:确定性地驱动真实的 flush + compaction 级联,
Expand All @@ -17,16 +15,13 @@ import (
func TestCompactionBench(t *testing.T) {
// 小 memtable + 小 compaction 阈值,逼出频繁 flush 与多级 compaction。
dir := t.TempDir()
oldPath, oldComp := config.G.SSTablePath, config.G.MaxCompactionSize
config.G.SSTablePath = dir
config.G.MaxCompactionSize = 4
defer func() { config.G.SSTablePath, config.G.MaxCompactionSize = oldPath, oldComp }()
opts := Options{Dir: dir, MaxCompactionSize: 4}

ResetCompactionStats()

// 不经 NewEngine(避免启动异步 FlushWorker/ListenCompactCh 造成非确定性),
// 只用 sst,手动驱动 flush 与 compaction。
mt := newBareMemTable(NewSSTable())
mt := newBareMemTable(NewSSTable(opts), opts)

const (
flushes = 300
Expand Down Expand Up @@ -90,7 +85,7 @@ func TestCompactionBench(t *testing.T) {
}

// === 模拟重启:LoadSSTableMetaList 把所有文件 Level 归 0(level 未持久化)===
sst2 := NewSSTable()
sst2 := NewSSTable(opts)
sst2.LoadSSTableMetaList()
time.Sleep(150 * time.Millisecond) // 等异步预热 goroutine 落定,避免与后续 compaction 竞争
dist2, total2 := levelDistribution(sst2)
Expand All @@ -101,7 +96,7 @@ func TestCompactionBench(t *testing.T) {
}
t.Logf("文件总数=%d per-level=%s ← %s", total2, dist2, collapsed)

mt2 := newBareMemTable(sst2)
mt2 := newBareMemTable(sst2, opts)
before := ReadCompactionStats()
t0 := time.Now()
mt2.CompactSSTable(0)
Expand Down
17 changes: 4 additions & 13 deletions storage/compaction_failure_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,8 +4,6 @@ import (
"fmt"
"os"
"testing"

"github.com/NeverENG/BanDB/config"
)

// writeSSTables 写出 n 个各含若干 key 的 SSTable,返回它们的路径。
Expand Down Expand Up @@ -39,11 +37,9 @@ func writeSSTables(t *testing.T, ss *SSTable, n, perFile int) []string {
// 故 root 下会失效——检测到仍能创建文件时直接跳过,而不是给出假绿。
func TestCompactionFailureKeepsSourceFiles(t *testing.T) {
dir := t.TempDir()
oldSST := config.G.SSTablePath
config.G.SSTablePath = dir
t.Cleanup(func() { config.G.SSTablePath = oldSST })
opts := Options{Dir: dir}

ss := NewSSTable()
ss := NewSSTable(opts)
paths := writeSSTables(t, ss, 3, 50)

if err := os.Chmod(dir, 0o500); err != nil { // r-x:可读可遍历,不可创建新文件
Expand Down Expand Up @@ -87,14 +83,9 @@ func TestCompactionFailureKeepsSourceFiles(t *testing.T) {
// 在合并失败时不得删除任何源文件,且全部 key 仍可读。
func TestCompactSSTableKeepsSourcesOnMergeFailure(t *testing.T) {
dir := t.TempDir()
oldSST, oldCompact := config.G.SSTablePath, config.G.MaxCompactionSize
config.G.SSTablePath = dir
config.G.MaxCompactionSize = 2 // 两个文件即触发合并
t.Cleanup(func() {
config.G.SSTablePath, config.G.MaxCompactionSize = oldSST, oldCompact
})
opts := Options{Dir: dir, MaxCompactionSize: 2} // 两个文件即触发合并

e := NewEngine()
e := NewEngine(opts)
t.Cleanup(func() { e.Close() })
paths := writeSSTables(t, e.sst, 3, 30)

Expand Down
11 changes: 3 additions & 8 deletions storage/compaction_level_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -4,21 +4,16 @@ import (
"fmt"
"testing"
"time"

"github.com/NeverENG/BanDB/config"
)

// TestCompaction_LevelPersistedAcrossRestart 是「重启不塌缩」的回归守卫:
// 造出跨多个 level 的文件后,模拟重启(新建 SSTable + LoadSSTableMetaList),
// 断言 per-level 分布被保留,而非全部塌缩到 L0(后者会在重启后触发全量重写)。
func TestCompaction_LevelPersistedAcrossRestart(t *testing.T) {
dir := t.TempDir()
oldPath, oldComp := config.G.SSTablePath, config.G.MaxCompactionSize
config.G.SSTablePath = dir
config.G.MaxCompactionSize = 4
defer func() { config.G.SSTablePath, config.G.MaxCompactionSize = oldPath, oldComp }()
opts := Options{Dir: dir, MaxCompactionSize: 4}

mt := newBareMemTable(NewSSTable())
mt := newBareMemTable(NewSSTable(opts), opts)
val := make([]byte, 32)
global := 0
for f := 0; f < 40; f++ {
Expand All @@ -40,7 +35,7 @@ func TestCompaction_LevelPersistedAcrossRestart(t *testing.T) {
}

// 模拟重启:全新 SSTable 从磁盘恢复。
sst2 := NewSSTable()
sst2 := NewSSTable(opts)
sst2.LoadSSTableMetaList()
time.Sleep(150 * time.Millisecond)
after, totalAfter := levelDistribution(sst2)
Expand Down
Loading
Loading