///|
/// A durable-state vocabulary for resumable batch de-identification.
pub(all) enum CheckpointState {
  Pending
  Processing
  Succeeded
  Failed
  Skipped
} derive(Debug, Eq)

///|
pub(all) struct CheckpointRecord {
  item_id : String
  state : CheckpointState
  attempts : Int
  input_checksum : String
  output_checksum : String
  finding_count : Int
  error_message : String
  updated_sequence : Int
} derive(Debug, Eq)

///|
pub(all) struct BatchCheckpoint {
  mut batch_id : String
  mut policy_checksum : String
  records : Map[String, CheckpointRecord]
  mut sequence : Int
} derive(Debug)

///|
pub(all) struct ResumePlan {
  batch_id : String
  pending : Array[BatchItem]
  retryable : Array[BatchItem]
  completed : Int
  skipped : Int
  checksum : String
} derive(Debug)

///|
pub fn checkpoint_state_name(state : CheckpointState) -> String {
  match state {
    Pending => "pending"
    Processing => "processing"
    Succeeded => "succeeded"
    Failed => "failed"
    Skipped => "skipped"
  }
}

///|
pub fn BatchCheckpoint::new(
  batch_id : String,
  policy_checksum : String,
) -> BatchCheckpoint {
  { batch_id, policy_checksum, records: Map([]), sequence: 0 }
}

///|
pub fn BatchCheckpoint::empty(batch_id : String) -> BatchCheckpoint {
  BatchCheckpoint::new(batch_id, "")
}

///|
pub fn BatchCheckpoint::size(self : BatchCheckpoint) -> Int {
  self.records.length()
}

///|
pub fn BatchCheckpoint::get(
  self : BatchCheckpoint,
  item_id : String,
) -> CheckpointRecord? {
  self.records.get(item_id)
}

///|
pub fn BatchCheckpoint::contains(
  self : BatchCheckpoint,
  item_id : String,
) -> Bool {
  self.records.contains(item_id)
}

///|
pub fn BatchCheckpoint::next_sequence(self : BatchCheckpoint) -> Int {
  self.sequence + 1
}

///|
fn checkpoint_record(
  checkpoint : BatchCheckpoint,
  item_id : String,
  state : CheckpointState,
  input_checksum : String,
  output_checksum : String,
  finding_count : Int,
  error_message : String,
) -> CheckpointRecord {
  {
    item_id,
    state,
    attempts: match checkpoint.records.get(item_id) {
      Some(previous) => previous.attempts
      None => 0
    },
    input_checksum,
    output_checksum,
    finding_count,
    error_message,
    updated_sequence: checkpoint.next_sequence(),
  }
}

///|
pub fn BatchCheckpoint::mark_pending(
  self : BatchCheckpoint,
  item : BatchItem,
) -> Unit {
  self.sequence += 1
  self.records[item.id] = checkpoint_record(
    self,
    item.id,
    Pending,
    stable_hash(item.text),
    "",
    0,
    "",
  )
}

///|
pub fn BatchCheckpoint::mark_processing(
  self : BatchCheckpoint,
  item : BatchItem,
) -> Unit {
  self.sequence += 1
  let existing = checkpoint_record(
    self,
    item.id,
    Processing,
    stable_hash(item.text),
    "",
    0,
    "",
  )
  self.records[item.id] = { ..existing, attempts: existing.attempts + 1 }
}

///|
pub fn BatchCheckpoint::mark_succeeded(
  self : BatchCheckpoint,
  item : BatchItem,
  result : DeidResult,
) -> Unit {
  self.sequence += 1
  self.records[item.id] = checkpoint_record(
    self,
    item.id,
    Succeeded,
    stable_hash(item.text),
    stable_hash(result.text),
    result.findings.length(),
    "",
  )
}

///|
pub fn BatchCheckpoint::mark_failed(
  self : BatchCheckpoint,
  item : BatchItem,
  message : String,
) -> Unit {
  self.sequence += 1
  let record = checkpoint_record(
    self,
    item.id,
    Failed,
    stable_hash(item.text),
    "",
    0,
    message,
  )
  self.records[item.id] = { ..record, attempts: record.attempts + 1 }
}

///|
pub fn BatchCheckpoint::mark_skipped(
  self : BatchCheckpoint,
  item : BatchItem,
  reason : String,
) -> Unit {
  self.sequence += 1
  self.records[item.id] = checkpoint_record(
    self,
    item.id,
    Skipped,
    stable_hash(item.text),
    "",
    0,
    reason,
  )
}

///|
pub fn BatchCheckpoint::state(
  self : BatchCheckpoint,
  item_id : String,
) -> CheckpointState {
  match self.records.get(item_id) {
    Some(record) => record.state
    None => Pending
  }
}

///|
pub fn BatchCheckpoint::is_complete(
  self : BatchCheckpoint,
  item_id : String,
) -> Bool {
  match self.records.get(item_id) {
    Some(record) => record.state == Succeeded || record.state == Skipped
    None => false
  }
}

///|
pub fn BatchCheckpoint::is_retryable(
  self : BatchCheckpoint,
  item_id : String,
  max_attempts : Int,
) -> Bool {
  match self.records.get(item_id) {
    Some(record) =>
      record.state == Failed &&
      (max_attempts <= 0 || record.attempts < max_attempts)
    None => false
  }
}

///|
pub fn BatchCheckpoint::records_by_state(
  self : BatchCheckpoint,
  state : CheckpointState,
) -> Array[CheckpointRecord] {
  self.records.values().filter(fn(record) { record.state == state }).to_array()
}

///|
pub fn BatchCheckpoint::successful_ids(self : BatchCheckpoint) -> Array[String] {
  self.records_by_state(Succeeded).map(fn(record) { record.item_id })
}

///|
pub fn BatchCheckpoint::failed_ids(self : BatchCheckpoint) -> Array[String] {
  self.records_by_state(Failed).map(fn(record) { record.item_id })
}

///|
pub fn BatchCheckpoint::pending_ids(self : BatchCheckpoint) -> Array[String] {
  self.records_by_state(Pending).map(fn(record) { record.item_id })
}

///|
pub fn BatchCheckpoint::state_counts(
  self : BatchCheckpoint,
) -> Map[String, Int] {
  let counts : Map[String, Int] = Map([])
  for record in self.records.values() {
    let key = checkpoint_state_name(record.state)
    counts[key] = counts.get_or_default(key, 0) + 1
  }
  counts
}

///|
pub fn BatchCheckpoint::attempt_count(self : BatchCheckpoint) -> Int {
  self.records.values().fold(init=0, (sum, record) => sum + record.attempts)
}

///|
pub fn BatchCheckpoint::finding_count(self : BatchCheckpoint) -> Int {
  self.records
  .values()
  .fold(init=0, (sum, record) => sum + record.finding_count)
}

///|
pub fn BatchCheckpoint::checksum(self : BatchCheckpoint) -> String {
  let lines = [self.batch_id, self.policy_checksum]
  for id, record in self.records {
    lines.push(
      [
        id,
        checkpoint_state_name(record.state),
        record.input_checksum,
        record.output_checksum,
        "\{record.attempts}",
      ].join(":"),
    )
  }
  lines.sort()
  stable_hash(lines.join("\n"))
}

///|
pub fn checkpoint_record_json(record : CheckpointRecord) -> String {
  "{" +
  "\"item_id\":\{json_escape(record.item_id)}," +
  "\"state\":\{json_escape(checkpoint_state_name(record.state))}," +
  "\"attempts\":\{record.attempts}," +
  "\"input_checksum\":\{json_escape(record.input_checksum)}," +
  "\"output_checksum\":\{json_escape(record.output_checksum)}," +
  "\"finding_count\":\{record.finding_count}," +
  "\"error\":\{json_escape(record.error_message)}," +
  "\"sequence\":\{record.updated_sequence}" +
  "}"
}

///|
pub fn BatchCheckpoint::to_json(self : BatchCheckpoint) -> String {
  let records = self.records.values().to_array()
  records.sort_by(fn(left, right) {
    left.updated_sequence - right.updated_sequence
  })
  "{" +
  "\"batch_id\":\{json_escape(self.batch_id)}," +
  "\"policy_checksum\":\{json_escape(self.policy_checksum)}," +
  "\"sequence\":\{self.sequence}," +
  "\"checksum\":\{json_escape(self.checksum())}," +
  "\"records\":[" +
  records.map(checkpoint_record_json).join(",") +
  "]}"
}

///|
pub fn BatchCheckpoint::to_lines(self : BatchCheckpoint) -> Array[String] {
  let lines = [
    "checkpoint-version=1",
    "batch-id=\{self.batch_id}",
    "policy-checksum=\{self.policy_checksum}",
    "sequence=\{self.sequence}",
    "checksum=\{self.checksum()}",
  ]
  for record in self.records.values() {
    lines.push(
      [
        record.item_id,
        checkpoint_state_name(record.state),
        "\{record.attempts}",
        record.input_checksum,
        record.output_checksum,
        "\{record.finding_count}",
        record.error_message,
        "\{record.updated_sequence}",
      ].join("\t"),
    )
  }
  lines
}

///|
pub fn BatchCheckpoint::to_text(self : BatchCheckpoint) -> String {
  self.to_lines().join("\n")
}

///|
pub fn checkpoint_state_from_name(name : String) -> CheckpointState {
  match name {
    "processing" => Processing
    "succeeded" => Succeeded
    "failed" => Failed
    "skipped" => Skipped
    _ => Pending
  }
}

///|
pub fn checkpoint_from_lines(lines : Array[String]) -> BatchCheckpoint {
  let batch_id = ""
  let policy_checksum = ""
  let checkpoint = BatchCheckpoint::new(batch_id, policy_checksum)
  for line in lines {
    let parts = line.split("\t").to_array()
    if parts.length() >= 8 {
      let record : CheckpointRecord = {
        item_id: parts[0].to_owned(),
        state: checkpoint_state_from_name(parts[1].to_owned()),
        attempts: decimal_value(parts[2].to_owned()),
        input_checksum: parts[3].to_owned(),
        output_checksum: parts[4].to_owned(),
        finding_count: decimal_value(parts[5].to_owned()),
        error_message: parts[6].to_owned(),
        updated_sequence: decimal_value(parts[7].to_owned()),
      }
      checkpoint.records[record.item_id] = record
      if record.updated_sequence > checkpoint.sequence {
        checkpoint.sequence = record.updated_sequence
      }
    } else if line.has_prefix("batch-id=") {
      checkpoint.batch_id = line[9:].to_owned()
    } else if line.has_prefix("policy-checksum=") {
      checkpoint.policy_checksum = line[16:].to_owned()
    }
  }
  checkpoint
}

///|
pub fn build_resume_plan(
  checkpoint : BatchCheckpoint,
  items : Array[BatchItem],
  max_attempts : Int,
) -> ResumePlan {
  let pending = []
  let retryable = []
  let mut completed = 0
  let mut skipped = 0
  for item in items {
    match checkpoint.records.get(item.id) {
      Some(record) if record.state == Succeeded => completed += 1
      Some(record) if record.state == Skipped => skipped += 1
      Some(_) if checkpoint.is_retryable(item.id, max_attempts) =>
        retryable.push(item)
      Some(_) => pending.push(item)
      None => pending.push(item)
    }
  }
  let ids = pending.map(fn(item) { item.id }).join(",") +
    "|" +
    retryable.map(fn(item) { item.id }).join(",")
  {
    batch_id: checkpoint.batch_id,
    pending,
    retryable,
    completed,
    skipped,
    checksum: stable_hash(ids),
  }
}

///|
pub fn ResumePlan::total_work(self : ResumePlan) -> Int {
  self.pending.length() + self.retryable.length()
}

///|
pub fn ResumePlan::is_empty(self : ResumePlan) -> Bool {
  self.total_work() == 0
}

///|
pub fn ResumePlan::summary(self : ResumePlan) -> String {
  [
    "batch_id=\{self.batch_id}",
    "pending=\{self.pending.length()}",
    "retryable=\{self.retryable.length()}",
    "completed=\{self.completed}",
    "skipped=\{self.skipped}",
    "checksum=\{self.checksum}",
  ].join("\n")
}

///|
pub fn ResumePlan::to_json(self : ResumePlan) -> String {
  "{" +
  "\"batch_id\":\{json_escape(self.batch_id)}," +
  "\"pending\":\{self.pending.length()}," +
  "\"retryable\":\{self.retryable.length()}," +
  "\"completed\":\{self.completed}," +
  "\"skipped\":\{self.skipped}," +
  "\"checksum\":\{json_escape(self.checksum)}" +
  "}"
}

///|
pub fn checkpoint_record_is_consistent(
  item : BatchItem,
  record : CheckpointRecord,
) -> Bool {
  record.item_id == item.id &&
  record.input_checksum == stable_hash(item.text) &&
  record.attempts >= 0 &&
  record.finding_count >= 0
}

///|
pub fn checkpoint_consistency_issues(
  checkpoint : BatchCheckpoint,
  items : Array[BatchItem],
) -> Array[String] {
  let issues = []
  let seen : Map[String, Unit] = Map([])
  for item in items {
    if seen.contains(item.id) {
      issues.push("duplicate item id: \{item.id}")
    }
    seen[item.id] = ()
    match checkpoint.get(item.id) {
      Some(record) if !checkpoint_record_is_consistent(item, record) =>
        issues.push("checksum mismatch: \{item.id}")
      _ => ()
    }
  }
  for id, _ in checkpoint.records {
    if !seen.contains(id) {
      issues.push("checkpoint record has no input: \{id}")
    }
  }
  issues
}

///|
pub fn checkpoint_is_consistent(
  checkpoint : BatchCheckpoint,
  items : Array[BatchItem],
) -> Bool {
  checkpoint_consistency_issues(checkpoint, items).is_empty()
}

///|
pub fn checkpoint_progress_percent(
  checkpoint : BatchCheckpoint,
  total : Int,
) -> Int {
  if total <= 0 {
    100
  } else {
    let completed = checkpoint.records_by_state(Succeeded).length() +
      checkpoint.records_by_state(Skipped).length()
    completed * 100 / total
  }
}

///|
pub fn checkpoint_eta_bucket(
  checkpoint : BatchCheckpoint,
  total : Int,
) -> String {
  let progress = checkpoint_progress_percent(checkpoint, total)
  if progress >= 100 {
    "complete"
  } else if progress >= 75 {
    "finishing"
  } else if progress >= 25 {
    "in_progress"
  } else {
    "starting"
  }
}

///|
pub fn checkpoint_summary(checkpoint : BatchCheckpoint, total : Int) -> String {
  [
    "batch_id=\{checkpoint.batch_id}",
    "records=\{checkpoint.size()}",
    "progress=\{checkpoint_progress_percent(checkpoint, total)}%",
    "bucket=\{checkpoint_eta_bucket(checkpoint, total)}",
    "attempts=\{checkpoint.attempt_count()}",
    "findings=\{checkpoint.finding_count()}",
    "checksum=\{checkpoint.checksum()}",
  ].join("\n")
}