///| KingbaseES client: PostgreSQL wire protocol 3.0, simple query, text COPY.

///|

///| Blocking sockets, no async runtime. Failures raise, and the SQLSTATE code of

///|
/// a server error stays in the message text.

///|

///| ## What the mode does not decide

///|

///| `database_mode` fixes the dialect family, and three more settings change the

///| answer inside one family. `connect` reads them once, so the rest of a program

///| asks a `Dialect` instead of testing the mode:

///|

///| ```moonbit

///| let client = @kb.connect(@kb.new_config(

///|   "db.example.com", 54321, "app", secret, "app",

///| ))

///| println(client.mode().name())               // "sqlserver" in that mode

///| println(client.dialect().timestamp_type())  // "datetime" there, "timestamp" in pg

///|
/// ```
pub using @dialect {
  mode_of,
  type Mode,
  type Dialect,
  type ColumnKind,
  type CaseFold,
}

///|
pub struct Config {
  host : String
  port : Int
  user : String
  password : String
  database : String
  /// Read `database_mode` and the settings that refine it after
  /// authentication. One extra query per connection.
  detect_dialect : Bool
  /// Used when detection is off, or when the instance does not report the mode.
  fallback_mode : Mode?
}

///|
/// A configuration that detects the dialect of the target.
pub fn new_config(
  host : String,
  port : Int,
  user : String,
  password : String,
  database : String,
) -> Config {
  {
    host,
    port,
    user,
    password,
    database,
    detect_dialect: true,
    fallback_mode: None,
  }
}

///|
/// A configuration that skips the detection query, for a connection used once by
/// a program that already knows the mode.
pub fn new_config_fixed_mode(
  host : String,
  port : Int,
  user : String,
  password : String,
  database : String,
  mode : Mode,
) -> Config {
  {
    host,
    port,
    user,
    password,
    database,
    detect_dialect: false,
    fallback_mode: Some(mode),
  }
}

///|
pub(all) struct ServerError {
  severity : String
  code : String
  message : String
}

///|
pub fn ServerError::describe(self : ServerError) -> String {
  "[\{self.code}] \{self.severity}: \{self.message}"
}

///|
/// One result set: column names and text-format values (`None` = SQL NULL).

///|
/// Every value is the server's own text output. The client parses no type, so a
/// mode that prints a timestamp differently cannot break a read.
pub struct ResultSet {
  columns : Array[String]
  rows : Array[Array[String?]]
  command : String
}

///|
pub fn ResultSet::row_count(self : ResultSet) -> Int {
  self.rows.length()
}

///|
pub fn ResultSet::column_count(self : ResultSet) -> Int {
  self.columns.length()
}

///|
pub fn ResultSet::cell(self : ResultSet, row : Int, col : Int) -> String {
  match self.rows.get(row) {
    Some(r) =>
      match r.get(col) {
        Some(Some(s)) => s
        _ => "NULL"
      }
    None => ""
  }
}

///|
pub fn ResultSet::first_text(self : ResultSet, col : Int) -> String {
  self.cell(0, col)
}

///|
/// The first column of the first row, or "" when the set is empty.
pub fn ResultSet::scalar(self : ResultSet) -> String {
  if self.rows.length() == 0 {
    ""
  } else {
    self.cell(0, 0)
  }
}

///|
/// Rows as `(first column, second column)` pairs, for setting and catalog reads.
pub fn ResultSet::pairs(self : ResultSet) -> Array[(String, String)] {
  let out : Array[(String, String)] = []
  for r in self.rows {
    match r.get(0) {
      Some(Some(k)) => {
        let v = match r.get(1) {
          Some(Some(s)) => s
          _ => ""
        }
        out.push((k, v))
      }
      _ => ()
    }
  }
  out
}

///|
/// How many rows the statement affected, read from the command tag.

///|
/// measured tags on these instances: `SELECT 1`, `INSERT 0 1`, `UPDATE 3`,
/// `DELETE 1000`, `COPY 5000`, and `SHOW` for a setting. A tag without a trailing
/// number reports no count, so this returns 0 rather than guessing.

///|
/// The number is the server's own, so a DML check does not have to read the table
/// back to know what changed — and `INSERT` reports `0 n`, which is why the scan
/// takes the trailing digits and not the first ones.
pub fn ResultSet::affected(self : ResultSet) -> Int {
  let tag = self.command
  // A command tag is ASCII, and on the native target `String` indexes by UTF-16
  // code unit, so the comparisons go through `to_int`.
  let mut stop = tag.length()
  while stop > 0 && tag[stop - 1].to_int() == 32 {
    stop = stop - 1
  }
  let mut start = stop
  while start > 0 {
    let c = tag[start - 1].to_int()
    if c < 48 || c > 57 {
      break
    }
    start = start - 1
  }
  if start == stop {
    0
  } else {
    let mut v = 0
    let mut i = start
    while i < stop {
      v = v * 10 + (tag[i].to_int() - 48)
      i = i + 1
    }
    v
  }
}

///|
pub struct Client {
  sock : @sys.Socket
  cfg : Config
  params : Map[String, String]
  mut backend_pid : Int
  mut transaction_status : Byte
  mut dialect : Dialect
}

///|
/// Turns an IO error into a raised message with its context.
fn[T] unwrap(r : Result[T, @sys.IOError], context : String) -> T raise {
  match r {
    Ok(v) => v
    Err(e) => fail("\{context}: \{e.message()}")
  }
}

///|
fn send_raw(c : Client, data : Bytes, context : String) -> Unit raise {
  let _ = unwrap(c.sock.write(data), context)
}

///|
fn send_message(
  c : Client,
  tag : Byte,
  body : Bytes,
  context : String,
) -> Unit raise {
  send_raw(c, @wire.header(tag, body.length()), context)
  send_raw(c, body, context)
}

///|
fn read_message(c : Client) -> @sys.Message raise {
  unwrap(c.sock.read_message(), "read")
}

///|
/// Text up to the next NUL: error fields, command tags, column names.
fn text_at(body : Bytes, offset : Int) -> String {
  @sys.text_at(body, offset)
}

///|
/// SCRAM payloads contain no embedded NUL, so they run to the end of the body.
fn text_to_end(body : Bytes, offset : Int) -> String {
  @sys.text_slice(body, offset, body.length() - offset)
}

///|
fn server_error_from(body : Bytes) -> ServerError {
  // ErrorReport is a sequence of  records terminated by
  // an empty record; severity, SQLSTATE and message are enough for reporting.
  let n = body.length()
  let mut i = 0
  let mut severity = ""
  let mut code = ""
  let mut message = ""
  while i < n {
    let field = body[i]
    if field == b'\x00' {
      break
    }
    let text = text_at(body, i + 1)
    if field == b'S' && severity == "" {
      severity = text
    } else if field == b'C' {
      code = text
    } else if field == b'M' {
      message = if message == "" { text } else { message + " " + text }
    }
    i = i + 1 + text.length() + 1
  }
  { severity, code, message, }
}

///|
/// Raises immediately when the message is an ErrorResponse.
fn expect_ok(m : @sys.Message) -> @sys.Message raise {
  if m.tag == b'E' {
    fail(server_error_from(m.body).describe())
  }
  m
}

///|
/// Consumes ParameterStatus / BackendKeyData until ReadyForQuery.
fn pump(c : Client) -> Unit raise {
  while true {
    let m = read_message(c)
    if m.tag == b'S' {
      let k = text_at(m.body, 0)
      c.params[k] = text_at(m.body, k.length() + 1)
    } else if m.tag == b'K' {
      c.backend_pid = @sys.be32(m.body, 0)
    } else if m.tag == b'E' {
      fail(server_error_from(m.body).describe())
    } else if m.tag == b'Z' {
      c.transaction_status = m.body[0]
      return
    }
  } nobreak {
    fail("unreachable")
  }
}

///| Connects, negotiates SSL, sends the startup packet, completes

///| authentication, then reads the dialect of the target.

///|

///| This KingbaseES answers `N` to SSLRequest and stays plaintext. A server that

///| answers `S` is refused rather than continued in plaintext, because a silent

///|
/// downgrade would be invisible to the caller.
pub fn connect(cfg : Config) -> Client raise {
  let sock = unwrap(@sys.connect(cfg.host, cfg.port), "connect")
  let start_mode = match cfg.fallback_mode {
    Some(m) => m
    None => @dialect.unknown_mode()
  }
  let c : Client = {
    sock,
    cfg,
    params: Map([]),
    backend_pid: 0,
    transaction_status: b'I',
    dialect: @dialect.new_dialect([], start_mode),
  }
  let w = @wire.new_writer(8)
  w.int32(8)
  w.int32(80877103)
  send_raw(c, w.payload(), "ssl request")
  let reply = unwrap(c.sock.read_byte(), "ssl reply")
  if reply != b'N' {
    fail(
      "server answered '\\{reply.to_char()}' to SSLRequest but this client is plaintext-only",
    )
  }
  startup(c)
  authenticate(c)
  pump(c)
  if cfg.detect_dialect {
    read_dialect(c, start_mode)
  }
  c
}

///|
/// Reads the settings that fix the dialect.

///|
/// Three stages, cheapest first, and each keeps whatever the previous one
/// missed: the mode alone, because every later query is shaped by it; one
/// catalog query for all names; then `show` for each name still missing.

///|
/// measured across all four modes: the `pg` mode ships no `sys_*` view at all
/// (`sys_settings` answers 42P01) and exposes `pg_settings`, while the other
/// three answer `sys_settings`. `sql_mode` and `ora_input_emptystr_isnull` are
/// hidden from the view wherever it exists. `quoted_identifier` does not exist in
/// `pg` mode and `DateFormat` not in `pg` or `mysql`, so one combined
/// `current_setting` list fails as a whole on a single undefined name — that is
/// why the last stage asks one name at a time. Reading all ten names separately
/// costs ten round trips of about a millisecond, and only on a mode with no view.
fn read_dialect(c : Client, fallback : Mode) -> Unit {
  let detected = read_mode(c)
  let mode = match detected {
    Some(m) => m
    None => fallback
  }
  let rows = collect_settings(c, mode)
  c.dialect = @dialect.new_dialect(rows, mode)
}

///|
/// The mode, as one setting. `show` answers it where `current_setting` does not.
fn read_mode(c : Client) -> Mode? {
  let attempts : Array[String] = [
    "select current_setting('database_mode')", "show database_mode",
  ]
  let mut found : Mode? = None
  for sql in attempts {
    match found {
      Some(_) => ()
      None =>
        match c.try_query(sql) {
          Ok(rs) => found = @dialect.mode_of(rs.scalar())
          Err(_) => ()
        }
    }
  }
  found
}

///|
fn collect_settings(c : Client, mode : Mode) -> Array[(String, String)] {
  let names = @dialect.setting_names()
  let mut rows : Array[(String, String)] = []
  for view in @dialect.settings_views(mode) {
    if rows.length() > 0 {
      break
    }
    match c.try_query(@dialect.sql_in_list(view)) {
      Ok(rs) => rows = rs.pairs()
      Err(_) => ()
    }
  }
  let have = Map([])
  for pair in rows {
    let (k, _) = pair
    have[k] = "seen"
  }
  for n in names {
    match have.get(n) {
      Some(_) => ()
      None => {
        let v = read_setting(c, n)
        if v != "" {
          rows.push((n, v))
        }
      }
    }
  }
  rows
}

///|
/// One setting through `show`, or an empty string when this mode has no such
/// setting.
fn read_setting(c : Client, name : String) -> String {
  match c.try_query("show " + name) {
    Ok(rs) => rs.scalar()
    Err(_) => ""
  }
}

///|
fn startup(c : Client) -> Unit raise {
  let w = @wire.new_writer(96)
  w.int32(196608) // protocol version 3.0
  w.cstring("user")
  w.cstring(c.cfg.user)
  w.cstring("database")
  w.cstring(c.cfg.database)
  w.cstring("client_encoding")
  w.cstring("UTF8")
  w.byte(b'\x00')
  let body = w.payload()
  // StartupMessage carries no tag byte; its length prefix covers itself plus the
  // version word and every parameter.
  let len = @wire.new_writer(4)
  len.int32(body.length() + 4)
  send_raw(c, len.payload(), "startup length")
  send_raw(c, body, "startup")
}

///|
fn authenticate(c : Client) -> Unit raise {
  while true {
    let m = expect_ok(read_message(c))
    if m.tag != b'R' {
      continue
    }
    let code = @sys.be32(m.body, 0)
    if code == 0 {
      return
    } else if code == 10 {
      sasl(c, m)
      return
    } else if code == 2 {
      // AuthenticationCleartextPassword
      let w = @wire.new_writer(c.cfg.password.length() + 8)
      w.text(c.cfg.password)
      w.byte(b'\x00')
      send_message(c, b'p', w.payload(), "password")
    } else {
      fail(
        "unsupported auth method \{code}; this client implements SCRAM-SHA-256 and cleartext",
      )
    }
  } nobreak {
    fail("unreachable")
  }
}

///|
fn sasl(c : Client, first_msg : @sys.Message) -> Unit raise {
  let mech = text_at(first_msg.body, 4)
  if mech != "SCRAM-SHA-256" {
    fail(
      "server offered SASL mechanism '\{mech}', only SCRAM-SHA-256 is supported",
    )
  }
  let state = scram_client_first(18)
  send_message(
    c,
    b'p',
    scram_initial_response(state.client_first),
    "sasl initial",
  )
  let cont = expect_ok(read_message(c))
  if cont.tag != b'R' || @sys.be32(cont.body, 0) != 11 {
    fail("expected SASLContinue")
  }
  let server_first = text_to_end(cont.body, 4)
  let reply = match scram_respond(state, server_first, c.cfg.password) {
    Ok(r) => r
    Err(e) => fail("scram: \{e}")
  }
  // SASLResponse carries raw SCRAM text: no NUL terminator (the server rejects
  // one as a malformed message).
  let w = @wire.new_writer(reply.client_final.length() + 8)
  w.text(reply.client_final)
  send_message(c, b'p', w.payload(), "sasl response")
  while true {
    let m = expect_ok(read_message(c))
    if m.tag == b'R' {
      let code = @sys.be32(m.body, 0)
      if code == 12 {
        let final_text = text_to_end(m.body, 4)
        if !final_text.has_prefix(scram_expected_server_final(reply)) {
          fail("server signature mismatch")
        }
      } else if code == 0 {
        return
      } else {
        fail("unexpected SASL message code \{code}")
      }
    } else if m.tag == b'Z' {
      return
    }
  } nobreak {
    fail("unreachable")
  }
}

///|
/// Sends a simple query and collects rows until ReadyForQuery.

///|
/// The simple protocol is the only form this client sends, because the target
/// rejects Parse and Bind with SQLSTATE 08P01. A mode that takes the extended
/// protocol reports it through `Dialect::extended_protocol_supported`.
pub fn Client::query(self : Client, sql : String) -> ResultSet raise {
  let w = @wire.new_writer(sql.length() + 8)
  w.cstring(sql)
  send_message(self, b'Q', w.payload(), "query")
  collect(self)
}

///|
/// Runs a statement, discards its rows and returns the command tag.
pub fn Client::execute(self : Client, sql : String) -> String raise {
  self.query(sql).command
}

///|
/// Runs a query and returns its rows, or the server error.

///|
/// A capability probe needs a result for every attempt, so one refusal must not
/// end the run. `query` raises the `describe` text of a `ServerError`, which
/// keeps the SQLSTATE in the leading brackets, so read it back from there.
pub fn Client::try_query(
  self : Client,
  sql : String,
) -> Result[ResultSet, ServerError] {
  let mut err : ServerError? = None
  let mut out : ResultSet? = None
  try self.query(sql) catch {
    @builtin.Failure(msg) => err = Some(error_from_text(msg))
    _ => err = Some({ severity: "", code: "", message: "unexpected failure", })
  } noraise {
    rs => out = Some(rs)
  }
  match err {
    Some(e) => Err(e)
    None =>
      match out {
        Some(rs) => Ok(rs)
        None => Err({ severity: "", code: "", message: "no result", })
      }
  }
}

///|
/// Runs a statement and returns its command tag, or the server error.
pub fn Client::try_execute(
  self : Client,
  sql : String,
) -> Result[String, ServerError] {
  match self.try_query(sql) {
    Ok(rs) => Ok(rs.command)
    Err(e) => Err(e)
  }
}

///|
fn error_from_text(msg : String) -> ServerError {
  // A raised `Failure` text starts with the source location of the `fail` call,
  // then the describe form "[SQLSTATE] SEVERITY: text". The code is searched
  // for, because the location part varies between call sites.
  match msg.find("[") {
    Some(i) if i + 7 <= msg.length() => {
      let code = msg[i + 1:i + 6].to_owned()
      let rest = msg[i + 7:].to_owned()
      { severity: "", code, message: rest, }
    }
    _ => { severity: "", code: "", message: msg, }
  }
}

///|
/// Sends one statement through the extended query protocol and reports what came
/// back.

///|
/// Parse, Bind, Describe, Execute and Sync go out as one batch, the way a driver
/// that supports prepared statements sends them. The success value is the sequence
/// of message tags the server sent before ReadyForQuery, so a refusal is visible
/// as its SQLSTATE instead of an empty result.

///|
/// This is a capability probe. The client itself runs statements through the simple
/// protocol, because that path needs one message per query and works on every mode
/// measured here.
pub fn Client::extended_try(
  self : Client,
  sql : String,
) -> Result[String, ServerError] raise {
  let parse = @wire.new_writer(sql.length() + 8)
  parse.cstring(sql)
  parse.int16(0) // no parameter type OIDs
  send_message(self, b'P', parse.payload(), "parse")
  let bind = @wire.new_writer(16)
  bind.cstring("") // unnamed portal
  bind.cstring("") // the statement Parse just defined
  bind.int16(0) // no parameter format codes
  bind.int16(0) // no parameter values
  bind.int16(0) // no result format codes, so every column is text
  send_message(self, b'B', bind.payload(), "bind")
  let describe = @wire.new_writer(8)
  describe.byte(b'S') // describe the prepared statement
  describe.cstring("")
  send_message(self, b'D', describe.payload(), "describe")
  let execute = @wire.new_writer(8)
  execute.cstring("") // run the unnamed portal
  execute.int32(0) // no row limit
  send_message(self, b'E', execute.payload(), "execute")
  send_message(self, b'S', @wire.new_writer(4).payload(), "sync")
  read_extended_reply(self)
}

///|
/// Reads replies up to ReadyForQuery. The first ErrorResponse, if any, becomes the
/// failure; otherwise the tags seen so far are returned.
fn read_extended_reply(c : Client) -> Result[String, ServerError] raise {
  let tags = @buffer.Buffer(size_hint=32)
  let mut err : ServerError? = None
  while true {
    let m = read_message(c)
    if m.tag == b'Z' {
      c.transaction_status = m.body[0]
      let text = @sys.buffer_text(tags)
      return match err {
        Some(e) => Err(e)
        None => Ok(text)
      }
    }
    if m.tag == b'E' {
      match err {
        None => err = Some(server_error_from(m.body))
        _ => ()
      }
    }
    // Message tags are ASCII, so the byte is the text.
    tags.write_byte(m.tag)
  } nobreak {
    fail("unreachable")
  }
}

///|
/// Reads messages until ReadyForQuery, discarding them.

///|
/// Used after an ErrorResponse, so a caller that catches the failure can keep
/// using the same connection instead of reading the tail of the refused
/// response as the next result.
fn drain_to_ready(c : Client) -> Unit raise {
  while true {
    let m = read_message(c)
    if m.tag == b'Z' {
      c.transaction_status = m.body[0]
      return
    }
  } nobreak {
    fail("unreachable")
  }
}

///|
fn collect(c : Client) -> ResultSet raise {
  let columns : Array[String] = []
  let rows : Array[Array[String?]] = []
  let mut command = ""
  while true {
    let m = read_message(c)
    if m.tag == b'T' {
      columns.clear()
      let n = @sys.be16(m.body, 0)
      let mut off = 2
      let mut i = 0
      while i < n {
        let name = text_at(m.body, off)
        columns.push(name)
        // name + NUL, then table OID(4) + column id(2) + type OID(4) +
        // typlen(2) + typmod(4) + format code(2) = 18 bytes of metadata
        off = off + name.length() + 1 + 18
        i = i + 1
      }
    } else if m.tag == b'D' {
      let n = @sys.be16(m.body, 0)
      let mut off = 2
      let row : Array[String?] = []
      let mut i = 0
      while i < n {
        let len = @sys.be32(m.body, off)
        off = off + 4
        if len < 0 {
          row.push(None)
        } else {
          row.push(Some(@sys.text_slice(m.body, off, len)))
          off = off + len
        }
        i = i + 1
      }
      rows.push(row)
    } else if m.tag == b'C' {
      command = text_at(m.body, 0)
    } else if m.tag == b'E' {
      // The server sends ReadyForQuery after the error. Read to it before
      // raising, or the next statement on this connection would read the tail
      // of this response and return the row count of the statement before it.
      // measured: without this drain, one refused statement shifted every
      // later result by one query.
      let e = server_error_from(m.body)
      drain_to_ready(c)
      fail(e.describe())
    } else if m.tag == b'Z' {
      c.transaction_status = m.body[0]
      let rs : ResultSet = { columns, rows, command, }
      return rs
    }
  } nobreak {
    fail("unreachable")
  }
}

///|
/// The dialect of the connected instance.
pub fn Client::dialect(self : Client) -> Dialect {
  self.dialect
}

///|
/// The compatibility mode of the connected instance.
pub fn Client::mode(self : Client) -> Mode {
  self.dialect.mode
}

///|
/// A `ParameterStatus` value the backend sent, e.g. `server_version`.
pub fn Client::setting(self : Client, key : String) -> String {
  match self.params.get(key) {
    Some(v) => v
    None => ""
  }
}

///|
/// `I` when the backend is idle, `T` inside a transaction, `E` after a failure.
pub fn Client::status_byte(self : Client) -> Byte {
  self.transaction_status
}

///|
/// The process id from `BackendKeyData`.
pub fn Client::pid(self : Client) -> Int {
  self.backend_pid
}

///|
/// Sends the terminate message and closes the socket.
pub fn Client::close(self : Client) -> Unit {
  let _ = self.sock.write(@wire.header(b'X', 0))
  self.sock.close()
}

///|
/// Monotonic microseconds, for timing queries and batches.

///|
/// This is the only clock a caller should use for timings: it does not move when
/// the wall clock is corrected.
pub fn now_us() -> Int64 {
  @sys.now_us()
}

///|
/// Text of everything written into a buffer, decoded as UTF-8.

///|
/// `Buffer::to_string` reinterprets bytes as code units on the native target, so
/// report and log text must decode explicitly.
pub fn buffer_text(buf : @buffer.Buffer) -> String {
  @sys.buffer_text(buf)
}

///|
/// Waits for `ms` milliseconds, e.g. between COPY rounds in a paced load.
pub fn sleep_ms(ms : Int) -> Unit {
  @sys.sleep_ms(ms)
}