///|
let ed_p : BigInt = @bigint.BigInt::from_string(
  "7fffffffffffffffffffffffffffffffffffffffffffffffffffffffffffffed",
  radix=16,
)

///|
let ed_d : BigInt = @bigint.BigInt::from_string(
  "52036cee2b6ffe738cc740797779e89800700a4d4141d8ab75eb4dca135978a3",
  radix=16,
)

///|
let ed_L : BigInt = @bigint.BigInt::from_string(
  "1000000000000000000000000000000014def9dea2f79cd65812631a5cf5d3ed",
  radix=16,
)

///|
let ed_Gx : BigInt = @bigint.BigInt::from_string(
  "216936d3cd6e53fec0a4e231fdd6dc5c692cc7609525a7b2c9562d608f25d51a",
  radix=16,
)

///|
let ed_Gy : BigInt = @bigint.BigInt::from_string(
  "6666666666666666666666666666666666666666666666666666666666666658",
  radix=16,
)

///|
pub(all) struct ExtPoint {
  x : BigInt
  y : BigInt
  z : BigInt
  t : BigInt
}

///|
fn field_add(a : BigInt, b : BigInt) -> BigInt {
  (a + b) % ed_p
}

///|
fn field_sub(a : BigInt, b : BigInt) -> BigInt {
  (a - b + ed_p) % ed_p
}

///|
pub fn field_mul(a : BigInt, b : BigInt) -> BigInt {
  a * b % ed_p
}

///|
pub fn field_inv(a : BigInt) -> BigInt {
  let p_minus_2 = ed_p - @bigint.BigInt::from_int(2)
  a.pow(p_minus_2, modulus=ed_p)
}

///|
fn field_sqrt(a : BigInt) -> BigInt? {
  // p mod 8 = 5, so sqrt(a) = a^((p+3)/8) mod p
  let exp = (ed_p + @bigint.BigInt::from_int(3)) / @bigint.BigInt::from_int(8)
  let candidate = a.pow(exp, modulus=ed_p)
  if field_mul(candidate, candidate) == a % ed_p {
    Some(candidate)
  } else {
    // Try multiplying by sqrt(-1)
    let sqrt_m1 = @bigint.BigInt::from_int(2).pow(
      (ed_p - @bigint.BigInt::from_int(1)) / @bigint.BigInt::from_int(4),
      modulus=ed_p,
    )
    let candidate2 = field_mul(candidate, sqrt_m1)
    if field_mul(candidate2, candidate2) == a % ed_p {
      Some(candidate2)
    } else {
      None
    }
  }
}

///|
fn point_zero() -> ExtPoint {
  ExtPoint::{
    x: @bigint.BigInt::from_int(0),
    y: @bigint.BigInt::from_int(1),
    z: @bigint.BigInt::from_int(1),
    t: @bigint.BigInt::from_int(0),
  }
}

///|
pub fn base_point() -> ExtPoint {
  ExtPoint::{
    x: ed_Gx,
    y: ed_Gy,
    z: @bigint.BigInt::from_int(1),
    t: field_mul(ed_Gx, ed_Gy),
  }
}

///|
fn point_add(p1 : ExtPoint, p2 : ExtPoint) -> ExtPoint {
  let a = field_mul(field_sub(p1.y, p1.x), field_sub(p2.y, p2.x))
  let b = field_mul(field_add(p1.y, p1.x), field_add(p2.y, p2.x))
  let c = field_mul(
    field_mul(p1.t, p2.t),
    field_mul(ed_d, @bigint.BigInt::from_int(2)),
  )
  let d = field_mul(p1.z, field_mul(p2.z, @bigint.BigInt::from_int(2)))
  let e = field_sub(b, a)
  let f = field_sub(d, c)
  let g = field_add(d, c)
  let h = field_add(b, a)
  ExtPoint::{
    x: field_mul(e, f),
    y: field_mul(g, h),
    z: field_mul(f, g),
    t: field_mul(e, h),
  }
}

///|
fn point_double(p : ExtPoint) -> ExtPoint {
  let aa = field_mul(p.x, p.x)
  let bb = field_mul(p.y, p.y)
  let cc = field_mul(field_mul(p.z, p.z), @bigint.BigInt::from_int(2))
  // D = a * A = -A (since a = -1 for Ed25519)
  let dd = field_sub(@bigint.BigInt::from_int(0), aa)
  let e = field_sub(
    field_mul(field_add(p.x, p.y), field_add(p.x, p.y)),
    field_add(aa, bb),
  )
  let g = field_add(dd, bb)
  let f = field_sub(g, cc)
  let h = field_sub(dd, bb)
  ExtPoint::{
    x: field_mul(e, f),
    y: field_mul(g, h),
    z: field_mul(f, g),
    t: field_mul(e, h),
  }
}

///|
pub fn scalar_mult(scalar : BigInt, point : ExtPoint) -> ExtPoint {
  let mut result = point_zero()
  let mut temp = point
  let bits = scalar.bit_length()
  let scalar_bytes = bigint_to_le_bytes(scalar, (bits + 7) / 8)
  for i = 0; i < bits; i = i + 1 {
    let byte_idx = i / 8
    let bit_idx = i % 8
    let bit = (scalar_bytes[byte_idx].to_int() >> bit_idx) & 1
    if bit == 1 {
      result = point_add(result, temp)
    }
    temp = point_double(temp)
  }
  result
}

///|
pub fn bigint_to_le_bytes(n : BigInt, len : Int) -> Array[Byte] {
  let be = n.to_octets(length=len)
  let le : Array[Byte] = Array::make(len, b'\x00')
  for i = 0; i < len; i = i + 1 {
    le[i] = be[len - 1 - i]
  }
  le
}

///|
pub fn le_bytes_to_bigint(bytes : Array[Byte]) -> BigInt {
  let len = bytes.length()
  let be : Array[Byte] = Array::make(len, b'\x00')
  for i = 0; i < len; i = i + 1 {
    be[i] = bytes[len - 1 - i]
  }
  @bigint.BigInt::from_octets(Bytes::from_array(be))
}

///|
pub fn point_encode(p : ExtPoint) -> Array[Byte] {
  let zinv = field_inv(p.z)
  let x = field_mul(p.x, zinv)
  let y = field_mul(p.y, zinv)
  let encoded = bigint_to_le_bytes(y, 32)
  let x_bytes = bigint_to_le_bytes(x, 32)
  if (x_bytes[0].to_int() & 1) == 1 {
    encoded[31] = (encoded[31].to_int() | 0x80).to_byte()
  }
  encoded
}

///|
fn point_decode(bytes : Array[Byte]) -> ExtPoint? {
  if bytes.length() != 32 {
    return None
  }
  let x_sign = (bytes[31].to_int() >> 7) & 1
  let y_bytes : Array[Byte] = Array::make(32, b'\x00')
  for i = 0; i < 32; i = i + 1 {
    y_bytes[i] = bytes[i]
  }
  y_bytes[31] = (y_bytes[31].to_int() & 0x7F).to_byte()
  let y = le_bytes_to_bigint(y_bytes)
  if y >= ed_p {
    return None
  }
  let y2 = field_mul(y, y)
  let u = field_sub(y2, @bigint.BigInt::from_int(1))
  let v = field_add(field_mul(ed_d, y2), @bigint.BigInt::from_int(1))
  let v_inv = field_inv(v)
  let x2 = field_mul(u, v_inv)
  if x2 == @bigint.BigInt::from_int(0) {
    if x_sign != 0 {
      return None
    }
    return Some(ExtPoint::{
      x: @bigint.BigInt::from_int(0),
      y,
      z: @bigint.BigInt::from_int(1),
      t: @bigint.BigInt::from_int(0),
    })
  }
  match field_sqrt(x2) {
    None => None
    Some(x) => {
      let x_le = bigint_to_le_bytes(x, 32)
      let x_parity = x_le[0].to_int() & 1
      let x_final = if x_parity == x_sign {
        x
      } else {
        field_sub(@bigint.BigInt::from_int(0), x)
      }
      Some(ExtPoint::{
        x: x_final,
        y,
        z: @bigint.BigInt::from_int(1),
        t: field_mul(x_final, y),
      })
    }
  }
}