feat(raft): classify writes, reject them on follower nodes

This commit is contained in:
2026-07-30 19:07:58 +03:00
parent 8a35a838d0
commit 38c1c01841
4 changed files with 36 additions and 3 deletions
+13 -3
View File
@@ -20,11 +20,13 @@ import ../query/lexer
import ../query/parser
import ../query/ast
import ../query/executor
import ../query/exec/params
import ../storage/lsm
import ../storage/gate
import ../core/mvcc
import ../core/disttxn
import ../core/replication
import ../core/raft
import ../core/sharding
import ../core/gossip
import ../protocol/ratelimit
@@ -41,6 +43,7 @@ type
txnManager*: TxnManager
distTxnManager*: DistTxnManager
replicationManager*: ReplicationManager
raftNode*: RaftNode
shardRouter*: ShardRouter
clusterMembership*: ClusterMembership
gossipProtocol*: GossipProtocol
@@ -204,7 +207,8 @@ proc valueToWire(val: string, colType: string): WireValue =
return WireValue(kind: fkString, strVal: val)
proc executeQuery(db: LSMTree, ctx: ExecutionContext, query: string, params: seq[WireValue] = @[],
replication: ReplicationManager = nil): (bool, QueryResult, string) =
replication: ReplicationManager = nil,
raftNode: RaftNode = nil): (bool, QueryResult, string) =
## All storage access is under the global StorageGate so HTTP worker threads
## and the TCP event loop never touch ORC-managed LSM/executor state concurrently.
withStorageGate:
@@ -215,6 +219,12 @@ proc executeQuery(db: LSMTree, ctx: ExecutionContext, query: string, params: seq
if astNode.stmts.len == 0:
return (true, QueryResult(), "")
# C3b: writes go through the Raft log — only the leader may accept them.
if raftNode != nil and isWrite(astNode.stmts[0]):
if raftNode.state != rsLeader:
let who = if raftNode.leaderId.len > 0: raftNode.leaderId else: "none elected"
return (false, QueryResult(), "not leader; leader is '" & who & "'")
let res = executor.executeQuery(ctx, astNode, params)
if res.success:
# Ship written key-value pairs to replicas
@@ -564,7 +574,7 @@ proc handleClient(server: Server, client: AsyncSocket, clientId: int) {.async.}
if shardCheck:
let startTicks = getMonoTime().ticks()
let (success, result, errorMsg) = executeQuery(connCtx.db, connCtx, queryStr, replication=server.replicationManager)
let (success, result, errorMsg) = executeQuery(connCtx.db, connCtx, queryStr, replication=server.replicationManager, raftNode=server.raftNode)
let durationMs = int((getMonoTime().ticks() - startTicks) div 1_000_000)
if durationMs >= slowThreshold:
@@ -584,7 +594,7 @@ proc handleClient(server: Server, client: AsyncSocket, clientId: int) {.async.}
info("[" & $clientId & "] QueryParams: " & queryStr & " (" & $params.len & " params)")
let startTicks = getMonoTime().ticks()
let (success, result, errorMsg) = executeQuery(connCtx.db, connCtx, queryStr, params, replication=server.replicationManager)
let (success, result, errorMsg) = executeQuery(connCtx.db, connCtx, queryStr, params, replication=server.replicationManager, raftNode=server.raftNode)
let durationMs = int((getMonoTime().ticks() - startTicks) div 1_000_000)
if durationMs >= slowThreshold: