# Module 4.3: SIR disease transmission
# Adapted from the earlier course SIR.R example by Stephen Davies,
# University of Mary Washington, accompanying Shiflet and Shiflet.
# This formulation uses beta * S * I, without division by N.

delta_t <- 0.01                  # days per step
end_time <- 28                   # days
time <- seq(0, end_time, by = delta_t)

transmission_coefficient <- 0.00218  # per person per day
recovery_rate <- 0.5                # per day

initial_susceptible <- 762
initial_infectious <- 1
initial_recovered <- 0
total_population <- initial_susceptible +
  initial_infectious + initial_recovered

susceptible <- numeric(length(time))
infectious <- numeric(length(time))
recovered <- numeric(length(time))
susceptible[1] <- initial_susceptible
infectious[1] <- initial_infectious
recovered[1] <- initial_recovered

for (i in 2:length(time)) {
  susceptible_old <- susceptible[i - 1]
  infectious_old <- infectious[i - 1]
  recovered_old <- recovered[i - 1]

  infection_flow <- transmission_coefficient *
    susceptible_old * infectious_old
  recovery_flow <- recovery_rate * infectious_old

  susceptible[i] <- susceptible_old - infection_flow * delta_t
  infectious[i] <- infectious_old +
    (infection_flow - recovery_flow) * delta_t
  recovered[i] <- recovered_old + recovery_flow * delta_t
}

# Verify the first step and total population.
results <- data.frame(time, susceptible, infectious, recovered)
results$total <- results$susceptible +
  results$infectious + results$recovered
head(results)
range(results$total)
max(abs(results$total - total_population))

plot(
  time, susceptible,
  type = "l", lwd = 3, col = "blue",
  ylim = c(0, total_population),
  xlab = "Time (days)", ylab = "People",
  main = "A simulated SIR outbreak"
)
lines(time, infectious, lwd = 3, col = "red")
lines(time, recovered, lwd = 3, lty = 2, col = "darkgreen")
legend(
  "right",
  legend = c("Susceptible", "Infectious", "Recovered"),
  col = c("blue", "red", "darkgreen"),
  lty = c(1, 1, 2), lwd = c(3, 3, 3), bty = "n"
)

# Peak and cumulative infections by day 28
peak_position <- which.max(results$infectious)
results[peak_position, ]
last_position <- length(time)
ever_infected <- infectious[last_position] + recovered[last_position]
fraction_infected <- ever_infected / total_population
ever_infected
fraction_infected

# Thresholds and reproduction numbers for beta * S * I
susceptible_threshold <- recovery_rate / transmission_coefficient
basic_reproduction_number <- transmission_coefficient *
  total_population / recovery_rate
results$effective_reproduction_number <- transmission_coefficient *
  results$susceptible / recovery_rate
susceptible_threshold
basic_reproduction_number
results[c(1, peak_position, last_position), ]

# Experiment: halve transmission from the beginning.
reduced_transmission <- transmission_coefficient / 2
reduced_transmission * initial_susceptible / recovery_rate

reduced_susceptible <- numeric(length(time))
reduced_infectious <- numeric(length(time))
reduced_recovered <- numeric(length(time))
reduced_susceptible[1] <- initial_susceptible
reduced_infectious[1] <- initial_infectious
reduced_recovered[1] <- initial_recovered

for (i in 2:length(time)) {
  susceptible_old <- reduced_susceptible[i - 1]
  infectious_old <- reduced_infectious[i - 1]
  recovered_old <- reduced_recovered[i - 1]

  infection_flow <- reduced_transmission *
    susceptible_old * infectious_old
  recovery_flow <- recovery_rate * infectious_old

  reduced_susceptible[i] <- susceptible_old - infection_flow * delta_t
  reduced_infectious[i] <- infectious_old +
    (infection_flow - recovery_flow) * delta_t
  reduced_recovered[i] <- recovered_old + recovery_flow * delta_t
}

plot(
  time, infectious,
  type = "l", lwd = 3, col = "red",
  ylim = c(0, max(infectious, reduced_infectious)),
  xlab = "Time (days)", ylab = "Infectious people",
  main = "Comparing transmission rates"
)
lines(time, reduced_infectious, lwd = 3, lty = 2, col = "blue")
legend(
  "topright",
  legend = c("Original transmission", "Half the transmission coefficient"),
  col = c("red", "blue"), lty = c(1, 2), lwd = c(3, 3), bty = "n"
)

reduced_peak_position <- which.max(reduced_infectious)
data.frame(
  scenario = c("Original", "Reduced transmission"),
  peak_infectious = c(infectious[peak_position],
    reduced_infectious[reduced_peak_position]),
  peak_day = c(time[peak_position], time[reduced_peak_position]),
  infected_by_day_28 = c(ever_infected,
    reduced_infectious[last_position] + reduced_recovered[last_position])
)
