#!/usr/bin/env python

"""
Copyright (c) 2006-2026 sqlmap developers (https://sqlmap.org)
See the file 'LICENSE' for copying permission
"""

"""
Minimal pure-python PostgreSQL frontend/backend protocol v3 client (stdlib only).

Covers the whole PostgreSQL-wire family (PostgreSQL, CockroachDB, CrateDB, Redshift, Greenplum, Vertica).
Auth: trust / cleartext / MD5 / SCRAM-SHA-256 (modern default). Uses the *simple query* protocol, whose
per-message implicit transaction auto-commits - so it is immune to the aborted-transaction poisoning and
commit-before-fetch pitfalls that bite the stateful native drivers. Binary (bytea) values arrive as the
server's readable '\\xHEX' text (text result format), so no memoryview/blob corruption either.
"""

import base64
import binascii
import hashlib
import hmac
import os
import socket
import struct

from extra.dbwire import DataError
from extra.dbwire import IntegrityError
from extra.dbwire import InterfaceError
from extra.dbwire import NotSupportedError
from extra.dbwire import OperationalError
from extra.dbwire import connection_lost
from extra.dbwire import handshake_done
from extra.dbwire import keepalive
from extra.dbwire import recvn
from extra.dbwire import ProgrammingError

_PROTOCOL_VERSION = 196608  # 3.0
_MAX_MESSAGE_LENGTH = 0x40000000  # 1 GB - guard against a hostile/corrupt length triggering an unbounded read
_OID_BYTEA = 17  # bytea arrives as the server's text form; decode to bytes so it hexes like the native driver

def _decode_bytea(raw):
    # PG text output: modern 'hex' = b'\\x<hexdigits>'; legacy 'escape' = octal \ooo + literal bytes
    if raw[:2] == b"\\x":
        try:
            return binascii.unhexlify(raw[2:])
        except (binascii.Error, ValueError):
            return raw.decode("utf-8", "replace")
    src, out, i, n = bytearray(raw), bytearray(), 0, len(raw)
    while i < n:
        if src[i] == 0x5c and i + 1 < n:  # backslash
            nxt = src[i + 1]
            if nxt == 0x5c:
                out.append(0x5c); i += 2
            elif 0x30 <= nxt <= 0x37 and i + 3 < n:  # \ooo octal
                out.append(((nxt - 48) << 6) | ((src[i + 2] - 48) << 3) | (src[i + 3] - 48)); i += 4
            else:
                out.append(nxt); i += 2
        else:
            out.append(src[i]); i += 1
    return bytes(out)

# SQLSTATE class (first 2 chars) -> DB-API exception, so callers can distinguish (mirrors psycopg2)
_SQLSTATE_CLASS = {
    "22": DataError, "23": IntegrityError,
    "08": OperationalError, "28": OperationalError, "53": OperationalError,
    "57": OperationalError, "58": OperationalError,
}

def _xor(a, b):
    # byte-wise XOR of two equal-length byte strings (Python 2 and 3 safe)
    if str is bytes:  # Python 2: iterating bytes yields 1-char strings
        return b"".join(chr(ord(x) ^ ord(y)) for x, y in zip(a, b))
    return bytes(x ^ y for x, y in zip(a, b))

def _read_message(sock):
    mtype = recvn(sock, 1)
    (length,) = struct.unpack("!I", recvn(sock, 4))
    if length < 4 or length > _MAX_MESSAGE_LENGTH:
        raise InterfaceError("invalid backend message length (%d)" % length)
    return mtype, recvn(sock, length - 4)

def _send(sock, mtype, payload):
    try:
        sock.sendall((mtype or b"") + struct.pack("!I", len(payload) + 4) + payload)
    except (socket.error, OSError) as ex:
        raise connection_lost(ex)

def _error_message(payload):
    # ErrorResponse/NoticeResponse: series of (byte field-code, cstring value), terminated by a NUL byte.
    # Returns (human message, SQLSTATE). Tolerant of a truncated/unterminated stream (find() not index()).
    fields, off = {}, 0
    while off < len(payload) and payload[off:off + 1] != b"\x00":
        code = payload[off:off + 1]
        end = payload.find(b"\x00", off + 1)
        if end == -1:
            break
        fields[code] = payload[off + 1:end].decode("utf-8", "replace")
        off = end + 1
    return fields.get(b"M", "unknown error"), fields.get(b"C", "")

def _raise_server_error(message, sqlstate):
    raise _SQLSTATE_CLASS.get((sqlstate or "")[:2], ProgrammingError)("(remote) %s" % message)

class Cursor(object):
    def __init__(self, connection):
        self.connection = connection
        self.description = None
        self.rowcount = -1
        self._rows = []
        self._pos = 0

    def execute(self, query, params=None):
        if params is not None:
            raise NotSupportedError("parameter binding is not supported; pass a fully-formed query string")
        self.description, self.rowcount, self._rows, self._pos = None, -1, [], 0  # reset before (a failed) query
        self.description, self._rows, self._pos, self.rowcount = self.connection._simple_query(query)
        return self

    def fetchall(self):
        retVal = self._rows[self._pos:]
        self._pos = len(self._rows)
        return retVal

    def fetchone(self):
        if self._pos >= len(self._rows):
            return None
        retVal = self._rows[self._pos]
        self._pos += 1
        return retVal

    def close(self):
        self._rows = []

class Connection(object):
    def __init__(self, sock):
        self._sock = sock
        self._txn_status = b"I"  # last ReadyForQuery transaction status: I(dle) / T(ransaction) / E(rror)

    def cursor(self):
        return Cursor(self)

    def commit(self):
        pass  # simple-query protocol commits each statement implicitly

    def rollback(self):
        pass

    def close(self):
        try:
            _send(self._sock, b"X", b"")  # Terminate
        except Exception:
            pass
        try:
            self._sock.close()
        except Exception:
            pass

    def _clear_aborted(self):
        # a prior statement left an aborted transaction block ('E'): every further statement errors with
        # 25P02 until it is rolled back. Clear it so the reused connection recovers (psycopg2 rollback semantics).
        _send(self._sock, b"Q", b"ROLLBACK\x00")
        while True:
            mtype, payload = _read_message(self._sock)
            if mtype == b"Z":
                self._txn_status = payload[:1] or b"I"
                break

    def _simple_query(self, query):
        if self._txn_status == b"E":
            self._clear_aborted()
        _send(self._sock, b"Q", query.encode("utf-8") + b"\x00")

        description, rows, rowcount, error = None, [], -1, None
        while True:
            mtype, payload = _read_message(self._sock)
            try:
                if mtype == b"T":  # RowDescription (a new result set: reset rows so we return only the last one)
                    (count,) = struct.unpack("!H", payload[:2])
                    description, rows, rowcount, off = [], [], -1, 2
                    for _ in range(count):
                        end = payload.index(b"\x00", off)
                        name = payload[off:end].decode("utf-8", "replace")
                        off = end + 1
                        (typeoid,) = struct.unpack("!I", payload[off + 6:off + 10])
                        off += 18  # tableoid4 colno2 typeoid4 typelen2 typmod4 format2
                        description.append((name, typeoid, None, None, None, None, None))
                elif mtype == b"D":  # DataRow
                    (count,) = struct.unpack("!H", payload[:2])
                    off, row = 2, []
                    for col in range(count):
                        (vlen,) = struct.unpack("!i", payload[off:off + 4])
                        off += 4
                        if vlen == -1:
                            row.append(None)
                        else:
                            if off + vlen > len(payload):
                                raise InterfaceError("truncated DataRow")
                            raw = payload[off:off + vlen]
                            off += vlen
                            if description and col < len(description) and description[col][1] == _OID_BYTEA:
                                row.append(_decode_bytea(raw))  # bytes so sqlmap hex-encodes it (like the native driver)
                            else:
                                try:
                                    row.append(raw.decode("utf-8"))
                                except UnicodeDecodeError:
                                    row.append(raw)  # non-UTF-8 (e.g. a SQL_ASCII db): keep bytes (hex-encoded), not lossy U+FFFD
                    rows.append(tuple(row))
                elif mtype == b"C":  # CommandComplete ("SELECT 3", "INSERT 0 1", ...)
                    tag = payload[:-1].decode("utf-8", "replace").split()
                    if tag and tag[-1].isdigit():
                        rowcount = int(tag[-1])
                elif mtype == b"G":  # CopyInResponse - server now waits for client CopyData; refuse to avoid a deadlock
                    _send(self._sock, b"f", b"COPY FROM STDIN is not supported\x00")  # CopyFail
                elif mtype == b"E":  # ErrorResponse
                    error = _error_message(payload)
                elif mtype == b"Z":  # ReadyForQuery (end of response); payload byte = transaction status
                    self._txn_status = payload[:1] or b"I"
                    break
                # ParameterStatus(S)/NoticeResponse(N)/EmptyQueryResponse(I)/CopyData(d)/CopyDone(c)/... ignored
            except (struct.error, IndexError, ValueError) as ex:
                raise InterfaceError("malformed backend message: %s" % ex)
        if error is not None:
            _raise_server_error(*error)
        return description, rows, 0, rowcount

def _authenticate(sock, user, password):
    client_nonce = cfirst_bare = salted = auth_message = None
    while True:
        mtype, payload = _read_message(sock)
        if mtype in (b"N", b"S"):  # NoticeResponse / ParameterStatus may legally precede AuthenticationOk
            continue
        if mtype == b"E":
            _raise_server_error_as_operational(payload)
        if mtype != b"R":
            raise InterfaceError("unexpected message %r during authentication" % mtype)
        (code,) = struct.unpack("!I", payload[:4])
        if code == 0:  # AuthenticationOk (also the trust case)
            return
        elif code == 3:  # cleartext password
            _send(sock, b"p", (password or "").encode("utf-8") + b"\x00")
        elif code == 5:  # MD5 password
            salt = payload[4:8]
            inner = hashlib.md5((password or "").encode("utf-8") + (user or "").encode("utf-8")).hexdigest()
            token = b"md5" + hashlib.md5(inner.encode("ascii") + salt).hexdigest().encode("ascii")
            _send(sock, b"p", token + b"\x00")
        elif code == 10:  # SASL (SCRAM-SHA-256)
            if not hasattr(hashlib, "pbkdf2_hmac"):
                raise NotSupportedError("SCRAM-SHA-256 authentication requires Python >= 2.7.8 (hashlib.pbkdf2_hmac)")
            client_nonce = base64.b64encode(os.urandom(18)).decode("ascii")
            cfirst_bare = "n=,r=%s" % client_nonce
            client_first = "n,," + cfirst_bare
            _send(sock, b"p", b"SCRAM-SHA-256\x00" + struct.pack("!I", len(client_first)) + client_first.encode("ascii"))
        elif code == 11:  # SASLContinue (server-first)
            try:
                server_first = payload[4:].decode("ascii")
                attrs = dict(kv.split("=", 1) for kv in server_first.split(","))
                snonce, salt, iterations = attrs["r"], base64.b64decode(attrs["s"]), int(attrs["i"])
            except (KeyError, ValueError, binascii.Error, UnicodeDecodeError) as ex:
                raise OperationalError("malformed SCRAM server-first message (%s)" % ex)
            # RFC 5802 5.1: the server nonce MUST start with the client nonce and MUST add material of its
            # own. Skipping this lets anything that can answer the TCP connection replay a recorded
            # server-first and drive the exchange - and dbwire has no TLS layer underneath to catch it.
            if not client_nonce or not snonce.startswith(client_nonce) or len(snonce) <= len(client_nonce):
                raise OperationalError("SCRAM server nonce does not extend the client nonce (rogue server?)")
            if iterations < 4096:   # RFC 5802 recommends >= 4096; a tiny count cheapens an offline attack
                raise OperationalError("SCRAM iteration count %d is too low" % iterations)
            salted = hashlib.pbkdf2_hmac("sha256", (password or "").encode("utf-8"), salt, iterations)
            client_key = hmac.new(salted, b"Client Key", hashlib.sha256).digest()
            stored_key = hashlib.sha256(client_key).digest()
            client_final_noproof = "c=biws,r=%s" % snonce
            auth_message = "%s,%s,%s" % (cfirst_bare, server_first, client_final_noproof)
            client_sig = hmac.new(stored_key, auth_message.encode("ascii"), hashlib.sha256).digest()
            proof = base64.b64encode(_xor(client_key, client_sig)).decode("ascii")
            _send(sock, b"p", ("%s,p=%s" % (client_final_noproof, proof)).encode("ascii"))
        elif code == 12:  # SASLFinal (server-final): verify the server too, or the handshake is one-way
            # Without this the client proves itself to the server and simply trusts whatever answers back.
            # ServerSignature = HMAC(ServerKey, AuthMessage) can only be produced by a peer that holds the
            # stored credentials, so comparing it is what makes the exchange mutual (RFC 5802 5, 5.1).
            if salted is None or auth_message is None:
                raise OperationalError("unexpected SCRAM server-final message")
            try:
                attrs = dict(kv.split("=", 1) for kv in payload[4:].decode("ascii").split(","))
            except (ValueError, UnicodeDecodeError) as ex:
                raise OperationalError("malformed SCRAM server-final message (%s)" % ex)
            if "e" in attrs:
                raise OperationalError("SCRAM authentication failed (%s)" % attrs["e"])
            try:
                signature = base64.b64decode(attrs["v"])
            except (KeyError, binascii.Error, ValueError) as ex:
                raise OperationalError("malformed SCRAM server signature (%s)" % ex)
            server_key = hmac.new(salted, b"Server Key", hashlib.sha256).digest()
            expected = hmac.new(server_key, auth_message.encode("ascii"), hashlib.sha256).digest()
            if not hmac.compare_digest(signature, expected):
                raise OperationalError("SCRAM server signature mismatch (rogue server?)")
        else:
            raise InterfaceError("unsupported authentication request %d" % code)

def _raise_server_error_as_operational(payload):
    message, _ = _error_message(payload)
    raise OperationalError("(remote) %s" % message)

def connect(host=None, port=5432, user=None, password=None, database=None, connect_timeout=None, **kwargs):
    try:
        sock = socket.create_connection((host or "localhost", int(port or 5432)), timeout=connect_timeout)
        keepalive(sock)
    except (socket.error, socket.timeout) as ex:
        raise OperationalError("could not connect to '%s:%s' (%s)" % (host, port, ex))

    params = b""
    for key, value in (("user", user or ""), ("database", database or user or ""), ("client_encoding", "UTF8")):
        params += key.encode("ascii") + b"\x00" + ("%s" % value).encode("utf-8") + b"\x00"
    params += b"\x00"

    try:
        _send(sock, b"", struct.pack("!I", _PROTOCOL_VERSION) + params)  # StartupMessage
        _authenticate(sock, user, password)
        while True:  # drain until ReadyForQuery (ParameterStatus/BackendKeyData/NoticeResponse)
            mtype, payload = _read_message(sock)
            if mtype == b"E":
                _raise_server_error_as_operational(payload)
            if mtype == b"Z":
                break
        handshake_done(sock)
    except Exception:  # any setup failure (DB-API or otherwise) must still close the socket
        try:
            sock.close()
        except Exception:
            pass
        raise

    return Connection(sock)
