///|
fn finish_reason_text(reason : @llama.FinishReason) -> String {
  match reason {
    @llama.EndOfGeneration => "end of generation"
    @llama.MaxTokens => "maximum token count"
    @llama.ContextFull => "context full"
  }
}

///|
fn run(
  model_path : StringView,
  strategy : @llama.SamplingStrategy,
) -> Unit raise {
  let model = @llama.Model::load(model_path)
  let session = model.new_session(context_size=128)
  let options = @llama.GenerationOptions::new(max_tokens=32, strategy~)
  let completion = session.complete("Once upon a time", options~)
  println(completion.text)
  println("Prompt tokens: \{completion.prompt_tokens}")
  println("Generated tokens: \{completion.generated_tokens}")
  println("Finish reason: \{finish_reason_text(completion.finish_reason)}")
}

///|
fn parse_strategy(args : Array[String]) -> @llama.SamplingStrategy? {
  if args.length() == 2 {
    Some(@llama.SamplingStrategy::greedy())
  } else if args.length() == 3 {
    match args[2] {
      "greedy" => Some(@llama.SamplingStrategy::greedy())
      "random" => Some(@llama.SamplingStrategy::random(seed=42U))
      _ => None
    }
  } else {
    None
  }
}

///|
fn main {
  let args = @env.args()
  guard parse_strategy(args) is Some(strategy) else {
    println(
      "Usage: moon run main --target native -- <model.gguf> [greedy|random]",
    )
    return
  }
  run(args[1], strategy) catch {
    error => println("Error: \{error.to_string()}")
  }
}
