use std::sync::Arc;
use std::sync::atomic::Ordering::Relaxed;
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicU64};
use std::time::Duration;
use yo_common::lock::Lock;
use yo_common::{Code, Error, Result};
use crate::reply::Out;
use super::Server;
use super::args::{self, Args};
const WATCH: Duration = Duration::from_millis(10);
const FOREVER: u64 = u64::MAX >> 1;
#[derive(Clone, Copy, PartialEq, Eq)]
#[repr(u8)]
pub(super) enum Stage {
None = 0,
WaitForSync = 1,
InProgress = 2,
}
impl Stage {
fn from(n: u8) -> Stage {
match n {
1 => Stage::WaitForSync,
2 => Stage::InProgress,
_ => Stage::None,
}
}
fn word(self) -> &'static str {
match self {
Stage::None => "no-failover",
Stage::WaitForSync => "waiting-for-sync",
Stage::InProgress => "failover-in-progress",
}
}
}
#[derive(Default)]
pub(crate) struct Failover {
stage: AtomicU8,
target: Lock<Option<(String, u16)>>,
end_ms: AtomicU64,
force: AtomicBool,
epoch: AtomicU64,
}
impl Server {
#[must_use]
pub(crate) fn failing_over(&self) -> bool {
self.failover.stage.load(Relaxed) != Stage::None as u8
}
#[must_use]
pub(super) fn failover_stage(&self) -> Stage {
Stage::from(self.failover.stage.load(Relaxed))
}
#[must_use]
pub(super) fn failover_word(&self) -> &'static str {
self.failover_stage().word()
}
}
fn clear(server: &Server) {
server.failover.epoch.fetch_add(1, Relaxed);
server.failover.end_ms.store(0, Relaxed);
server.failover.force.store(false, Relaxed);
{
let mut target = server.failover.target.lock();
yo_alloc::allow(|| *target = None);
}
server.failover.stage.store(Stage::None as u8, Relaxed);
server.unpause();
}
pub(super) fn landed(server: &Server) {
if server.failover_stage() == Stage::InProgress {
clear(server);
}
}
pub(super) fn abort(server: &Arc<Server>) {
if server.failover_stage() == Stage::InProgress {
server.stop_following();
}
clear(server);
}
pub(super) fn execute(server: &Server, args: Args<'_>, out: &mut Out) -> Result<()> {
if args.len() == 2 && args.get(1).eq_ignore_ascii_case(b"abort") {
if !server.failing_over() {
return Err(Error::new(Code::Invalid, "No failover in progress."));
}
let Some(shared) = server.myself() else {
return Err(embedded());
};
abort(&shared);
out.ok();
return Ok(());
}
let mut timeout = 0i64;
let mut force = false;
let mut target: Option<(&[u8], i64)> = None;
let mut at = 1;
while at < args.len() {
let word = args.get(at);
if word.eq_ignore_ascii_case(b"timeout") && at + 1 < args.len() && timeout == 0 {
timeout = args.int(at + 1)?;
if timeout <= 0 {
return Err(Error::new(
Code::Invalid,
"FAILOVER timeout must be greater than 0",
));
}
at += 2;
} else if word.eq_ignore_ascii_case(b"to") && at + 2 < args.len() && target.is_none() {
let port = args.int(at + 2)?;
target = Some((args.get(at + 1), port));
at += 3;
} else if word.eq_ignore_ascii_case(b"force") && !force {
force = true;
at += 1;
} else {
return Err(args::syntax());
}
}
if server.failing_over() {
return Err(Error::new(Code::Invalid, "FAILOVER already in progress."));
}
if server.following() {
return Err(Error::new(
Code::Invalid,
"FAILOVER is not valid when server is a replica.",
));
}
if server.replica_count() == 0 {
return Err(Error::new(
Code::Invalid,
"FAILOVER requires connected replicas.",
));
}
if force && (timeout == 0 || target.is_none()) {
return Err(Error::new(
Code::Invalid,
"FAILOVER with force option requires both a timeout and target HOST and IP.",
));
}
let named = match target {
Some((host, port)) => {
let host = String::from_utf8_lossy(host).into_owned();
let port = u16::try_from(port).unwrap_or(0);
match server.replica_online_at(&host, port) {
None => {
return Err(Error::new(
Code::Invalid,
"FAILOVER target HOST and PORT is not a replica.",
));
}
Some(false) => {
return Err(Error::new(
Code::Invalid,
"FAILOVER target replica is not online.",
));
}
Some(true) => Some((host, port)),
}
}
None => None,
};
let Some(shared) = server.myself() else {
return Err(embedded());
};
{
let mut held = server.failover.target.lock();
yo_alloc::allow(|| *held = named);
}
if timeout > 0 {
let end = server.clock.now_ms() + timeout as u64;
server.failover.end_ms.store(end, Relaxed);
}
server.failover.force.store(force, Relaxed);
server
.failover
.stage
.store(Stage::WaitForSync as u8, Relaxed);
server.pause(FOREVER, false);
let epoch = server.failover.epoch.load(Relaxed);
let watching = Arc::clone(&shared);
yo_alloc::allow(|| {
let _ = std::thread::Builder::new()
.name(String::from("yo-failover"))
.spawn(move || watch(&watching, epoch));
});
out.ok();
Ok(())
}
fn embedded() -> Error {
Error::new(
Code::Invalid,
"FAILOVER is not available on an embedded server",
)
}
fn watch(server: &Arc<Server>, epoch: u64) {
while server.failover.epoch.load(Relaxed) == epoch
&& server.failover_stage() == Stage::WaitForSync
&& !server.stopping()
{
step(server);
std::thread::sleep(WATCH);
}
}
fn step(server: &Arc<Server>) {
let end = server.failover.end_ms.load(Relaxed);
if end != 0 && end <= server.clock.now_ms() {
if server.failover.force.load(Relaxed) {
hand_over(server);
} else {
abort(server);
}
return;
}
let named = { server.failover.target.lock().clone() };
match named {
Some((host, port)) => {
if server.replica_caught_up(&host, port) {
hand_over(server);
}
}
None => {
if let Some(at) = server.first_caught_up() {
{
let mut held = server.failover.target.lock();
yo_alloc::allow(|| *held = Some(at));
}
hand_over(server);
}
}
}
}
fn hand_over(server: &Arc<Server>) {
let Some((host, port)) = server.failover.target.lock().clone() else {
return;
};
server
.failover
.stage
.store(Stage::InProgress as u8, Relaxed);
server.follow_for_failover(&host, port);
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_three_stages_have_the_words_redis_prints() {
assert_eq!(Stage::None.word(), "no-failover");
assert_eq!(Stage::WaitForSync.word(), "waiting-for-sync");
assert_eq!(Stage::InProgress.word(), "failover-in-progress");
}
#[test]
fn a_stage_that_is_not_one_of_the_three_reads_as_no_failover() {
assert!(Stage::from(0) == Stage::None);
assert!(Stage::from(1) == Stage::WaitForSync);
assert!(Stage::from(2) == Stage::InProgress);
assert!(Stage::from(9) == Stage::None);
}
#[test]
fn a_server_nobody_is_failing_over_says_so() {
let server = Server::new();
assert!(!server.failing_over());
assert_eq!(server.failover_word(), "no-failover");
}
}