diff --git a/src/barabadb/core/server.nim b/src/barabadb/core/server.nim index d7dc24b..3613e76 100644 --- a/src/barabadb/core/server.nim +++ b/src/barabadb/core/server.nim @@ -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: diff --git a/src/barabadb/query/exec/params.nim b/src/barabadb/query/exec/params.nim index cd195b7..7dc64b5 100644 --- a/src/barabadb/query/exec/params.nim +++ b/src/barabadb/query/exec/params.nim @@ -178,3 +178,12 @@ proc isDDL*(stmt: Node): bool = result = true else: result = false + +proc isWrite*(stmt: Node): bool = + ## True for statements that mutate stored data. `nkCommitTxn` is included + ## because COMMIT emits the transaction's buffered kvPairs. + case stmt.kind + of nkInsert, nkUpdate, nkDelete, nkMerge, nkCommitTxn: + result = true + else: + result = false diff --git a/src/baradadb.nim b/src/baradadb.nim index b3348a0..d198621 100644 --- a/src/baradadb.nim +++ b/src/baradadb.nim @@ -341,6 +341,7 @@ proc main() = var raftNode = newRaftNode(config.raftNodeId, raftPeers, config.raftPort, dataDir = raftDataDir) raftNode.peerAddrs = config.raftPeerAddrs + tcpServer.raftNode = raftNode # C3b: executeQuery rejects writes on followers # Wire state machine to apply committed entries to the default database let defaultDbInfo = getDatabaseInfo(registry, "default") raftNode.applyCommand = proc(cmd: string, data: seq[byte]) {.gcsafe.} = diff --git a/tests/bugfix_test.nim b/tests/bugfix_test.nim index 1f2d46d..12ab37d 100644 --- a/tests/bugfix_test.nim +++ b/tests/bugfix_test.nim @@ -3,6 +3,7 @@ import std/strutils import std/os import std/tables import ../src/barabadb/query/[parser, executor, lexer, ast] +import ../src/barabadb/query/exec/params import ../src/barabadb/core/types import ../src/barabadb/core/config import ../src/barabadb/storage/lsm @@ -383,3 +384,15 @@ suite "Raft peer address parsing": delEnv("BARADB_RAFT_PEERS") check msg.len > 0 check bad in msg + +suite "Raft write classification": + + test "isWrite classifies DML and COMMIT": + check isWrite(parse("INSERT INTO t (id) VALUES (1)").stmts[0]) + check isWrite(parse("UPDATE t SET id = 2").stmts[0]) + check isWrite(parse("DELETE FROM t WHERE id = 1").stmts[0]) + check isWrite(parse("COMMIT").stmts[0]) + check not isWrite(parse("SELECT * FROM t").stmts[0]) + check not isWrite(parse("CREATE TABLE t (id INT)").stmts[0]) + check not isWrite(parse("BEGIN").stmts[0]) + check not isWrite(parse("ROLLBACK").stmts[0])