// lindividual/der — ASN.1 DER parser for MoonBit
// 
// Implements ITU-T X.690 Distinguished Encoding Rules.
// Supports PKCS#1, PKCS#8, SEC 1, SubjectPublicKeyInfo.

///|
/// Universal ASN.1 tag classes used in DER.
pub(all) enum Tag {
  Boolean // 0x01
  Integer // 0x02
  BitString // 0x03
  OctetString // 0x04
  Null // 0x05
  Oid // 0x06
  Utf8String // 0x0C
  PrintableString // 0x13
  UtcTime // 0x17
  GeneralizedTime // 0x18
  Sequence // 0x30
  Set // 0x31
} derive(Eq)

///|
/// A raw TLV triple from DER data.
pub(all) struct Tlv {
  tag : Tag
  value : Bytes
}

///|
/// Errors for malformed DER input.
pub suberror DerError {
  Invalid(String)
}

// ── Core parsing ────────────────────────────────────────────

///|
pub fn read_tlv(data : Bytes, pos : Int) -> (Tlv, Int) raise DerError {
  let len = data.length()
  if pos >= len {
    raise Invalid("unexpected eof")
  }
  let tag_byte = data[pos]
  let tag = byte_to_tag(tag_byte)
  let mut p = pos + 1
  if p >= len {
    raise Invalid("unexpected eof")
  }

  let content_len = if data[p] < b'\x80' {
    let l = data[p].to_int()
    p = p + 1
    l
  } else {
    let num_bytes = (data[p] & b'\x7f').to_int()
    if num_bytes == 0 || num_bytes > 4 {
      raise Invalid("length too long")
    }
    p = p + 1
    if p + num_bytes > len {
      raise Invalid("unexpected eof")
    }
    let mut l = 0
    for i = 0; i < num_bytes; i = i + 1 {
      l = (l << 8) | data[p + i].to_int()
    }
    p = p + num_bytes
    l
  }
  if p + content_len > len {
    raise Invalid("unexpected eof")
  }
  ({ tag, value: data[p:p + content_len].to_owned() }, p + content_len)
}

///|
pub fn skip_tlv(data : Bytes, pos : Int) -> Int raise DerError {
  let (_, next) = read_tlv(data, pos)
  next
}

///|
pub fn read_integer_bytes(
  data : Bytes,
  pos : Int,
) -> (Bytes, Int) raise DerError {
  let (tlv, next) = read_tlv(data, pos)
  if tlv.tag != Integer {
    raise Invalid("expected Integer")
  }
  let val = tlv.value
  if val.length() > 1 && val[0] == b'\x00' {
    (val[1:].to_owned(), next)
  } else {
    (val, next)
  }
}

///|
pub fn read_oid(data : Bytes, pos : Int) -> (String, Int) raise DerError {
  let (tlv, next) = read_tlv(data, pos)
  if tlv.tag != Oid {
    raise Invalid("expected OID")
  }
  let bytes = tlv.value
  let len = bytes.length()
  if len < 2 {
    raise Invalid("bad oid")
  }

  let first = bytes[0].to_int()
  let x = first / 40
  let y = first % 40
  let mut result = x.to_string() + "." + y.to_string()

  let mut i = 1
  while i < len {
    let mut val : Int64 = 0
    while i < len && (bytes[i] & b'\x80') != b'\x00' {
      val = (val << 7) | (bytes[i] & b'\x7f').to_int64()
      i = i + 1
    }
    if i < len {
      val = (val << 7) | (bytes[i] & b'\x7f').to_int64()
      i = i + 1
    }
    result = result + "." + val.to_string()
  }
  (result, next)
}

///|
pub fn enter_sequence(data : Bytes, pos : Int) -> (Int, Int) raise DerError {
  let (tlv, next) = read_tlv(data, pos)
  if tlv.tag != Sequence {
    raise Invalid("expected Sequence")
  }
  (next - tlv.value.length(), next)
}

///|
/// Strip PEM armor and base64-decode.
pub fn from_pem(pem : String) -> Bytes {
  let mut result = pem
  while result.contains("-----") {
    let start = result.find("-----").unwrap_or(0)
    let after = result[start + 5:]
    let end = after.find("-----").unwrap_or(after.length())
    let before = result[:start]
    let rest = after[end + 5:]
    result = before.to_owned() + rest.to_owned()
  }
  @base64.decode_lossy(result, ignore_whitespace=true)
}

// ── Known OIDs ──────────────────────────────────────────────

// ── PKCS#8 ──────────────────────────────────────────────────

///|
pub fn unwrap_pkcs8(data : Bytes) -> Bytes raise DerError {
  let (inner_start, _) = enter_sequence(data, 0)
  let mut pos = inner_start
  pos = skip_tlv(data, pos)
  pos = skip_tlv(data, pos)
  let (tlv, _) = read_tlv(data, pos)
  if tlv.tag != OctetString {
    raise Invalid("expected OctetString")
  }
  tlv.value
}

// ── RSA ─────────────────────────────────────────────────────

///|
pub fn parse_rsa_key(data : Bytes) -> (Bytes, Bytes) raise DerError {
  try {
    let inner = unwrap_pkcs8(data)
    parse_rsa_key(inner)
  } catch {
    _ => {
      let (inner_start, _) = enter_sequence(data, 0)
      let mut pos = inner_start
      pos = skip_tlv(data, pos)
      let (n, next) = read_integer_bytes(data, pos)
      pos = next
      pos = skip_tlv(data, pos)
      let (d, next) = read_integer_bytes(data, pos)
      pos = next
      (n, d)
    }
  }
}

// ── EC ──────────────────────────────────────────────────────

///|
pub(all) struct EcPrivateKey {
  curve : String
  d : Bytes
  pub_x : Bytes?
  pub_y : Bytes?
} derive(Debug)

///|
pub fn parse_ec_key(data : Bytes) -> EcPrivateKey raise DerError {
  let der = unwrap_pkcs8(data) catch { _ => data }
  let (inner_start, _) = enter_sequence(der, 0)
  let mut pos = inner_start

  pos = skip_tlv(der, pos) // version

  let (pk_tlv, next) = read_tlv(der, pos)
  if pk_tlv.tag != OctetString {
    raise Invalid("expected OctetString")
  }
  let d = pk_tlv.value
  pos = next

  let mut curve = ""
  let mut pub_x : Bytes? = None
  let mut pub_y : Bytes? = None

  while pos < der.length() {
    let tag_byte = der[pos]
    if tag_byte == b'\xA0' {
      let (params_tlv, np) = read_tlv(der, pos)
      let (oid_str, _) = read_oid(params_tlv.value, 0)
      curve = oid_str
      pos = np
    } else if tag_byte == b'\xA1' {
      let (pub_tlv, np) = read_tlv(der, pos)
      let bits = pub_tlv.value
      if bits.length() > 1 && bits[0] == b'\x00' {
        let key_bytes = bits[1:]
        let half = key_bytes.length() / 2
        pub_x = Some(key_bytes[:half].to_owned())
        pub_y = Some(key_bytes[half:].to_owned())
      }
      pos = np
    } else {
      break
    }
  }

  if curve == "" {
    curve = "1.2.840.10045.3.1.7"
  }
  { curve, d, pub_x, pub_y }
}


///|
pub(all) struct PublicKeyInfo {
  algorithm : String
  key : Bytes
} derive(Debug)

///|
pub fn parse_public_key(data : Bytes) -> PublicKeyInfo raise DerError {
  let (inner_start, _) = enter_sequence(data, 0)
  let mut pos = inner_start

  let (algo_tlv, next) = read_tlv(data, pos)
  let (oid_str, _) = read_oid(algo_tlv.value, 0)
  pos = next

  let (key_tlv, next) = read_tlv(data, pos)
  if key_tlv.tag != BitString {
    raise Invalid("expected BitString")
  }
  let bits = key_tlv.value
  let key = if bits.length() > 1 && bits[0] == b'\x00' {
    bits[1:]
  } else {
    bits
  }

  { algorithm: oid_str, key: key.to_owned() }
}

// ── Internal ────────────────────────────────────────────────

///|
fn byte_to_tag(b : Byte) -> Tag {
  match b {
    b'\x01' => Boolean
    b'\x02' => Integer
    b'\x03' => BitString
    b'\x04' => OctetString
    b'\x05' => Null
    b'\x06' => Oid
    b'\x0C' => Utf8String
    b'\x13' => PrintableString
    b'\x17' => UtcTime
    b'\x18' => GeneralizedTime
    b'\x30' => Sequence
    b'\x31' => Set
    _ => Utf8String
  }
}

///|
pub impl Show for Tag with fn output(self, logger) {
  let s = match self {
    Boolean => "Boolean"
    Integer => "Integer"
    BitString => "BitString"
    OctetString => "OctetString"
    Null => "Null"
    Oid => "Oid"
    Utf8String => "Utf8String"
    PrintableString => "PrintableString"
    UtcTime => "UtcTime"
    GeneralizedTime => "GeneralizedTime"
    Sequence => "Sequence"
    Set => "Set"
  }
  logger.write_string(s)
}