
# packages ----------------------------------------------------------------

library(tidyverse)
library(deSolve)
library(patchwork)

DataEVD <- read.csv('./EVD.csv')[,2:3]
names(DataEVD) <- c('date', 'cumlative_cases')
DataEVD <- DataEVD |> 
  mutate(date =as.Date(date)) |> 
  complete(date = seq.Date(min(date), max(date), by='day'),
           fill = list(cumlative_cases = NA)) |> 
  fill(cumlative_cases, .direction = 'down') |> 
  mutate(year = year(date),
         time = as.integer(date - min(date))) |> 
  left_join(data.frame(year = 2013:2016,
                       N = c(11055430, 11333365, 11625998, 11930985)),
            by = 'year')


# model -------------------------------------------------------------------

sir_model <- function(time, state, parms) {
  with(as.list(c(state, parms)), {
    N <- S + I + R
    dS <- -beta * S * I / N
    dC <- -dS
    dI <- beta * S * I / N - gamma * I
    dR <- gamma * I
    return(list(c(dS, dI, dC, dR)))
  })
}

state <- c(S = DataEVD$N[1],
           I = DataEVD$cumlative_cases[1],
           C = DataEVD$cumlative_cases[1],
           R = 0)
times <- DataEVD$time

# Sensitivity -------------------------------------------------------------

paradata <- expand.grid(
  gamma = 1/seq(9.4-5.5, 9.4+5.5, 0.5),
  r0 = seq(1.5, 2.5, 0.1)
) |> 
  mutate(beta = r0*gamma)
paradata$i <- as.integer(rownames(paradata))

## model function

sir_fun <- function(i){
  sir_out <- ode(y = state,
                 times = times,
                 func = sir_model,
                 parms = c(
                   beta = paradata$beta[i],
                   gamma = paradata$gamma[i]
                 ),
                 method = 'rk4') |> 
    as.data.frame() |> 
    left_join(DataEVD,
              by = c(time = 'time')) |> 
    arrange(time)
  
  cumu_cases <- max(sir_out$C)
  peak_cases <- max(diff(sir_out$C))
  duration <- max(which(diff(sir_out$C) >= 1))

  return(c(paradata$r0[i], paradata$gamma[i], cumu_cases, peak_cases, duration))
  # 
  # sir_out$r0 <- paradata$r0[i]
  # sir_out$beta <- paradata$beta[i]
  # sir_out$gamma <- paradata$gamma[i]
  # sir_out$i <- i
  # return(sir_out)
}

outcome <- lapply(paradata$i, sir_fun)
outcome <- do.call('rbind', outcome) |> 
  as.data.frame()
names(outcome) <- c('r0', 'gamma', 'cumu', 'peak', 'duration')

fig1 <- ggplot(data = outcome)+
  geom_contour_filled(
    mapping = aes(x = r0,
                  y = gamma,
                  z = cumu)
  )+
  scale_x_continuous(expand = c(0, 0))+
  scale_y_continuous(expand = c(0, 0))+
  scale_fill_viridis_d(option = 'D', direction = -1)+
  theme_classic()+
  labs(x = "Basic Reproductive Number",
       y = "Recovery Rate",
       fill = "Cumulative Cases",
       title = 'A')

fig2 <- ggplot(data = outcome)+
  geom_contour_filled(
    mapping = aes(x = r0,
                  y = gamma,
                  z = peak)
  )+
  scale_x_continuous(expand = c(0, 0))+
  scale_y_continuous(expand = c(0, 0))+
  scale_fill_viridis_d(option = 'A', direction = -1)+
  theme_classic()+
  labs(x = "Basic Reproductive Number",
       y = "Recovery Rate",
       fill = "Peak Incidence",
       title = 'B')

fig1 + fig2

ggsave(filename = 'figure 4.png',
       dpi = 300,
       width = 10,
       height = 4)







