feat(raft): optional TLS on raft transport (server + dialer)
This commit is contained in:
@@ -12,6 +12,7 @@ import std/endians
|
|||||||
import std/os
|
import std/os
|
||||||
import logging
|
import logging
|
||||||
import ../protocol/wire
|
import ../protocol/wire
|
||||||
|
import ../protocol/ssl
|
||||||
|
|
||||||
type
|
type
|
||||||
RaftState* = enum
|
RaftState* = enum
|
||||||
@@ -734,13 +735,16 @@ type
|
|||||||
running*: bool
|
running*: bool
|
||||||
peerSockets*: Table[string, AsyncSocket]
|
peerSockets*: Table[string, AsyncSocket]
|
||||||
timer*: ElectionTimer
|
timer*: ElectionTimer
|
||||||
|
## Optional TLS context; nil = plaintext (default, pre-TLS behavior).
|
||||||
|
tls*: TLSContext
|
||||||
|
|
||||||
proc newRaftNetwork*(node: RaftNode): RaftNetwork =
|
proc newRaftNetwork*(node: RaftNode, tls: TLSContext = nil): RaftNetwork =
|
||||||
RaftNetwork(
|
RaftNetwork(
|
||||||
node: node,
|
node: node,
|
||||||
running: false,
|
running: false,
|
||||||
peerSockets: initTable[string, AsyncSocket](),
|
peerSockets: initTable[string, AsyncSocket](),
|
||||||
timer: newElectionTimer(node, node.electionTimeout),
|
timer: newElectionTimer(node, node.electionTimeout),
|
||||||
|
tls: tls,
|
||||||
)
|
)
|
||||||
|
|
||||||
const RaftConnectTimeoutMs = 200
|
const RaftConnectTimeoutMs = 200
|
||||||
@@ -759,6 +763,12 @@ proc connectToPeer(net: RaftNetwork, peerId: string) {.async.} =
|
|||||||
if not ok:
|
if not ok:
|
||||||
sock.close()
|
sock.close()
|
||||||
return
|
return
|
||||||
|
if net.tls != nil:
|
||||||
|
try:
|
||||||
|
net.tls.wrapClient(sock)
|
||||||
|
except CatchableError:
|
||||||
|
try: sock.close() except CatchableError: discard
|
||||||
|
return
|
||||||
net.peerSockets[peerId] = sock
|
net.peerSockets[peerId] = sock
|
||||||
except CatchableError:
|
except CatchableError:
|
||||||
if sock != nil:
|
if sock != nil:
|
||||||
@@ -866,6 +876,13 @@ proc run*(net: RaftNetwork) {.async.} =
|
|||||||
while net.running:
|
while net.running:
|
||||||
try:
|
try:
|
||||||
let client = await net.socket.accept()
|
let client = await net.socket.accept()
|
||||||
|
if net.tls != nil:
|
||||||
|
try:
|
||||||
|
net.tls.wrapServer(client)
|
||||||
|
except CatchableError:
|
||||||
|
# Handshake failed (e.g. plaintext dial) — drop, no protocol effect.
|
||||||
|
client.close()
|
||||||
|
continue
|
||||||
asyncCheck net.receiveLoop(client)
|
asyncCheck net.receiveLoop(client)
|
||||||
except CatchableError:
|
except CatchableError:
|
||||||
break
|
break
|
||||||
|
|||||||
@@ -29,10 +29,14 @@ proc newTLSConfig*(certFile: string, keyFile: string, caFile: string = "",
|
|||||||
proc newTLSContext*(config: TLSConfig): TLSContext =
|
proc newTLSContext*(config: TLSConfig): TLSContext =
|
||||||
result = TLSContext(config: config)
|
result = TLSContext(config: config)
|
||||||
if fileExists(config.certFile) and fileExists(config.keyFile):
|
if fileExists(config.certFile) and fileExists(config.keyFile):
|
||||||
|
# caFile is only honored by newContext when verifyPeer is true
|
||||||
|
# (verifyMode != CVerifyNone); a missing CA file then raises IOError,
|
||||||
|
# which is the desired fail-closed behavior.
|
||||||
result.sslCtx = newContext(
|
result.sslCtx = newContext(
|
||||||
certFile = config.certFile,
|
certFile = config.certFile,
|
||||||
keyFile = config.keyFile,
|
keyFile = config.keyFile,
|
||||||
verifyMode = if config.verifyPeer: CVerifyPeer else: CVerifyNone,
|
verifyMode = if config.verifyPeer: CVerifyPeer else: CVerifyNone,
|
||||||
|
caFile = config.caFile,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
raise newException(IOError, "TLS certificate or key file not found: " &
|
raise newException(IOError, "TLS certificate or key file not found: " &
|
||||||
@@ -40,7 +44,11 @@ proc newTLSContext*(config: TLSConfig): TLSContext =
|
|||||||
|
|
||||||
proc wrapClient*(tls: TLSContext, socket: AsyncSocket) {.inline.} =
|
proc wrapClient*(tls: TLSContext, socket: AsyncSocket) {.inline.} =
|
||||||
if tls.sslCtx != nil:
|
if tls.sslCtx != nil:
|
||||||
tls.sslCtx.wrapSocket(socket)
|
# wrapConnectedSocket (asyncnet overload) sets connect state; the
|
||||||
|
# handshake itself is driven lazily by the first send/recv. Plain
|
||||||
|
# wrapSocket leaves the SSL handle in SSL_ST_BEFORE and the first
|
||||||
|
# SSL_write fails with "uninitialized".
|
||||||
|
tls.sslCtx.wrapConnectedSocket(socket, handshakeAsClient)
|
||||||
|
|
||||||
proc wrapServer*(tls: TLSContext, socket: AsyncSocket) {.inline.} =
|
proc wrapServer*(tls: TLSContext, socket: AsyncSocket) {.inline.} =
|
||||||
if tls.sslCtx != nil:
|
if tls.sslCtx != nil:
|
||||||
|
|||||||
+10
-1
@@ -338,12 +338,21 @@ proc main() =
|
|||||||
var raftNet: RaftNetwork = nil
|
var raftNet: RaftNetwork = nil
|
||||||
if config.raftEnabled:
|
if config.raftEnabled:
|
||||||
info("Starting Raft node " & config.raftNodeId & " on port " & $config.raftPort)
|
info("Starting Raft node " & config.raftNodeId & " on port " & $config.raftPort)
|
||||||
|
var raftTls: TLSContext = nil
|
||||||
if config.raftTlsEnabled:
|
if config.raftTlsEnabled:
|
||||||
if config.raftTlsCertFile.len == 0 or config.raftTlsKeyFile.len == 0 or
|
if config.raftTlsCertFile.len == 0 or config.raftTlsKeyFile.len == 0 or
|
||||||
not fileExists(config.raftTlsCertFile) or not fileExists(config.raftTlsKeyFile):
|
not fileExists(config.raftTlsCertFile) or not fileExists(config.raftTlsKeyFile):
|
||||||
raise newException(ValueError,
|
raise newException(ValueError,
|
||||||
"BARADB_RAFT_TLS_ENABLED=true but cert/key missing " &
|
"BARADB_RAFT_TLS_ENABLED=true but cert/key missing " &
|
||||||
"(BARADB_RAFT_TLS_CERT_FILE / BARADB_RAFT_TLS_KEY_FILE)")
|
"(BARADB_RAFT_TLS_CERT_FILE / BARADB_RAFT_TLS_KEY_FILE)")
|
||||||
|
if config.raftTlsVerifyPeer and config.raftTlsCaFile.len > 0 and
|
||||||
|
not fileExists(config.raftTlsCaFile):
|
||||||
|
raise newException(ValueError,
|
||||||
|
"BARADB_RAFT_TLS_VERIFY_PEER=true but CA file missing: " &
|
||||||
|
config.raftTlsCaFile & " (BARADB_RAFT_TLS_CA_FILE)")
|
||||||
|
raftTls = newTLSContext(newTLSConfig(
|
||||||
|
config.raftTlsCertFile, config.raftTlsKeyFile,
|
||||||
|
caFile = config.raftTlsCaFile, verifyPeer = config.raftTlsVerifyPeer))
|
||||||
let raftDataDir = config.dataDir / "raft"
|
let raftDataDir = config.dataDir / "raft"
|
||||||
createDir(raftDataDir) # idempotent; loadState reads from it, saveState writes
|
createDir(raftDataDir) # idempotent; loadState reads from it, saveState writes
|
||||||
# Raft convention: `peers` excludes the node itself (majority math and
|
# Raft convention: `peers` excludes the node itself (majority math and
|
||||||
@@ -379,7 +388,7 @@ proc main() =
|
|||||||
# Wire replication ↔ DistTxn
|
# Wire replication ↔ DistTxn
|
||||||
wireReplicationDistTxn(tcpServer.replicationManager, tcpServer.distTxnManager)
|
wireReplicationDistTxn(tcpServer.replicationManager, tcpServer.distTxnManager)
|
||||||
|
|
||||||
raftNet = newRaftNetwork(raftNode)
|
raftNet = newRaftNetwork(raftNode, raftTls)
|
||||||
asyncCheck raftNet.run()
|
asyncCheck raftNet.run()
|
||||||
|
|
||||||
# HTTP (hunos) after raft wiring so /metrics can see raftNode
|
# HTTP (hunos) after raft wiring so /metrics can see raftNode
|
||||||
|
|||||||
@@ -2607,6 +2607,96 @@ suite "Raft Network Transport":
|
|||||||
check replyMsg.kind == rmkRequestVoteReply
|
check replyMsg.kind == rmkRequestVoteReply
|
||||||
check replyMsg.success
|
check replyMsg.success
|
||||||
|
|
||||||
|
suite "Raft TLS Transport":
|
||||||
|
test "2-node election over TLS":
|
||||||
|
let certDir = getTempDir() / "baradb_test_raft_tls"
|
||||||
|
let (certPath, keyPath) = generateSelfSignedCert(certDir, "raft-tls.local")
|
||||||
|
if certPath.len == 0:
|
||||||
|
skip() # openssl unavailable
|
||||||
|
else:
|
||||||
|
let tls = newTLSContext(newTLSConfig(certPath, keyPath))
|
||||||
|
var n1 = newRaftNode("n1", @["n2"], raftPort = 29301)
|
||||||
|
var n2 = newRaftNode("n2", @["n1"], raftPort = 29302)
|
||||||
|
n1.electionTimeout = 150
|
||||||
|
n2.electionTimeout = 350
|
||||||
|
n1.peerAddrs["n2"] = ("127.0.0.1", 29302)
|
||||||
|
n2.peerAddrs["n1"] = ("127.0.0.1", 29301)
|
||||||
|
|
||||||
|
let net1 = newRaftNetwork(n1, tls)
|
||||||
|
let net2 = newRaftNetwork(n2, tls)
|
||||||
|
|
||||||
|
asyncCheck net1.run()
|
||||||
|
asyncCheck net2.run()
|
||||||
|
waitFor sleepAsync(50)
|
||||||
|
|
||||||
|
# No manual ticks — timerLoop drives the election over TLS.
|
||||||
|
var leaderCount = 0
|
||||||
|
var waited = 0
|
||||||
|
while waited < 3000:
|
||||||
|
leaderCount = 0
|
||||||
|
if n1.isLeader: inc leaderCount
|
||||||
|
if n2.isLeader: inc leaderCount
|
||||||
|
if leaderCount == 1: break
|
||||||
|
waitFor sleepAsync(100)
|
||||||
|
waited += 100
|
||||||
|
|
||||||
|
net1.stop()
|
||||||
|
net2.stop()
|
||||||
|
waitFor sleepAsync(50)
|
||||||
|
|
||||||
|
check leaderCount == 1
|
||||||
|
|
||||||
|
test "plaintext dial to a TLS raft port has no protocol effect":
|
||||||
|
let certDir = getTempDir() / "baradb_test_raft_tls"
|
||||||
|
let (certPath, keyPath) = generateSelfSignedCert(certDir, "raft-tls.local")
|
||||||
|
if certPath.len == 0:
|
||||||
|
skip() # openssl unavailable
|
||||||
|
else:
|
||||||
|
let tls = newTLSContext(newTLSConfig(certPath, keyPath))
|
||||||
|
var n = newRaftNode("srv", @["cli"], raftPort = 29311)
|
||||||
|
n.electionTimeout = 60000 # keep the server passive during the test
|
||||||
|
n.peerAddrs["cli"] = ("127.0.0.1", 29312)
|
||||||
|
let net = newRaftNetwork(n, tls)
|
||||||
|
asyncCheck net.run()
|
||||||
|
waitFor sleepAsync(50)
|
||||||
|
|
||||||
|
let termBefore = n.currentTerm
|
||||||
|
|
||||||
|
# A plaintext client sends a perfectly valid serialized raft frame; the
|
||||||
|
# bytes fail the TLS handshake, so nothing reaches the state machine.
|
||||||
|
let voteReq = RaftMessage(kind: rmkRequestVote, term: 42, senderId: "cli")
|
||||||
|
let data = serialize(voteReq)
|
||||||
|
var frame = newSeq[byte](4 + data.len)
|
||||||
|
frame[0] = byte(data.len shr 24)
|
||||||
|
frame[1] = byte(data.len shr 16)
|
||||||
|
frame[2] = byte(data.len shr 8)
|
||||||
|
frame[3] = byte(data.len)
|
||||||
|
for i in 0 ..< data.len:
|
||||||
|
frame[4 + i] = data[i]
|
||||||
|
|
||||||
|
let client = newAsyncSocket()
|
||||||
|
waitFor client.connect("127.0.0.1", Port(29311))
|
||||||
|
try:
|
||||||
|
waitFor client.send(cast[string](frame))
|
||||||
|
except CatchableError:
|
||||||
|
discard
|
||||||
|
waitFor sleepAsync(300)
|
||||||
|
|
||||||
|
# The server must have dropped the connection after the failed handshake.
|
||||||
|
var connectionDropped = false
|
||||||
|
try:
|
||||||
|
connectionDropped = (waitFor client.recv(1)).len == 0
|
||||||
|
except CatchableError:
|
||||||
|
connectionDropped = true
|
||||||
|
client.close()
|
||||||
|
net.stop()
|
||||||
|
waitFor sleepAsync(50)
|
||||||
|
|
||||||
|
check n.state == rsFollower
|
||||||
|
check n.currentTerm == termBefore
|
||||||
|
check n.votedFor == ""
|
||||||
|
check connectionDropped
|
||||||
|
|
||||||
suite "Raft SQL Write Path":
|
suite "Raft SQL Write Path":
|
||||||
test "leader append+commit wait round-trips through applyCommand":
|
test "leader append+commit wait round-trips through applyCommand":
|
||||||
proc scenario() =
|
proc scenario() =
|
||||||
|
|||||||
Reference in New Issue
Block a user