///|
pub(all) struct FluidState {
  density : Double
  velocity : Double
  pressure : Double
  energy : Double
} derive(Debug, ToJson)

///|
pub(all) struct FluidFlux {
  mass : Double
  momentum : Double
  energy : Double
} derive(Debug, ToJson)

///|
pub fn FluidState::new(
  density : Double,
  velocity : Double,
  pressure : Double,
  energy : Double,
) -> FluidState {
  { density, velocity, pressure, energy }
}

///|
pub fn sound_speed(state : FluidState, gamma : Double) -> Double {
  if state.density <= 0.0 || gamma <= 0.0 {
    0.0
  } else {
    (gamma * state.pressure / state.density).sqrt()
  }
}

///|
pub fn specific_internal_energy(state : FluidState) -> Double {
  safe_ratio(state.energy, state.density, 0.0) -
  0.5 * state.velocity * state.velocity
}

///|
pub fn euler_flux(state : FluidState, _gamma : Double) -> FluidFlux {
  let momentum = state.density * state.velocity
  {
    mass: momentum,
    momentum: momentum * state.velocity + state.pressure,
    energy: (state.energy + state.pressure) * state.velocity,
  }
}

///|
pub fn fluid_pressure(
  density : Double,
  temperature : Double,
  gas_constant : Double,
) -> Double {
  density * gas_constant * temperature
}

///|
pub fn fluid_energy(
  density : Double,
  velocity : Double,
  pressure : Double,
  gamma : Double,
) -> Double {
  if gamma <= 1.0 {
    0.5 * density * velocity * velocity
  } else {
    pressure / (gamma - 1.0) + 0.5 * density * velocity * velocity
  }
}

///|
pub fn fluid_state_from_primitive(
  density : Double,
  velocity : Double,
  pressure : Double,
  gamma : Double,
) -> FluidState {
  FluidState::new(
    density,
    velocity,
    pressure,
    fluid_energy(density, velocity, pressure, gamma),
  )
}

///|
pub fn fluid_update(
  state : FluidState,
  flux_difference : FluidFlux,
  dt : Double,
  cell_width : Double,
) -> FluidState {
  let scale = safe_ratio(dt, cell_width, 0.0)
  {
    density: state.density - scale * flux_difference.mass,
    velocity: state.velocity -
    scale * flux_difference.momentum / state.density.max(1.0e-30),
    pressure: state.pressure,
    energy: state.energy - scale * flux_difference.energy,
  }
}

///|
pub fn lax_friedrichs_flux(
  left : FluidState,
  right : FluidState,
  gamma : Double,
  wave_speed : Double,
) -> FluidFlux {
  let a = euler_flux(left, gamma)
  let b = euler_flux(right, gamma)
  let speed = wave_speed.abs()
  {
    mass: 0.5 * (a.mass + b.mass) - 0.5 * speed * (right.density - left.density),
    momentum: 0.5 * (a.momentum + b.momentum) -
    0.5 *
    speed *
    (right.density * right.velocity - left.density * left.velocity),
    energy: 0.5 * (a.energy + b.energy) -
    0.5 * speed * (right.energy - left.energy),
  }
}

///|
pub fn fluid_mach_number(state : FluidState, gamma : Double) -> Double {
  safe_ratio(state.velocity.abs(), sound_speed(state, gamma), 0.0)
}

///|
pub fn fluid_enthalpy(state : FluidState) -> Double {
  safe_ratio(state.energy + state.pressure, state.density, 0.0)
}

///|
pub fn fluid_compressibility(state : FluidState, gamma : Double) -> Double {
  safe_ratio(1.0, gamma * state.pressure.max(1.0e-30), 0.0)
}

///|
pub fn fluid_temperature(state : FluidState, gas_constant : Double) -> Double {
  safe_ratio(state.pressure, state.density * gas_constant, 0.0)
}

///|
pub fn fluid_state_is_physical(state : FluidState) -> Bool {
  state.density > 0.0 && state.pressure >= 0.0 && state.energy >= 0.0
}

///|
pub fn fluid_clamp_physical(state : FluidState) -> FluidState {
  {
    density: state.density.max(1.0e-30),
    velocity: state.velocity,
    pressure: state.pressure.max(0.0),
    energy: state.energy.max(0.0),
  }
}

///|
pub fn fluid_wave_speed(state : FluidState, gamma : Double) -> Double {
  state.velocity.abs() + sound_speed(state, gamma)
}

///|
pub fn fluid_cfl_dt(
  states : ArrayView[FluidState],
  gamma : Double,
  cell_width : Double,
  cfl : Double,
) -> Double {
  let mut maximum = 0.0
  for state in states {
    maximum = maximum.max(fluid_wave_speed(state, gamma))
  }
  safe_ratio(cfl * cell_width, maximum, 0.0)
}

///|
pub fn fluid_mean_density(states : ArrayView[FluidState]) -> Double {
  mean(states.map(fn(state) { state.density }))
}

///|
pub fn fluid_mean_velocity(states : ArrayView[FluidState]) -> Double {
  mean(states.map(fn(state) { state.velocity }))
}

///|
pub fn fluid_mean_pressure(states : ArrayView[FluidState]) -> Double {
  mean(states.map(fn(state) { state.pressure }))
}