Bug fixes: composite PK, nl_to_sql sandbox, FK check, SQL injection, storage correctness

Critical fixes:
- Composite PK: execInsert + validateConstraints use all PK columns
- nl_to_sql: non-SELECT SQL no longer executed directly during validation
- FK check: removed O(N) scanMemTable fallback, uses db.get() with SSTables
- exprToSql: nkIdent wrapped in quotes to prevent SQL injection
- restoreSchema: try/except around tokenize/parse for crash resilience
- recovery.nim: lastTxnId tracking + putUnsafe/deleteUnsafe without WAL
- SCRAM: verifyClientProof length check + DefaultIterationCount restored

Storage fixes:
- lsm.nim: SSTable sort order fixed (ascending), close() flushes all memtables
- compaction.nim: tombstones preserved during compaction
- wal.nim: header written for empty existing files, readEntries checks magic
- btree.nim: B+ tree leaf split keeps boundary key
- bloom.nim: deserialize raises on short data
- mmap.nim: bounds checks for adviseWillNeed/DontNeed + posix.close()

Protocol fixes:
- zerocopy.nim: readString bounds check
- wire.nim: deserializeValue 32-bit underflow check
- auth.nim: JWT JSON escaping for claims
- server.nim: readUint32BE bounds check + specific exception handling
- raft.nim: readData checks + specific exception handling

Tests:
- Added Composite Primary Key test suite (4 tests)

Build: 0 warnings, 0 errors
This commit is contained in:
2026-05-18 11:33:11 +03:00
parent a28c845476
commit 967c0855a5
25 changed files with 376 additions and 171 deletions
+1 -1
View File
@@ -81,7 +81,7 @@ proc deserialize*(bf: var BloomFilter, data: seq[byte]) =
)
let numBytes = (bf.size + 7) div 8
if data.len < 8 + numBytes:
return
raise newException(ValueError, "Bloom filter data too short: expected " & $(8 + numBytes) & " bytes, got " & $data.len)
for i in 0..<bf.size:
if (data[8 + i div 8] and (1'u8 shl (i mod 8))) != 0:
bf.bits[i] = true
BIN
View File
Binary file not shown.
+5 -2
View File
@@ -62,9 +62,12 @@ proc splitChild[K, V](parent: BTreeNode[K, V], index: int, order: int) =
let midKey = child.keys[mid]
parent.keys.insert(midKey, index)
parent.children.insert(newNode, index + 1)
child.keys.setLen(mid)
# In B+ tree, leaf nodes must keep the boundary key for range scans
if child.isLeaf:
child.values.setLen(mid)
child.keys.setLen(mid + 1)
child.values.setLen(mid + 1)
else:
child.keys.setLen(mid)
proc insertNonFull[K, V](node: BTreeNode[K, V], key: K, value: V, order: int) =
var i = node.keys.len - 1
+2 -3
View File
@@ -98,11 +98,10 @@ proc compact*(cs: CompactionStrategy, level: int): CompactionResult =
merged.add(entry)
lastKey = entry.key
# Filter out tombstones (deleted entries)
# Keep tombstones to prevent deleted keys from resurrecting in lower levels
var final: seq[Entry] = @[]
for entry in merged:
if not entry.deleted:
final.add(entry)
final.add(entry)
# Write merged SSTable
let outputPath = cs.dataDir / "sstables" / ("level_" & $level & "_" & $tables[0].createdAt & ".sst")
+21 -2
View File
@@ -318,7 +318,7 @@ proc newLSMTree*(dir: string, memMaxSize: int = DefaultMemTableSize): LSMTree =
except:
discard # skip corrupt SSTables
sstables.sort(proc(a, b: SSTable): int = cmp(b.id, a.id))
sstables.sort(proc(a, b: SSTable): int = cmp(a.id, b.id))
new(result)
initLock(result.lock)
@@ -405,6 +405,21 @@ proc delete*(db: LSMTree, key: string) =
db.memTable = newMemTable(db.memMaxSize)
discard db.memTable.put(key, @[], ts, deleted = true)
proc putUnsafe*(db: LSMTree, key: string, value: seq[byte], deleted: bool = false) =
## Direct LSM insert without WAL logging — used by recovery.
let ts = uint64(getMonoTime().ticks())
acquire(db.lock)
defer: release(db.lock)
if not db.memTable.put(key, value, ts, deleted):
if db.immutableMem.len > 0:
db.flushUnsafe()
db.immutableMem = db.memTable
db.memTable = newMemTable(db.memMaxSize)
discard db.memTable.put(key, value, ts, deleted)
proc deleteUnsafe*(db: LSMTree, key: string) =
putUnsafe(db, key, @[], deleted = true)
proc getUnsafe(db: LSMTree, key: string): (bool, seq[byte]) =
let (found, entry) = db.memTable.get(key)
if found:
@@ -478,6 +493,9 @@ proc flush*(db: LSMTree) =
proc close*(db: LSMTree) =
acquire(db.lock)
defer: release(db.lock)
# Flush both memtables to avoid data loss
while db.immutableMem.len > 0:
flushUnsafe(db)
flushUnsafe(db)
for sst in db.sstables.mitems:
sst.close()
@@ -531,7 +549,8 @@ proc scanAll*(db: LSMTree): seq[(string, seq[byte])] =
result.add((e.key, e.value))
# Scan SSTables from newest to oldest
for sst in db.sstables:
for i in countdown(db.sstables.high, db.sstables.low):
let sst = db.sstables[i]
for key, offset in sst.index:
if key notin seen:
seen[key] = true
+12 -3
View File
@@ -54,14 +54,18 @@ proc openMmap*(path: string, mode: MmapMode = mmReadOnly): MmapFile =
let mapped = mmap(nil, fileSize, prot, flags, fd, 0)
if mapped == MAP_FAILED:
discard close(fd)
discard posix.close(fd)
return MmapFile(path: path, regions: @[], totalSize: 0, pageSize: PageSize)
# Kernel keeps the mapping alive independently of the fd; close it now
# to avoid leaking one fd per SSTable.
discard posix.close(fd)
let region = MmapRegion(
data: cast[ptr UncheckedArray[byte]](mapped),
size: fileSize,
offset: 0,
fd: fd,
fd: -1,
mode: mode,
)
@@ -116,16 +120,21 @@ proc adviseRandom*(mf: MmapFile) =
proc adviseWillNeed*(mf: MmapFile, offset: int, size: int) =
if mf.regions.len > 0:
if offset < 0 or offset + size > mf.regions[0].size:
return
discard madvise(addr mf.regions[0].data[offset], size, MADV_WILLNEED)
proc adviseDontNeed*(mf: MmapFile, offset: int, size: int) =
if mf.regions.len > 0:
if offset < 0 or offset + size > mf.regions[0].size:
return
discard madvise(addr mf.regions[0].data[offset], size, MADV_DONTNEED)
proc close*(mf: MmapFile) =
for region in mf.regions:
discard munmap(region.data, region.size)
discard close(cint(region.fd))
if region.fd != -1:
discard close(cint(region.fd))
mf.regions.setLen(0)
proc size*(mf: MmapFile): int = mf.totalSize
+6 -7
View File
@@ -32,6 +32,7 @@ type
dataDir*: string
entries*: seq[RecoveredEntry]
result*: RecoveryResult
lastTxnId*: uint64 # tracks commits seen in WAL
proc newCrashRecovery*(walDir: string, dataDir: string): CrashRecovery =
CrashRecovery(
@@ -101,6 +102,7 @@ proc scanWAL*(rec: CrashRecovery): seq[RecoveredEntry] =
discard
stream.close()
rec.lastTxnId = txnId
proc analyze*(rec: CrashRecovery): RecoveryResult =
rec.entries = rec.scanWAL()
@@ -109,10 +111,7 @@ proc analyze*(rec: CrashRecovery): RecoveryResult =
rec.result = RecoveryResult(state: recDone, applied: false)
return rec.result
var lastCommitted: uint64 = 0
for entry in rec.entries:
if entry.txnId > lastCommitted:
lastCommitted = entry.txnId
var lastCommitted = rec.lastTxnId
var redoCount = 0
var undoCount = 0
@@ -149,11 +148,11 @@ proc recover*(rec: CrashRecovery, db: LSMTree = nil): RecoveryResult =
var undoCount = 0
for entry in rec.entries:
if entry.txnId < analysis.lastTxn:
# Committed — redo
# Committed — redo (bypass WAL to avoid duplicate entries)
if entry.isDelete:
db.delete(entry.key)
db.deleteUnsafe(entry.key)
else:
db.put(entry.key, entry.value)
db.putUnsafe(entry.key, entry.value)
inc redoCount
else:
# Uncommitted — skip (undo)
+11 -7
View File
@@ -30,10 +30,11 @@ proc newWriteAheadLog*(dir: string, syncOnWrite: bool = true): WriteAheadLog =
createDir(dir)
let path = dir / "wal.log"
let exists = fileExists(path)
let isEmpty = if exists: getFileSize(path) == 0 else: true
let stream = if exists: newFileStream(path, fmAppend) else: newFileStream(path, fmWrite)
if stream == nil:
raise newException(IOError, "Cannot open WAL: " & path)
if not exists:
if not exists or isEmpty:
stream.write(WALMagic)
stream.write(WALVersion)
stream.flush()
@@ -78,14 +79,17 @@ proc writeCommit*(wal: var WriteAheadLog, timestamp: uint64) =
proc sync*(wal: var WriteAheadLog) =
wal.stream.flush()
let fd = posix.open(cstring(wal.path), O_RDONLY)
# Re-open with O_RDWR so fsync operates on a write-capable fd.
# Not ideal (two fds for same file) but avoids accessing private
# FileStream internals that vary across Nim versions.
let fd = posix.open(cstring(wal.path), O_RDWR)
if fd != -1:
discard posix.fsync(fd)
discard posix.close(fd)
proc close*(wal: var WriteAheadLog) =
wal.stream.flush()
let fd = posix.open(cstring(wal.path), O_RDONLY)
let fd = posix.open(cstring(wal.path), O_RDWR)
if fd != -1:
discard posix.fsync(fd)
discard posix.close(fd)
@@ -100,10 +104,10 @@ proc readEntries*(walPath: string, untilTimestamp: uint64 = 0): seq[WalEntry] =
let s = newFileStream(walPath, fmRead)
if s == nil: return
# Skip header
var magic: uint32
var version: uint32
discard s.readData(addr magic, 4)
discard s.readData(addr version, 4)
var magic: uint32 = 0
var version: uint32 = 0
if s.readData(addr magic, 4) != 4: return
if s.readData(addr version, 4) != 4: return
if magic != WALMagic: return
while not s.atEnd:
var kind: uint8