///| Reproducible request-work benchmarks for an external batch model provider.

///| This module intentionally reports provider requests and algorithmic work,

///| never wall-clock speed. A host application may add latency measurements

///|
/// around the same callbacks when it owns the transport and hardware.
pub enum BenchmarkError {
  InvalidModelId
  InvalidOutputBudget
  BaselineFailure(TreeExperimentError)
  AdaptiveFailure(TreeExperimentError)
  OutputLengthMismatch
} derive(Eq, Debug)

///|
pub struct AdaptiveBenchmark {
  model_id : String
  prompt_tokens : Int
  output_tokens : Int
  baseline_target_requests : Int
  adaptive_target_requests : Int
  adaptive_draft_requests : Int
  adaptive_rounds : Int
  adaptive_candidate_nodes : Int
  adaptive_accepted_nodes : Int
  planned_depths : Array[Int]
}

///|
pub fn AdaptiveBenchmark::target_request_reduction(
  self : AdaptiveBenchmark,
) -> Int {
  self.baseline_target_requests - self.adaptive_target_requests
}

///|
pub fn AdaptiveBenchmark::target_request_ratio(
  self : AdaptiveBenchmark,
) -> Double {
  if self.baseline_target_requests == 0 {
    0.0
  } else {
    self.adaptive_target_requests.to_double() /
    self.baseline_target_requests.to_double()
  }
}

///|
pub fn AdaptiveBenchmark::planned_depths(
  self : AdaptiveBenchmark,
) -> Array[Int] {
  self.planned_depths.copy()
}

///|
/// Run a matched baseline and adaptive-tree experiment through the same pure
/// target batch provider. The baseline submits exactly one context per request;
/// adaptive scoring may submit many contexts per request. Providers must return
/// one logits row for each input context, in the same order.
pub fn benchmark_adaptive_batch_decoding(
  model_id : String,
  prompt : Array[Int],
  draft : (Array[Array[Int]]) -> Result[Array[Array[Double]], String],
  target : (Array[Array[Int]]) -> Result[Array[Array[Double]], String],
  policy : AdaptiveTreePolicy,
  output_tokens : Int,
  seed : Int,
) -> Result[AdaptiveBenchmark, BenchmarkError] {
  if model_id.length() == 0 {
    return Err(InvalidModelId)
  }
  if output_tokens <= 0 {
    return Err(InvalidOutputBudget)
  }
  let baseline_config = match
    TreeExperimentConfig::new(1, 1, 1, output_tokens, seed) {
    Ok(value) => value
    Err(_) => return Err(InvalidOutputBudget)
  }
  let one_context_target = fn(
    context : Array[Int],
  ) -> Result[Array[Double], String] {
    let rows = match target([context]) {
      Ok(value) => value
      Err(error) => return Err(error)
    }
    if rows.length() != 1 {
      return Err("TreeSpec batch provider returned the wrong row count")
    }
    Ok(rows[0])
  }
  let baseline = match
    decode_baseline_from_model(prompt, one_context_target, baseline_config) {
    Ok(value) => value
    Err(error) => return Err(BaselineFailure(error))
  }
  let adaptive = match
    decode_adaptive_tree_from_batch_models(
      prompt, draft, target, policy, output_tokens, seed,
    ) {
    Ok(value) => value
    Err(error) => return Err(AdaptiveFailure(error))
  }
  if baseline.generated.length() != output_tokens ||
    adaptive.run.generated.length() != output_tokens {
    return Err(OutputLengthMismatch)
  }
  Ok({
    model_id,
    prompt_tokens: prompt.length(),
    output_tokens,
    baseline_target_requests: baseline.target_queries,
    adaptive_target_requests: adaptive.run.target_queries,
    adaptive_draft_requests: adaptive.run.draft_queries,
    adaptive_rounds: adaptive.run.logical_target_batches,
    adaptive_candidate_nodes: adaptive.run.candidate_nodes,
    adaptive_accepted_nodes: adaptive.run.accepted_nodes,
    planned_depths: adaptive.planned_depths(),
  })
}

///|
/// Stable text output suitable for a checked-in benchmark record. The request
/// ratio is algorithmic work, not an assertion about latency or throughput.
pub fn AdaptiveBenchmark::render(self : AdaptiveBenchmark) -> String {
  let mut depths = ""
  for index in 0.. 0 {
      depths = depths + ","
    }
    depths = depths + self.planned_depths[index].to_string()
  }
  "TreeSpec adaptive batch benchmark\nmodel_id=" +
  self.model_id +
  "\nprompt_tokens=" +
  self.prompt_tokens.to_string() +
  "\noutput_tokens=" +
  self.output_tokens.to_string() +
  "\nbaseline_target_requests=" +
  self.baseline_target_requests.to_string() +
  "\nadaptive_target_requests=" +
  self.adaptive_target_requests.to_string() +
  "\nadaptive_draft_requests=" +
  self.adaptive_draft_requests.to_string() +
  "\nadaptive_rounds=" +
  self.adaptive_rounds.to_string() +
  "\nadaptive_candidate_nodes=" +
  self.adaptive_candidate_nodes.to_string() +
  "\nadaptive_accepted_nodes=" +
  self.adaptive_accepted_nodes.to_string() +
  "\ntarget_request_reduction=" +
  self.target_request_reduction().to_string() +
  "\ntarget_request_ratio=" +
  self.target_request_ratio().to_string() +
  "\nplanned_depths=[" +
  depths +
  "]\nNo latency or throughput claim is measured.\n"
}