"""Bounded HASHBATTLE engines. No wallet, signing, RPC or private dependencies.

CpuEngine(kit_dir) and GpuEngine(kit_dir) expose search(work, start_nonce)
-> (hit_or_None, hashes_done). A hit is {"nonce": int, "hash": 64 hex chars}
without 0x. The caller owns work freshness and advances by hashes_done.
"""
from pathlib import Path
import json
import os
import string
import subprocess
import tempfile

UINT64_MAX = (1 << 64) - 1
CPU_MAX_HASHES = 1 << 20


def _nonce(value):
    if type(value) is not int or not 0 <= value <= UINT64_MAX:
        raise ValueError("nonce must be an integer in [0, 2**64 - 1]")
    return value


def _work_bytes(work):
    result = []
    for field, size in (("miner", 20), ("prevHex", 32), ("anchorHash", 32), ("targetHex", 32)):
        value = work.get(field)
        if (not isinstance(value, str) or not value.startswith("0x")
                or len(value) != 2 + size * 2
                or any(c not in string.hexdigits for c in value[2:])):
            raise ValueError(f"{field} must be 0x-prefixed {size}-byte hex")
        result.append(bytes.fromhex(value[2:]))
    return tuple(result)


class CpuEngine:
    """Bundled proven C solver; gcc needed on first use (or source updates).

    max_hashes is optional for small test/interactive windows. Defaults to 2**20.
    No Python crypto and no GPU import are needed in CPU mode.
    """

    def __init__(self, kit_dir, *, max_hashes=CPU_MAX_HASHES):
        if type(max_hashes) is not int or not 1 <= max_hashes <= UINT64_MAX:
            raise ValueError("max_hashes must be a positive uint64 integer")
        self.kit_dir = Path(kit_dir).resolve()
        self.solver = self.kit_dir / "keccak256_miner"
        self.max_hashes = max_hashes

    def _build(self):
        source = self.kit_dir / "keccak256_miner.c"
        if (self.solver.is_file() and os.access(self.solver, os.X_OK)
                and self.solver.stat().st_mtime_ns >= source.stat().st_mtime_ns):
            return
        # Atomic installation also permits separate miner processes to start together.
        fd, name = tempfile.mkstemp(prefix=".keccak-build-", dir=self.kit_dir)
        os.close(fd)
        temporary = Path(name)
        try:
            command = ["gcc", "-O3", "-march=native", "-std=c11", "-D_POSIX_C_SOURCE=200809L",
                       "-Wall", "-Wextra", "-Werror", "-pedantic", str(source), "-o", str(temporary)]
            result = subprocess.run(command, capture_output=True, text=True, check=False)
            if result.returncode:
                raise RuntimeError(f"C solver build failed: {result.stderr.strip()}")
            temporary.chmod(0o700)
            temporary.replace(self.solver)
        finally:
            temporary.unlink(missing_ok=True)

    def digests(self, work, nonces):
        """Batch reference digests from the C --test mode; used to verify GPU work."""
        miner, previous, anchor, _target = _work_bytes(work)
        preimages = [miner + _nonce(n).to_bytes(32, "big") + previous + anchor for n in nonces]
        if not preimages:
            return []
        self._build()
        result = subprocess.run([str(self.solver), "--test"],
                                input="".join(p.hex() + "\n" for p in preimages),
                                capture_output=True, text=True, check=False)
        lines = result.stdout.splitlines()
        if result.returncode or len(lines) != len(preimages):
            raise RuntimeError(f"C reference hash failed: {result.stderr.strip()}")
        return [bytes.fromhex(line) for line in lines]

    def search(self, work, start_nonce):
        """Scan at most max_hashes, clipped at uint64 end; a hit stops C early."""
        start_nonce = _nonce(start_nonce)
        _work_bytes(work)
        count = min(self.max_hashes, UINT64_MAX - start_nonce + 1)
        self._build()
        with tempfile.TemporaryDirectory(prefix="hb-cpu-") as directory:
            outfile = Path(directory) / "result.json"
            command = [str(self.solver), work["miner"], work["prevHex"], work["anchorHash"],
                       work["targetHex"], str(start_nonce), str(outfile), str(count)]
            result = subprocess.run(command, capture_output=True, text=True, check=False)
            if result.returncode not in (0, 1):
                raise RuntimeError(f"C solver failed: {result.stderr.strip()}")
            output = json.loads(outfile.read_text())
        if result.returncode == 1:
            return None, output["hashes"]
        return {"nonce": output["nonce"], "hash": output["hash"]}, output["hashes"]


SLOTS = 65536
# Short dispatches are friendlier to a co-tenant GPU than the operator's batch=1024.
GPU_BATCH = 16
_SELF_TEST_NONCE = (1 << 32) + 123
_SELF_TEST_WORK = {
    "miner": "0x1234567890abcdef1234567890abcdef12345678",
    "prevHex": "0x" + bytes(range(32)).hex(),
    "anchorHash": "0x" + bytes(range(32, 64)).hex(),
    "targetHex": "0x" + "00" * 32,
}
_SELF_TEST_HASH = "edcec6584ac3b3e3121ac8a8ac6022d2ddd5f320ad45af26a019b352c49fc5ff"


def _hardware_adapter(adapters):
    """Only positive hardware type identification; never a fallback request."""
    eligible = []
    for adapter in adapters:
        info = adapter.info
        kind = str(info.get("adapter_type", "")).lower().replace("_", "").replace("-", "")
        details = (str(info) + " " + str(adapter.summary)).lower()
        if kind not in ("discretegpu", "integratedgpu"):
            continue
        if any(word in details for word in (
                "llvmpipe", "lavapipe", "swiftshader", "softpipe", "software",
                "microsoft basic render", "warp")):
            continue
        eligible.append((0 if kind == "discretegpu" else 1, adapter))
    if not eligible:
        raise RuntimeError("No hardware GPU available (CPU/software/unknown adapters refused)")
    return min(eligible, key=lambda item: item[0])[1]


class GpuEngine:
    """Real discrete/integrated wgpu compute, using the proven public WGSL kernel.

    Requires wgpu >= 0.31 and a working hardware driver. Software adapters are
    refused. C verifies digest0 on every dispatch and rehashes all candidates.
    A startup nonce > 2**32 tests both halves before the first real search.
    Optional batch is 1..64 (default 16); each call scans SLOTS * batch hashes,
    clipped at uint64 end. Calls on an engine instance must be sequential.
    """

    def __init__(self, kit_dir, *, batch=GPU_BATCH):
        if type(batch) is not int or not 1 <= batch <= 64:
            raise ValueError("GPU batch must be an integer in [1, 64]")
        try:
            import wgpu
        except ImportError as error:
            raise RuntimeError("GPU mode requires wgpu>=0.31; install the GPU dependency") from error
        self.wgpu = wgpu
        self.kit_dir = Path(kit_dir).resolve()
        self.batch = batch
        self.reference = CpuEngine(self.kit_dir)
        self.adapter = _hardware_adapter(wgpu.gpu.enumerate_adapters_sync())
        print(f"GPU adapter: {self.adapter.summary} | {dict(self.adapter.info)}", flush=True)
        self.device = self.adapter.request_device_sync()
        module = self.device.create_shader_module(code=(self.kit_dir / "gpu_miner.wgsl").read_text())
        self.pipeline = self.device.create_compute_pipeline(
            layout="auto", compute={"module": module, "entry_point": "mine"})
        usage = wgpu.BufferUsage
        self.out_buf = self.device.create_buffer(
            size=SLOTS * 16, usage=usage.STORAGE | usage.COPY_SRC)
        self.dbg_buf = self.device.create_buffer(size=32, usage=usage.STORAGE | usage.COPY_SRC)
        self.out_read = self.device.create_buffer(
            size=SLOTS * 16, usage=usage.COPY_DST | usage.MAP_READ)
        self.dbg_read = self.device.create_buffer(size=32, usage=usage.COPY_DST | usage.MAP_READ)
        # Validate the C reference itself, then compare this real shader's digest0.
        if self.reference.digests(_SELF_TEST_WORK, [_SELF_TEST_NONCE])[0].hex() != _SELF_TEST_HASH:
            raise RuntimeError("C reference self-test mismatch")
        self._checked_search(_SELF_TEST_WORK, _SELF_TEST_NONCE, 1)

    def _read(self, buffer):
        buffer.map_sync(self.wgpu.MapMode.READ)
        try:
            return bytes(buffer.read_mapped())
        finally:
            buffer.unmap()

    def _dispatch(self, work, start_nonce, count):
        """One bounded GPU dispatch; no Python hashing or CPU mining fallback."""
        import struct
        account, previous, anchor, target = _work_bytes(work)
        active_slots = (count + self.batch - 1) // self.batch
        params = [0] * 34
        params[0:5] = [int.from_bytes(account[i:i+4], "little") for i in range(0, 20, 4)]
        params[5:13] = [int.from_bytes(previous[i:i+4], "little") for i in range(0, 32, 4)]
        params[13:21] = [int.from_bytes(anchor[i:i+4], "little") for i in range(0, 32, 4)]
        params[21:29] = [int.from_bytes(target[i:i+4], "big") for i in range(0, 32, 4)]
        params[29] = start_nonce & 0xffffffff
        params[30] = start_nonce >> 32
        params[31] = self.batch
        params[32] = count
        params[33] = active_slots
        raw = struct.pack("<34I", *params)
        usage = self.wgpu.BufferUsage
        param_buf = self.device.create_buffer(size=len(raw), usage=usage.STORAGE | usage.COPY_DST)
        try:
            self.device.queue.write_buffer(param_buf, 0, raw)
            bindings = [(0, param_buf, len(raw)), (1, self.out_buf, SLOTS * 16),
                        (2, self.dbg_buf, 32)]
            group = self.device.create_bind_group(
                layout=self.pipeline.get_bind_group_layout(0), entries=[
                    {"binding": index, "resource": {"buffer": buffer, "offset": 0, "size": size}}
                    for index, buffer, size in bindings])
            encoder = self.device.create_command_encoder()
            compute = encoder.begin_compute_pass()
            compute.set_pipeline(self.pipeline)
            compute.set_bind_group(0, group)
            compute.dispatch_workgroups((active_slots + 63) // 64)
            compute.end()
            encoder.copy_buffer_to_buffer(self.out_buf, 0, self.out_read, 0, SLOTS * 16)
            encoder.copy_buffer_to_buffer(self.dbg_buf, 0, self.dbg_read, 0, 32)
            self.device.queue.submit([encoder.finish()])
            output = self._read(self.out_read)
            debug = struct.unpack("<8I", self._read(self.dbg_read))
        finally:
            param_buf.destroy()
        digest0 = b"".join(word.to_bytes(4, "big") for word in debug)
        candidates = []
        # Only active slots are read: untouched tails from a prior dispatch are ignored.
        for low, high, found, _reserved in struct.iter_unpack("<4I", output[:active_slots * 16]):
            if found:
                candidates.append((high << 32) | low)
        return candidates, digest0

    def _checked_search(self, work, start_nonce, count):
        candidates, digest0 = self._dispatch(work, start_nonce, count)
        if any(n < start_nonce or n >= start_nonce + count for n in candidates):
            raise RuntimeError("GPU candidate outside dispatched nonce window")
        # One batched C process, including all hits, not one subprocess per slot.
        digests = self.reference.digests(work, [start_nonce, *candidates])
        if digest0 != digests[0]:
            raise RuntimeError("GPU self-test digest0 mismatch")
        target = _work_bytes(work)[3]
        hit = None
        for nonce, digest in zip(candidates, digests[1:]):
            if digest < target and (hit is None or nonce < hit["nonce"]):
                hit = {"nonce": nonce, "hash": digest.hex()}
        return hit, count

    def search(self, work, start_nonce):
        """Complete one GPU window (even with a hit); hashes_done is exact."""
        start_nonce = _nonce(start_nonce)
        _work_bytes(work)
        count = min(SLOTS * self.batch, UINT64_MAX - start_nonce + 1)
        return self._checked_search(work, start_nonce, count)
