# Load necessary libraries
library(dplyr)
library(MASS)  # For Mahalanobis distance calculation

# Create the data
data <- data.frame(
  person = c("Andy", "Betty", "Chad", "Doug", "Edith", "Fred", "Gina", "Hank", "Inez", "Janet"),
  y1 = c(200, 250, 150, 300, NA, NA, NA, NA, NA, NA),
  y0 = c(NA, NA, NA, NA, 225, 500, 200, 190, 180, 140),
  d = c(1, 1, 1, 1, 0, 0, 0, 0, 0, 0),
  age = c(23, 27, 29, 22, 27, 32, 26, 26, 28, 29)
)

# Add unit ID and create earnings variable
data <- data %>%
  dplyr::mutate(
    unit = row_number(),
    earnings = ifelse(d == 1, y1, y0)
  )

# Split data into treated and control groups
treated <- data %>% dplyr::filter(d == 1)
control <- data %>% dplyr::filter(d == 0)

# For a single variable, we'll use variance instead of covariance
var_age <- var(data$age)

# Initialize matrices to store distances and matches
distance_matrix <- matrix(NA, nrow = nrow(treated), ncol = nrow(control))
match1 <- numeric(nrow(treated))
match2 <- numeric(nrow(treated))

# Calculate distances (scaled by variance)
for (i in 1:nrow(treated)) {
  for (j in 1:nrow(control)) {
    distance_matrix[i, j] <- (treated$age[i] - control$age[j])^2 / var_age
  }
}

# Find matches (allowing for ties)
for (i in 1:nrow(treated)) {
  sorted_indices <- order(distance_matrix[i,])
  match1[i] <- sorted_indices[1]
  
  # Check if second-best match is equally good
  if (abs(distance_matrix[i, sorted_indices[1]] - 
          distance_matrix[i, sorted_indices[2]]) < 1e-10) {
    match2[i] <- sorted_indices[2]
  }
}

# Calculate matched outcomes
matched_data <- treated %>%
  dplyr::mutate(
    match11 = match1 + nrow(treated),  # Adjust indices to match Stata output
    match12 = ifelse(match2 > 0, match2 + nrow(treated), NA),
    y0_match = ifelse(
      !is.na(match12),
      (control$earnings[match1] + control$earnings[match2]) / 2,
      control$earnings[match1]
    ),
    te = earnings - y0_match
  )

# Display results
print("Matched Data:")
print(as.data.frame(matched_data)[, c("unit", "person", "y1", "y0", "d", "age", 
                                      "earnings", "match11", "match12", "y0_match", "te")])

# Calculate ATT
att <- mean(matched_data$te)
print(paste("ATT =", round(att, 2)))

# Summary statistics of treatment effects
print("Summary of Treatment Effects:")
print(summary(matched_data$te))