use core::cell::Cell;
use std::io::{ErrorKind, Read, Write};
use std::net::{TcpStream, ToSocketAddrs};
use std::sync::Arc;
use std::sync::atomic::Ordering::{Relaxed, Release};
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64};
use std::time::{Duration, Instant};
use yo_common::lock::Lock;
use yo_common::{Code, Error, Result};
use crate::proto::{Limits, Proto};
use crate::reply::Out;
use crate::request::{Argv, Step};
use super::args::{self, Args};
use super::repl::{self, ID_LEN};
use super::{Flow, Server, Session};
const DIAL_TIMEOUT: Duration = Duration::from_secs(5);
const POLL: Duration = Duration::from_millis(100);
const ACK_EVERY: Duration = Duration::from_secs(1);
const RETRY: Duration = Duration::from_millis(500);
const HANDSHAKE_TIMEOUT: Duration = Duration::from_secs(10);
const SYNC_TIMEOUT: Duration = Duration::from_secs(60);
const LINE_MAX: usize = 1024;
#[derive(Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
enum State {
None = 0,
Connect = 1,
Sync = 2,
Up = 3,
}
impl State {
fn from(n: u8) -> State {
match n {
1 => State::Connect,
2 => State::Sync,
3 => State::Up,
_ => State::None,
}
}
}
#[derive(Clone)]
struct Upstream {
host: String,
port: u16,
}
pub(crate) struct Follower {
upstream: Lock<Option<Upstream>>,
epoch: AtomicU64,
state: AtomicU8,
read_only: AtomicBool,
on: AtomicBool,
auth: Lock<(Vec<u8>, Vec<u8>)>,
port: AtomicU64,
last_io_ms: AtomicU64,
down_ms: AtomicU64,
resume: AtomicBool,
}
impl Default for Follower {
fn default() -> Follower {
Follower {
upstream: Lock::new(None),
epoch: AtomicU64::new(0),
state: AtomicU8::new(State::None as u8),
read_only: AtomicBool::new(true),
on: AtomicBool::new(false),
auth: Lock::new((Vec::new(), Vec::new())),
port: AtomicU64::new(0),
last_io_ms: AtomicU64::new(0),
down_ms: AtomicU64::new(0),
resume: AtomicBool::new(false),
}
}
}
thread_local! {
static APPLYING: Cell<bool> = const { Cell::new(false) };
}
#[must_use]
pub(crate) fn applying() -> bool {
APPLYING.get()
}
impl Server {
#[must_use]
pub(crate) fn following(&self) -> bool {
self.follow.on.load(Relaxed)
}
#[must_use]
pub(crate) fn read_only_replica(&self) -> bool {
self.follow.on.load(Relaxed) && self.follow.read_only.load(Relaxed)
}
pub fn announce_port(&self, port: u16) {
self.follow.port.store(u64::from(port), Relaxed);
}
#[must_use]
pub(crate) fn announced_port(&self) -> u16 {
self.follow.port.load(Relaxed) as u16
}
pub fn master_auth(&self, user: &[u8], pass: &[u8]) {
let mut auth = self.follow.auth.lock();
yo_alloc::allow(|| *auth = (user.to_vec(), pass.to_vec()));
}
pub(crate) fn with_master_auth<T>(&self, each: impl FnOnce(&[u8], &[u8]) -> T) -> T {
let auth = self.follow.auth.lock();
each(&auth.0, &auth.1)
}
pub fn set_replica_read_only(&self, yes: bool) {
self.follow.read_only.store(yes, Relaxed);
}
pub(crate) fn replica_read_only_setting(&self) -> bool {
self.follow.read_only.load(Relaxed)
}
pub fn follow_master(self: &Arc<Server>, host: &str, port: u16) {
self.is_behind();
let host = yo_alloc::allow(|| host.to_owned());
self.follow_now(Some(Upstream { host, port }));
}
#[must_use]
pub(crate) fn master_link_up(&self) -> bool {
State::from(self.follow.state.load(Relaxed)) == State::Up
}
pub(super) fn stop_following(self: &Arc<Server>) {
if self.following() {
self.promote();
}
self.follow_now(None);
}
pub(super) fn follow_for_failover(self: &Arc<Server>, host: &str, port: u16) {
self.follow.resume.store(true, Relaxed);
let host = yo_alloc::allow(|| host.to_owned());
self.follow_now(Some(Upstream { host, port }));
}
fn follow_now(self: &Arc<Server>, to: Option<Upstream>) {
let epoch = self.follow.epoch.fetch_add(1, Relaxed) + 1;
{
let mut upstream = self.follow.upstream.lock();
yo_alloc::allow(|| *upstream = to.clone());
}
let Some(to) = to else {
self.follow.on.store(false, Release);
self.follow.state.store(State::None as u8, Relaxed);
self.follow.resume.store(false, Relaxed);
return;
};
self.follow.on.store(true, Release);
self.follow.state.store(State::Connect as u8, Relaxed);
self.follow.down_ms.store(self.clock.now_ms(), Relaxed);
let server = Arc::clone(self);
yo_alloc::allow(|| {
let _ = std::thread::Builder::new()
.name(String::from("yo-replica"))
.spawn(move || link(&server, epoch, &to));
});
}
}
#[cfg(test)]
impl Server {
pub(super) fn pretend_following(&self, host: &str, port: u16, up: bool) {
{
let mut upstream = self.follow.upstream.lock();
*upstream = Some(Upstream {
host: host.to_owned(),
port,
});
}
self.follow.on.store(true, Release);
self.follow
.state
.store(if up { State::Up } else { State::Connect } as u8, Relaxed);
self.follow.last_io_ms.store(self.clock.now_ms(), Relaxed);
self.follow.down_ms.store(self.clock.now_ms(), Relaxed);
}
pub(super) fn pretend_master(&self) {
self.follow.on.store(false, Release);
self.follow.state.store(State::None as u8, Relaxed);
}
}
pub(super) fn replicaof(server: &Server, args: Args<'_>, out: &mut Out) -> Result<()> {
let name = if args.name().eq_ignore_ascii_case(b"slaveof") {
"slaveof"
} else {
"replicaof"
};
if args.len() != 3 {
return Err(args::wrong_arity(name));
}
if server.failing_over() {
return Err(Error::new(
Code::Invalid,
"REPLICAOF not allowed while failing over.",
));
}
let host = args.get(1);
let port = args.get(2);
let told = if host.eq_ignore_ascii_case(b"no") && port.eq_ignore_ascii_case(b"one") {
None
} else {
let Some(port) = core::str::from_utf8(port)
.ok()
.and_then(|w| w.parse::<u16>().ok())
else {
return Err(Error::new(Code::Invalid, "Invalid master port"));
};
Some(port)
};
let Some(shared) = server.myself() else {
return Err(Error::new(
Code::Invalid,
"REPLICAOF is not available on an embedded server",
));
};
let Some(port) = told else {
shared.stop_following();
out.ok();
return Ok(());
};
let host = yo_alloc::allow(|| String::from_utf8_lossy(host).into_owned());
{
let upstream = server.follow.upstream.lock();
let same = upstream
.as_ref()
.is_some_and(|at| at.port == port && at.host == host);
if same && server.following() {
out.simple(
b"OK REPLICAOF would result into synchronization with the master we are already connected with. No operation performed.",
);
return Ok(());
}
}
shared.follow_now(Some(Upstream { host, port }));
out.ok();
Ok(())
}
pub(super) const READONLY: &str = "READONLY You can't write against a read only replica.";
pub(super) fn info(server: &Server, s: &mut String) {
use core::fmt::Write as _;
let state = State::from(server.follow.state.load(Relaxed));
if state == State::None {
return;
}
let (host, port) = {
let upstream = server.follow.upstream.lock();
match upstream.as_ref() {
Some(up) => (up.host.clone(), up.port),
None => (String::new(), 0),
}
};
let now = server.clock.now_ms();
let up = state == State::Up;
let last = now.saturating_sub(server.follow.last_io_ms.load(Relaxed)) / 1000;
let offset = server.repl_offset();
let _ = write!(
s,
"master_host:{host}\r\nmaster_port:{port}\r\n\
master_link_status:{}\r\nmaster_last_io_seconds_ago:{}\r\n\
master_sync_in_progress:{}\r\n\
slave_read_repl_offset:{offset}\r\nslave_repl_offset:{offset}\r\n",
if up { "up" } else { "down" },
if up { last as i64 } else { -1 },
usize::from(state == State::Sync),
);
if !up {
let down = now.saturating_sub(server.follow.down_ms.load(Relaxed)) / 1000;
let _ = write!(s, "master_link_down_since_seconds:{down}\r\n");
}
let _ = write!(
s,
"slave_priority:100\r\nslave_read_only:{}\r\nreplica_announced:1\r\n",
usize::from(server.follow.read_only.load(Relaxed)),
);
}
#[must_use]
pub(super) fn role_word(server: &Server) -> &'static str {
if server.following() {
"slave"
} else {
"master"
}
}
pub(super) fn role(server: &Server, out: &mut Out) {
let (host, port) = {
let upstream = server.follow.upstream.lock();
match upstream.as_ref() {
Some(up) => (up.host.clone(), up.port),
None => (String::new(), 0),
}
};
out.array(5);
out.bulk(b"slave");
out.bulk(host.as_bytes());
out.int(i64::from(port));
out.bulk(match State::from(server.follow.state.load(Relaxed)) {
State::Up => b"connected".as_slice(),
State::Sync => b"sync".as_slice(),
_ => b"connect".as_slice(),
});
out.int(server.repl_offset() as i64);
}
fn link(server: &Arc<Server>, epoch: u64, to: &Upstream) {
while server.follow.epoch.load(Relaxed) == epoch && !server.stopping() {
let _ = once(server, epoch, to);
if server.follow.epoch.load(Relaxed) != epoch {
return;
}
if State::from(server.follow.state.load(Relaxed)) != State::Connect {
server.follow.state.store(State::Connect as u8, Relaxed);
server.follow.down_ms.store(server.clock.now_ms(), Relaxed);
}
std::thread::sleep(RETRY);
}
}
fn once(server: &Arc<Server>, epoch: u64, to: &Upstream) -> std::io::Result<()> {
let mut wire = dial(to)?;
handshake(server, &mut wire)?;
let handing_over = server.failover_stage() == super::failover::Stage::InProgress;
if server.follow.resume.load(Relaxed) {
let id = server.repl_id();
let from = (server.repl_offset() + 1).to_string();
if handing_over {
wire.write(&[b"PSYNC", &id, from.as_bytes(), b"FAILOVER"])?;
} else {
wire.write(&[b"PSYNC", &id, from.as_bytes()])?;
}
} else {
wire.write(&[b"PSYNC", b"?", b"-1"])?;
}
let head = wire.line(SYNC_TIMEOUT)?;
if head.starts_with(b"+FULLRESYNC ") {
full_resync(server, &mut wire, &head[12..])?;
server.follow.state.store(State::Up as u8, Relaxed);
server.follow.resume.store(true, Relaxed);
} else if head.starts_with(b"+CONTINUE") {
if let Some(id) = head.get(10..).and_then(fixed_id) {
server.adopt(id, server.repl_offset());
}
server.follow.state.store(State::Up as u8, Relaxed);
server.follow.resume.store(true, Relaxed);
} else {
if handing_over {
super::failover::abort(server);
}
return Err(broken("the master would not resynchronise"));
}
super::failover::landed(server);
stream(server, epoch, &mut wire)
}
fn full_resync(server: &Arc<Server>, wire: &mut Link, head: &[u8]) -> std::io::Result<()> {
let mut words = head.split(|&b| b == b' ');
let id = words
.next()
.and_then(fixed_id)
.ok_or_else(|| broken("the master named no replication id"))?;
let offset = words
.next()
.and_then(|w| core::str::from_utf8(w).ok())
.and_then(|w| w.trim().parse::<u64>().ok())
.ok_or_else(|| broken("the master named no offset"))?;
server.follow.state.store(State::Sync as u8, Relaxed);
let image = wire.payload()?;
server
.load_image(&image, true)
.map_err(|e| broken(&format!("the snapshot would not load: {e}")))?;
server.adopt(id, offset);
Ok(())
}
fn stream(server: &Arc<Server>, epoch: u64, wire: &mut Link) -> std::io::Result<()> {
let mut session = Session::new(server.next_client());
session.admit(true);
session.serve_master(true);
let mut out = Out::new(Proto::Resp2);
let mut argv = Argv::new();
let limits = Limits::default();
let mut acked = Instant::now();
let mut sent = 0u64;
APPLYING.set(true);
let ended = loop {
if server.follow.epoch.load(Relaxed) != epoch || server.stopping() {
break Ok(());
}
match argv.decode(wire.held(), &limits) {
Err(_) => break Err(broken("the master sent something that is not a command")),
Ok(Step::Incomplete) => {
if let Err(e) = wire.fill(POLL) {
break Err(e);
}
}
Ok(Step::Command { consumed }) => {
let getack = is_getack(&argv, wire.held());
if !getack {
apply(server, &mut session, &argv, wire.held(), &mut out);
}
let bytes = wire.take(consumed);
repl::relayed(server, bytes, session.db());
server
.follow
.last_io_ms
.store(server.clock.now_ms(), Relaxed);
if getack {
acked = Instant::now();
sent = server.repl_offset();
if let Err(e) = wire.ack(sent) {
break Err(e);
}
}
continue;
}
}
let now = server.repl_offset();
if acked.elapsed() >= ACK_EVERY || now != sent {
acked = Instant::now();
sent = now;
if let Err(e) = wire.ack(now) {
break Err(e);
}
}
};
APPLYING.set(false);
super::forget_session(server, &mut session);
ended
}
fn is_getack(argv: &Argv, buf: &[u8]) -> bool {
argv.len() == 3
&& argv
.arg(buf, 0)
.is_some_and(|w| w.eq_ignore_ascii_case(b"replconf"))
&& argv
.arg(buf, 1)
.is_some_and(|w| w.eq_ignore_ascii_case(b"getack"))
}
fn apply(server: &Server, session: &mut Session, argv: &Argv, buf: &[u8], out: &mut Out) {
loop {
out.clear();
let args = Args::new(argv, buf);
if super::execute(server, session, args, out) != Flow::Hold {
return;
}
std::thread::sleep(Duration::from_millis(1));
}
}
struct Link {
sock: TcpStream,
buf: Vec<u8>,
}
impl Link {
fn held(&self) -> &[u8] {
&self.buf
}
fn take(&mut self, n: usize) -> Vec<u8> {
self.buf.drain(..n).collect()
}
fn fill(&mut self, wait: Duration) -> std::io::Result<()> {
self.sock.set_read_timeout(Some(wait))?;
let mut chunk = [0u8; 16 * 1024];
match self.sock.read(&mut chunk) {
Ok(0) => Err(broken("the master closed the link")),
Ok(n) => {
self.buf.extend_from_slice(&chunk[..n]);
Ok(())
}
Err(e) if soft(&e) => Ok(()),
Err(e) => Err(e),
}
}
fn line(&mut self, wait: Duration) -> std::io::Result<Vec<u8>> {
let until = Instant::now() + wait;
loop {
if let Some(at) = self.buf.iter().position(|&b| b == b'\n') {
let mut line = self.take(at + 1);
while line.last().is_some_and(|&b| b == b'\n' || b == b'\r') {
line.pop();
}
if line.is_empty() {
continue;
}
return Ok(line);
}
if self.buf.len() > LINE_MAX {
return Err(broken("the master sent a line with no end to it"));
}
if Instant::now() >= until {
return Err(broken("the master did not answer"));
}
self.fill(POLL)?;
}
}
fn command(&mut self, parts: &[&[u8]]) -> std::io::Result<Vec<u8>> {
self.write(parts)?;
self.line(HANDSHAKE_TIMEOUT)
}
fn write(&mut self, parts: &[&[u8]]) -> std::io::Result<()> {
let mut wire = Vec::with_capacity(32);
wire.extend_from_slice(b"*");
wire.extend_from_slice(parts.len().to_string().as_bytes());
wire.extend_from_slice(b"\r\n");
for part in parts {
wire.extend_from_slice(b"$");
wire.extend_from_slice(part.len().to_string().as_bytes());
wire.extend_from_slice(b"\r\n");
wire.extend_from_slice(part);
wire.extend_from_slice(b"\r\n");
}
self.sock.write_all(&wire)
}
fn ack(&mut self, offset: u64) -> std::io::Result<()> {
self.write(&[b"REPLCONF", b"ACK", offset.to_string().as_bytes()])
}
fn payload(&mut self) -> std::io::Result<Vec<u8>> {
let head = self.line(SYNC_TIMEOUT)?;
let want = core::str::from_utf8(head.get(1..).unwrap_or_default())
.ok()
.and_then(|n| n.parse::<usize>().ok())
.filter(|_| head.first() == Some(&b'$'))
.ok_or_else(|| broken("the master did not say how long the snapshot is"))?;
let mut last = Instant::now();
while self.buf.len() < want {
let had = self.buf.len();
self.fill(POLL)?;
if self.buf.len() > had {
last = Instant::now();
} else if last.elapsed() >= SYNC_TIMEOUT {
return Err(broken("the master stopped part way through the snapshot"));
}
}
Ok(self.take(want))
}
}
fn soft(e: &std::io::Error) -> bool {
matches!(
e.kind(),
ErrorKind::WouldBlock | ErrorKind::TimedOut | ErrorKind::Interrupted
)
}
fn broken(why: &str) -> std::io::Error {
std::io::Error::other(String::from(why))
}
fn fixed_id(word: &[u8]) -> Option<[u8; ID_LEN]> {
let word = word.strip_suffix(b"\r").unwrap_or(word);
let word = word.get(..ID_LEN)?;
word.iter()
.all(u8::is_ascii_hexdigit)
.then(|| <[u8; ID_LEN]>::try_from(word).ok())
.flatten()
}
fn dial(to: &Upstream) -> std::io::Result<Link> {
let at = (to.host.as_str(), to.port)
.to_socket_addrs()?
.next()
.ok_or_else(|| broken("the master's address does not resolve"))?;
let sock = TcpStream::connect_timeout(&at, DIAL_TIMEOUT)?;
sock.set_nodelay(true)?;
Ok(Link {
sock,
buf: Vec::new(),
})
}
fn handshake(server: &Server, wire: &mut Link) -> std::io::Result<()> {
wire.command(&[b"PING"])?;
let (user, pass) = {
let auth = server.follow.auth.lock();
auth.clone()
};
if !pass.is_empty() {
let said = if user.is_empty() {
wire.command(&[b"AUTH", &pass])?
} else {
wire.command(&[b"AUTH", &user, &pass])?
};
if said.first() == Some(&b'-') {
return Err(broken("the master would not take the password"));
}
}
let port = server.follow.port.load(Relaxed).to_string();
wire.command(&[b"REPLCONF", b"listening-port", port.as_bytes()])?;
wire.command(&[b"REPLCONF", b"capa", b"psync2"])?;
server
.follow
.last_io_ms
.store(server.clock.now_ms(), Relaxed);
Ok(())
}
#[cfg(test)]
mod tests {
use super::{ID_LEN, State, fixed_id, soft};
#[test]
fn a_replication_id_is_forty_hex_characters_and_nothing_else() {
let good = b"0123456789abcdef0123456789abcdef01234567";
assert_eq!(fixed_id(good).unwrap(), *good);
let mut with_cr = good.to_vec();
with_cr.push(b'\r');
assert_eq!(fixed_id(&with_cr).unwrap(), *good);
let mut longer = good.to_vec();
longer.extend_from_slice(b"more");
assert_eq!(fixed_id(&longer).unwrap(), *good);
assert!(fixed_id(&good[..ID_LEN - 1]).is_none());
let mut wrong = good.to_vec();
wrong[7] = b'z';
assert!(fixed_id(&wrong).is_none());
assert!(fixed_id(b"").is_none());
}
#[test]
fn a_state_that_is_not_one_of_the_four_reads_as_no_master() {
assert!(State::from(1) == State::Connect);
assert!(State::from(2) == State::Sync);
assert!(State::from(3) == State::Up);
assert!(State::from(0) == State::None);
assert!(State::from(99) == State::None);
}
#[test]
fn a_read_that_timed_out_is_not_a_broken_link() {
use std::io::{Error, ErrorKind};
assert!(soft(&Error::from(ErrorKind::WouldBlock)));
assert!(soft(&Error::from(ErrorKind::TimedOut)));
assert!(soft(&Error::from(ErrorKind::Interrupted)));
assert!(!soft(&Error::from(ErrorKind::ConnectionReset)));
assert!(!soft(&Error::from(ErrorKind::UnexpectedEof)));
}
}