///|
pub suberror LiteRtError {
  LiteRtError(String)
}

///|
pub fn LiteRtError::to_string(self : LiteRtError) -> String {
  let LiteRtError(message) = self
  message
}

///|
pub enum LiteRtValueSource {
  Input
  Constant(Array[Float])
  Intermediate
}

///|
pub struct LiteRtValue {
  name_ : String
  shape_ : @shape.Shape
  source_ : LiteRtValueSource
}

///|
fn validate_name(name : String) -> Unit raise LiteRtError {
  if name == "" {
    raise LiteRtError("LiteRT value name must not be empty")
  }
}

///|
pub fn LiteRtValue::input(
  name : String,
  shape : @shape.Shape,
) -> LiteRtValue raise LiteRtError {
  validate_name(name)
  { name_: name, shape_: shape, source_: Input }
}

///|
pub fn LiteRtValue::constant(
  name : String,
  shape : @shape.Shape,
  values : Array[Float],
) -> LiteRtValue raise LiteRtError {
  validate_name(name)
  if shape.element_count() != values.length() {
    raise LiteRtError(
      "LiteRT constant \{name} has \{values.length()} values for shape \{shape.to_string()}",
    )
  }
  { name_: name, shape_: shape, source_: Constant(values.copy()) }
}

///|
pub fn LiteRtValue::intermediate(
  name : String,
  shape : @shape.Shape,
) -> LiteRtValue raise LiteRtError {
  validate_name(name)
  { name_: name, shape_: shape, source_: Intermediate }
}

///|
pub fn LiteRtValue::name(self : LiteRtValue) -> String {
  self.name_
}

///|
pub fn LiteRtValue::shape(self : LiteRtValue) -> @shape.Shape {
  self.shape_
}

///|
pub(all) enum LiteRtNode {
  Add(String, String, String)
  Sub(String, String, String)
  Mul(String, String, String)
  Div(String, String, String)
  ReduceMean(String, String, Array[Int], Bool)
  Gather(String, String, Array[Int], @shape.Shape, Int)
  Slice(String, String, Array[Int], Array[Int])
  Gelu(String, String)
  LayerNormalization(String, String, String, String, Array[Int], Float)
  Concat(String, String, String, Int)
  Matmul(String, String, String)
  Conv2d(String, String, String, @shape.Conv2dOptions)
  MaxPool2d(String, String, @shape.Pool2dOptions)
  AveragePool2d(String, String, @shape.Pool2dOptions)
  Sigmoid(String, String)
  Tanh(String, String)
  Clamp(String, String, Float, Float)
  Relu(String, String)
  Softmax(String, String, Int)
  Reshape(String, String)
  Transpose(String, String, Array[Int])
}

///|
fn input_names(node : LiteRtNode) -> Array[String] {
  match node {
    Add(lhs, rhs, _) => [lhs, rhs]
    Sub(lhs, rhs, _) => [lhs, rhs]
    Mul(lhs, rhs, _) => [lhs, rhs]
    Div(lhs, rhs, _) => [lhs, rhs]
    ReduceMean(input, _, _, _) => [input]
    Gather(input, _, _, _, _) => [input]
    Slice(input, _, _, _) => [input]
    Gelu(input, _) => [input]
    LayerNormalization(input, scale, bias, _, _, _) => [input, scale, bias]
    Concat(lhs, rhs, _, _) => [lhs, rhs]
    Matmul(lhs, rhs, _) => [lhs, rhs]
    Conv2d(input, filter, _, _) => [input, filter]
    MaxPool2d(input, _, _) => [input]
    AveragePool2d(input, _, _) => [input]
    Sigmoid(input, _) => [input]
    Tanh(input, _) => [input]
    Clamp(input, _, _, _) => [input]
    Relu(input, _) => [input]
    Softmax(input, _, _) => [input]
    Reshape(input, _) => [input]
    Transpose(input, _, _) => [input]
  }
}

///|
fn output_name(node : LiteRtNode) -> String {
  match node {
    Add(_, _, output) => output
    Sub(_, _, output) => output
    Mul(_, _, output) => output
    Div(_, _, output) => output
    ReduceMean(_, output, _, _) => output
    Gather(_, output, _, _, _) => output
    Slice(_, output, _, _) => output
    Gelu(_, output) => output
    LayerNormalization(_, _, _, output, _, _) => output
    Concat(_, _, output, _) => output
    Matmul(_, _, output) => output
    Conv2d(_, _, output, _) => output
    MaxPool2d(_, output, _) => output
    AveragePool2d(_, output, _) => output
    Sigmoid(_, output) => output
    Tanh(_, output) => output
    Clamp(_, output, _, _) => output
    Relu(_, output) => output
    Softmax(_, output, _) => output
    Reshape(_, output) => output
    Transpose(_, output, _) => output
  }
}

///|
pub struct LiteRtModel {
  values_ : Array[LiteRtValue]
  nodes_ : Array[LiteRtNode]
  output_names_ : Array[String]
}

///|
pub struct LiteRtLoweredValue[T] {
  name_ : String
  tensor_ : T
}

///|
pub fn[T] LiteRtLoweredValue::name(self : LiteRtLoweredValue[T]) -> String {
  self.name_
}

///|
pub fn[T] LiteRtLoweredValue::tensor(self : LiteRtLoweredValue[T]) -> T {
  self.tensor_
}

///|
fn value_for(
  values : Map[String, LiteRtValue],
  name : String,
) -> LiteRtValue raise LiteRtError {
  guard values.contains(name) else {
    raise LiteRtError("LiteRT graph refers to unknown value: \{name}")
  }
  values[name]
}

///|
fn shape_or_litert_error(
  operation : () -> @shape.Shape raise @shape.ShapeError,
) -> @shape.Shape raise LiteRtError {
  operation() catch {
    error => raise LiteRtError(error.to_string())
  }
}

///|
fn infer_shape(
  node : LiteRtNode,
  values : Map[String, LiteRtValue],
) -> @shape.Shape raise LiteRtError {
  match node {
    Add(lhs, rhs, _) => {
      let lhs_shape = value_for(values, lhs).shape_
      let rhs_shape = value_for(values, rhs).shape_
      shape_or_litert_error(() => @shape.Shape::broadcast(lhs_shape, rhs_shape))
    }
    Sub(lhs, rhs, _) => {
      let lhs_shape = value_for(values, lhs).shape_
      let rhs_shape = value_for(values, rhs).shape_
      shape_or_litert_error(() => @shape.Shape::broadcast(lhs_shape, rhs_shape))
    }
    Mul(lhs, rhs, _) => {
      let lhs_shape = value_for(values, lhs).shape_
      let rhs_shape = value_for(values, rhs).shape_
      shape_or_litert_error(() => @shape.Shape::broadcast(lhs_shape, rhs_shape))
    }
    Div(lhs, rhs, _) => {
      let lhs_shape = value_for(values, lhs).shape_
      let rhs_shape = value_for(values, rhs).shape_
      shape_or_litert_error(() => @shape.Shape::broadcast(lhs_shape, rhs_shape))
    }
    ReduceMean(input, _, axes, keep_dimensions) => {
      let input_shape = value_for(values, input).shape_
      shape_or_litert_error(() => input_shape.reduce(axes, keep_dimensions))
    }
    Gather(input, _, _, indices_shape, axis) => {
      let input_shape = value_for(values, input).shape_
      shape_or_litert_error(() => input_shape.gather(indices_shape, axis))
    }
    Slice(input, _, starts, sizes) => {
      let input_shape = value_for(values, input).shape_
      shape_or_litert_error(() => input_shape.slice(starts, sizes))
    }
    Gelu(input, _) => value_for(values, input).shape_
    LayerNormalization(input, scale, bias, _, axes, epsilon) => {
      let input_shape = value_for(values, input).shape_
      let scale_shape = value_for(values, scale).shape_
      let bias_shape = value_for(values, bias).shape_
      shape_or_litert_error(() => {
        @shape.Shape::layer_normalization(
          input_shape, scale_shape, bias_shape, axes, epsilon,
        )
      })
    }
    Concat(lhs, rhs, _, axis) => {
      let lhs_shape = value_for(values, lhs).shape_
      let rhs_shape = value_for(values, rhs).shape_
      shape_or_litert_error(() => {
        @shape.Shape::concat(lhs_shape, rhs_shape, axis)
      })
    }
    Matmul(lhs, rhs, _) => {
      let lhs_shape = value_for(values, lhs).shape_
      let rhs_shape = value_for(values, rhs).shape_
      shape_or_litert_error(() => @shape.Shape::matmul(lhs_shape, rhs_shape))
    }
    Conv2d(input, filter, _, options) => {
      let input_shape = value_for(values, input).shape_
      let filter_shape = value_for(values, filter).shape_
      shape_or_litert_error(() => {
        @shape.Shape::conv2d(input_shape, filter_shape, options)
      })
    }
    MaxPool2d(input, _, options) => {
      let input_shape = value_for(values, input).shape_
      shape_or_litert_error(() => @shape.Shape::pool2d(input_shape, options))
    }
    AveragePool2d(input, _, options) => {
      let input_shape = value_for(values, input).shape_
      shape_or_litert_error(() => @shape.Shape::pool2d(input_shape, options))
    }
    Sigmoid(input, _) => value_for(values, input).shape_
    Tanh(input, _) => value_for(values, input).shape_
    Clamp(input, _, minimum, maximum) => {
      if minimum > maximum {
        raise LiteRtError("LiteRT clamp minimum must not exceed maximum")
      }
      value_for(values, input).shape_
    }
    Relu(input, _) => value_for(values, input).shape_
    Softmax(input, _, axis) => {
      let shape = value_for(values, input).shape_
      shape.validate_axis(axis) catch {
        error => raise LiteRtError(error.to_string())
      }
      shape
    }
    Reshape(input, output) => {
      let input_shape = value_for(values, input).shape_
      let output_shape = value_for(values, output).shape_
      if input_shape.element_count() != output_shape.element_count() {
        raise LiteRtError(
          "LiteRT reshape from \{input} to \{output} changes element count",
        )
      }
      output_shape
    }
    Transpose(input, _, permutation) => {
      let input_shape = value_for(values, input).shape_
      shape_or_litert_error(() => input_shape.transpose(permutation))
    }
  }
}

///|
fn validate_node(
  node : LiteRtNode,
  values : Map[String, LiteRtValue],
  available : Map[String, Bool],
) -> Unit raise LiteRtError {
  for input in input_names(node) {
    let _ = value_for(values, input)
    if !available.contains(input) {
      raise LiteRtError("LiteRT node uses \{input} before it is available")
    }
  }
  let output = output_name(node)
  let output_value = value_for(values, output)
  match output_value.source_ {
    Intermediate => ()
    _ => raise LiteRtError("LiteRT node output \{output} must be intermediate")
  }
  if available.contains(output) {
    raise LiteRtError("LiteRT node overwrites value: \{output}")
  }
  let inferred = infer_shape(node, values)
  if !inferred.same_as(output_value.shape_) {
    raise LiteRtError(
      "LiteRT node output \{output} has shape \{output_value.shape_.to_string()}, expected \{inferred.to_string()}",
    )
  }
  available[output] = true
}

///|
pub fn LiteRtModel::new(
  values : Array[LiteRtValue],
  nodes : Array[LiteRtNode],
  output_names : Array[String],
) -> LiteRtModel raise LiteRtError {
  if values.is_empty() {
    raise LiteRtError("LiteRT graph must declare at least one value")
  }
  if output_names.is_empty() {
    raise LiteRtError("LiteRT graph must declare at least one output")
  }
  let values_by_name : Map[String, LiteRtValue] = Map([])
  let available : Map[String, Bool] = Map([])
  for value in values {
    validate_name(value.name_)
    if values_by_name.contains(value.name_) {
      raise LiteRtError("duplicate LiteRT value name: \{value.name_}")
    }
    values_by_name[value.name_] = value
    match value.source_ {
      Input | Constant(_) => available[value.name_] = true
      Intermediate => ()
    }
  }
  for node in nodes {
    validate_node(node, values_by_name, available)
  }
  let seen_outputs : Map[String, Bool] = Map([])
  for output in output_names {
    validate_name(output)
    let _ = value_for(values_by_name, output)
    if !available.contains(output) {
      raise LiteRtError("LiteRT output is not produced: \{output}")
    }
    if seen_outputs.contains(output) {
      raise LiteRtError("duplicate LiteRT output name: \{output}")
    }
    seen_outputs[output] = true
  }
  {
    values_: values.copy(),
    nodes_: nodes.copy(),
    output_names_: output_names.copy(),
  }
}

///|
pub fn LiteRtModel::inputs(self : LiteRtModel) -> Array[LiteRtValue] {
  self.values_.filter(fn(value) {
    match value.source_ {
      Input => true
      _ => false
    }
  })
}

///|
fn[T] tensor_for(
  tensors : Map[String, T],
  name : String,
) -> T raise LiteRtError {
  guard tensors.contains(name) else {
    raise LiteRtError("LiteRT lowering did not materialize value: \{name}")
  }
  tensors[name]
}

///|
fn[T : @tensor.TensorOps] lower_node(
  node : LiteRtNode,
  tensors : Map[String, T],
  output_shape : @shape.Shape,
) -> T raise {
  match node {
    Add(lhs, rhs, _) => tensor_for(tensors, lhs).add(tensor_for(tensors, rhs))
    Sub(lhs, rhs, _) => tensor_for(tensors, lhs).sub(tensor_for(tensors, rhs))
    Mul(lhs, rhs, _) => tensor_for(tensors, lhs).mul(tensor_for(tensors, rhs))
    Div(lhs, rhs, _) => tensor_for(tensors, lhs).div(tensor_for(tensors, rhs))
    ReduceMean(input, _, axes, keep_dimensions) =>
      tensor_for(tensors, input).reduce_mean(axes, keep_dimensions)
    Gather(input, _, indices, indices_shape, axis) =>
      tensor_for(tensors, input).gather(indices, indices_shape, axis)
    Slice(input, _, starts, sizes) =>
      tensor_for(tensors, input).slice(starts, sizes)
    Gelu(input, _) => tensor_for(tensors, input).gelu()
    LayerNormalization(input, scale, bias, _, axes, epsilon) =>
      tensor_for(tensors, input).layer_normalization(
        tensor_for(tensors, scale),
        tensor_for(tensors, bias),
        axes,
        epsilon,
      )
    Concat(lhs, rhs, _, axis) =>
      tensor_for(tensors, lhs).concat(tensor_for(tensors, rhs), axis)
    Matmul(lhs, rhs, _) =>
      tensor_for(tensors, lhs).matmul(tensor_for(tensors, rhs))
    Conv2d(input, filter, _, options) =>
      tensor_for(tensors, input).conv2d(tensor_for(tensors, filter), options)
    MaxPool2d(input, _, options) =>
      tensor_for(tensors, input).max_pool2d(options)
    AveragePool2d(input, _, options) =>
      tensor_for(tensors, input).average_pool2d(options)
    Sigmoid(input, _) => tensor_for(tensors, input).sigmoid()
    Tanh(input, _) => tensor_for(tensors, input).tanh()
    Clamp(input, _, minimum, maximum) =>
      tensor_for(tensors, input).clamp(minimum, maximum)
    Relu(input, _) => tensor_for(tensors, input).relu()
    Softmax(input, _, axis) => tensor_for(tensors, input).softmax(axis)
    Reshape(input, _) => tensor_for(tensors, input).reshape(output_shape)
    Transpose(input, _, permutation) =>
      tensor_for(tensors, input).transpose(permutation)
  }
}

///|
pub fn[T : @tensor.TensorOps] LiteRtModel::lower(
  self : LiteRtModel,
  make_input : (String, @shape.Shape) -> T raise,
  make_constant : (String, @shape.Shape, Array[Float]) -> T raise,
) -> Array[LiteRtLoweredValue[T]] raise {
  let tensors : Map[String, T] = Map([])
  let values_by_name : Map[String, LiteRtValue] = Map([])
  for value in self.values_ {
    values_by_name[value.name_] = value
    match value.source_ {
      Input => {
        let tensor = make_input(value.name_, value.shape_)
        if !tensor.shape().same_as(value.shape_) {
          raise LiteRtError(
            "LiteRT input factory returned wrong shape for \{value.name_}",
          )
        }
        tensors[value.name_] = tensor
      }
      Constant(values) => {
        let tensor = make_constant(value.name_, value.shape_, values.copy())
        if !tensor.shape().same_as(value.shape_) {
          raise LiteRtError(
            "LiteRT constant factory returned wrong shape for \{value.name_}",
          )
        }
        tensors[value.name_] = tensor
      }
      Intermediate => ()
    }
  }
  for node in self.nodes_ {
    let output = output_name(node)
    let expected = value_for(values_by_name, output)
    let tensor = lower_node(node, tensors, expected.shape_)
    if !tensor.shape().same_as(expected.shape_) {
      raise LiteRtError("LiteRT lowering produced wrong shape for \{output}")
    }
    tensors[output] = tensor
  }
  self.output_names_.map(fn(name) raise {
    { name_: name, tensor_: tensor_for(tensors, name) }
  })
}