from __future__ import annotations
from typing import cast, Callable, Type, TypeVar, Generic, Any
import contextlib, decimal, statistics, time, ctypes, array, collections, itertools
from tinygrad.helpers import PROFILE, getenv, from_mv, cpu_profile, ProfileRangeEvent, unwrap
from tinygrad.helpers import suppress_finalizing, TracingKey
from tinygrad.device import BufferStorage, Buffer, BufferSpec, Compiled, Allocator, ProfileDeviceEvent, ProfileProgramEvent, Program, TinyELF
from tinygrad.uop.ops import sym_infer, sint, UOp
from tinygrad.runtime.support.memory import BumpAllocator, MMIOInterface
from tinygrad.renderer import Renderer

SignalType = TypeVar('SignalType', bound='HCQSignal')
HCQDeviceType = TypeVar('HCQDeviceType', bound='HCQCompiled')
ProgramType = TypeVar('ProgramType', bound='HCQProgram')
ArgsStateType = TypeVar('ArgsStateType', bound='HCQArgsState')

class HCQBuffer:
  def __init__(self, va_addr:sint, size:int, meta:Any=None, _base:HCQBuffer|None=None, view:MMIOInterface|None=None, owner:Any=None):
    self.va_addr, self.size, self.meta, self._base, self.view = va_addr, size, meta, _base, view
    self._devs, self.owner = ([owner] if owner is not None else []), owner
    self._mappings:dict[Compiled, HCQBuffer] = {} # mapping to the other devices

  def offset(self, offset:int=0, size:int|None=None) -> HCQBuffer:
    return HCQBuffer(self.va_addr+offset, size or (self.size - offset), owner=self.owner, meta=self.meta,
      _base=self._base or self, view=(self.view.view(offset=offset, size=size) if self.view is not None else None))

  def cpu_view(self) -> MMIOInterface:
    assert self.view is not None, "buffer has no cpu_view"
    return self.view

  @property
  def base(self) -> HCQBuffer: return self._base or self

  @property
  def mappings(self): return self._mappings if self._base is None else self._base._mappings

  @property
  def mapped_devs(self): return self._devs if self._base is None else self._base._devs

class HWQueue(Generic[SignalType, HCQDeviceType, ProgramType, ArgsStateType]):
  """
  A base class for hardware command queues in the HCQ (Hardware Command Queue) API.
  """

  def __init__(self):
    self._q:Any = []
    self.binded_device:HCQDeviceType|None = None
    self.q_sints:list[tuple[int, int]] = []
    self.mv_sints:list[tuple[MMIOInterface, int, int, int|None]] = []
    self.syms:list[sint] = []
    self._prev_resolved_syms:list[int|None] = []

  def _new_sym(self, sym:sint) -> int:
    if sym not in self.syms:
      self.syms.append(sym)
      self._prev_resolved_syms.append(None)
    return self.syms.index(sym)

  def q(self, *values):
    """
    Enqueues values in the queue.

    Args:
      values: The values to enqueue in the queue.
    """

    for v in values:
      if isinstance(v, UOp):
        self.q_sints.append((len(self._q), self._new_sym(v)))
        self._q.append(0xbadc0ded)
      else: self._q.append(v)

  # *** common commands  ***

  def timestamp(self, signal:SignalType):
    """
    Enqueues a timestamp command which records the current time in a signal after all previously enqueued commands are completed.

    Args:
      signal: The signal to store the timestamp
    """

  def signal(self, signal:SignalType, value:sint):
    """
    Enqueues a signal command which sets the signal to the given value, ensuring all previous operations are completed.

    Args:
      signal: The signal to set
      value: The value to set the signal to
    """

  def wait(self, signal:SignalType, value:sint):
    """
    Enqueues a wait command which halts execution until the signal is greater than or equal to a specific value.

    Args:
      signal: The signal to wait on
      value: The value to wait for
    """

  # *** commands for compute queues ***

  def memory_barrier(self):
    """
    Enqueues a memory barrier command to ensure memory coherence between agents. Only on compute queues.
    """

  def exec(self, prg:ProgramType, args_state:ArgsStateType, global_size:tuple[sint, ...], local_size:tuple[sint, ...]):
    """
    Enqueues an execution command for a kernel program. Only on compute queues.

    Args:
      prg: The program to execute
      args_state: The args state to execute program with
      global_size: The global work size
      local_size: The local work size
    """

  def write(self, b:HCQBuffer, val:sint, b64:bool=False):
    """
    Enqueues a command to write a value to a buffer address after all previously enqueued commands are completed.

    Args:
      b: The buffer to write to
      val: The value to write
      b64: If True, write a 64-bit value; otherwise write 32-bit
    """
    raise NotImplementedError("write not implemented")

  def poll_bit(self, b:HCQBuffer, val:sint, mask:int):
    """
    Enqueues a poll command which halts execution until (mem[b] & mask) == val.
    val must be 0 or mask (i.e. checks if masked bits are all clear or all set).

    Args:
      b: The buffer to poll
      val: The expected value after masking (0 or mask)
      mask: The bit mask to test
    """
    raise NotImplementedError("poll_bit not implemented")

  # *** commands for copy queues ***

  def copy(self, dest:HCQBuffer, src:HCQBuffer, copy_size:int):
    """
    Enqueues a copy command to transfer data. Only on copy queues.

    Args:
      dest: The destination buffer of the copy
      src: The source buffer of the copy
      copy_size: The size of data to copy
    """

  # *** submit and bind commands  ***

  def bind(self, dev:HCQDeviceType):
    """
    Associates the queue with a specific device for optimized execution.

    This optional method allows backend implementations to tailor the queue for efficient use on the given device. When implemented, it can eliminate
    the need to copy queues into the device, thereby enhancing performance.

    Args:
      dev: The target device for queue optimization.

    Note:
      Implementing this method is optional but recommended for performance gains.
    """

  def bind_args_state(self, args_state:ArgsStateType):
    for vals, mem, fmt in args_state.bind_data: self.bind_sints_to_mem(*vals, mem=mem, fmt=fmt)

  def bind_sints(self, *vals:sint, mem:MMIOInterface, struct_t:Type[ctypes.Structure], start_field:str, fmt, mask:int|None=None):
    self.bind_sints_to_mem(*vals, mem=mem, fmt=fmt, mask=mask, offset=getattr(struct_t, start_field).offset)

  def bind_sints_to_mem(self, *vals:sint, mem:MMIOInterface, fmt, mask:int|None=None, offset:int=0):
    mv = mem.view(offset=offset, size=len(vals)*8, fmt=fmt)
    for i, val in enumerate(vals):
      if isinstance(val, int): mv[i] = val if mask is None else ((mv[i] & ~mask) | val)
      else: self.mv_sints.append((mv, i, self._new_sym(val), mask))

  def _apply_var_vals(self, var_vals:dict[str, int]):
    resolved_syms: list[int|None] = [sym_infer(sym, var_vals) for sym in self.syms]

    for off, sym_idx in self.q_sints:
      if self._prev_resolved_syms[sym_idx] == resolved_syms[sym_idx]: continue
      self._q[off] = resolved_syms[sym_idx]

    for mv, off, sym_idx, mask in self.mv_sints:
      if self._prev_resolved_syms[sym_idx] == resolved_syms[sym_idx]: continue
      mv[off] = resolved_syms[sym_idx] if mask is None else ((mv[off] & ~mask) | resolved_syms[sym_idx])

    self._prev_resolved_syms = resolved_syms

  def submit(self, dev:HCQDeviceType, var_vals:dict[str, int]|None=None):
    """
    Submits the command queue to a specific device for execution.

    Args:
      dev: The device to submit the queue to
    """

    if var_vals is not None: self._apply_var_vals(var_vals)
    self._submit(dev)
    return self
  def _submit(self, dev:HCQDeviceType): raise NotImplementedError("need _submit")

class HCQSignal(Generic[HCQDeviceType]):
  def __init__(self, base_buf:HCQBuffer, value:int=0, owner:HCQDeviceType|None=None, is_timeline:bool=False, timestamp_divider=1000, virt=False):
    self.base_buf, self.owner, self.is_timeline = base_buf, owner, is_timeline
    self.should_return = isinstance(self.base_buf.va_addr, int) and self.owner is not None and not virt
    self.timestamp_divider:decimal.Decimal = decimal.Decimal(timestamp_divider)
    if isinstance(self.base_buf.va_addr, int) and not virt: self.value = value

  def __del__(self):
    if self.should_return: HCQCompiled.signal_pool[unwrap(self.owner).peer_group].append(self.base_buf)

  @property
  def value_addr(self) -> sint: return self.base_buf.va_addr

  @property
  def timestamp_addr(self) -> sint: return self.base_buf.va_addr + 8

  @property
  def value(self) -> int: return self.base_buf.cpu_view().view(0, 8, 'Q')[0]

  @value.setter
  def value(self, new_value:int): self.base_buf.cpu_view().view(0, 8, 'Q')[0] = new_value

  @property
  def timestamp(self) -> decimal.Decimal:
    """
    Get the timestamp field of the signal.

    This property provides read-only access to the signal's timestamp.

    Returns:
      The timestamp in microseconds.
    """
    return self.base_buf.cpu_view().view(8, 8, 'Q')[0] / self.timestamp_divider

  def _sleep(self, time_spent_since_last_sleep_ms:int):
    """
    Optional function which can implement sleep functionality for the signal.
    Raises RuntimeError if a fault is detected.
    """

  def wait(self, value:int, timeout:int|None=None):
    """
    Waits the signal is greater than or equal to a specific value.

    Args:
      value: The value to wait for.
      timeout: Maximum time to wait in milliseconds. Defaults to 30s.
    """
    timeout = timeout or getenv("HCQDEV_WAIT_TIMEOUT_MS", 30000)
    start_time = int(time.perf_counter() * 1000)
    while (not_passed:=(prev_value:=self.value) < value) and (cur_time:=int(time.perf_counter() * 1000)) - start_time < timeout:
      self._sleep(cur_time - start_time)
      if self.value != prev_value: start_time = int(time.perf_counter() * 1000) # progress was made, reset timer
    if not_passed and self.value < value: raise RuntimeError(f"Wait timeout: {timeout} ms! (the signal is not set to {value}, but {self.value})")

@contextlib.contextmanager
def hcq_profile(dev:HCQCompiled, enabled, desc, queue_type:Callable[[], HWQueue]|None=None, queue:HWQueue|None=None, dev_suff:str|None=None,
                profile_key:bytes|None=None):
  st, en = (dev.new_signal(), dev.new_signal()) if enabled else (None, None)
  assert queue is not None or queue_type is not None, "Either queue or queue_type must be provided"

  if enabled and queue is not None: queue.timestamp(st)
  elif enabled and queue_type is not None:
    queue_type().wait(dev.timeline_signal, dev.timeline_value - 1).timestamp(st).signal(dev.timeline_signal, dev.next_timeline()).submit(dev)

  try: yield (st, en)
  finally:
    if enabled and queue is not None: queue.timestamp(en)
    elif enabled and queue_type is not None:
      queue_type().wait(dev.timeline_signal, dev.timeline_value - 1).timestamp(en).signal(dev.timeline_signal, dev.next_timeline()).submit(dev)

    if enabled and PROFILE: dev.sig_prof_records.append((unwrap(st), unwrap(en), desc, f"{dev.device}:{dev_suff}" if dev_suff else dev.device,
                                                         profile_key))

class HCQArgsState(Generic[ProgramType]):
  def __init__(self, buf:HCQBuffer, prg:ProgramType, bufs:tuple[HCQBuffer, ...], vals:tuple[sint|None, ...]=()):
    self.buf, self.prg, self.bufs, self.vals = buf, prg, bufs, vals
    self.bind_data:list[tuple[tuple[sint, ...], MMIOInterface, str]] = []

  def bind_sints_to_buf(self, *vals:sint, buf:HCQBuffer, fmt, offset=0): self.bind_data.append((vals, buf.cpu_view().view(offset=offset), fmt))

class CLikeArgsState(HCQArgsState[ProgramType]):
  def __init__(self, buf:HCQBuffer, prg:ProgramType, bufs:tuple[HCQBuffer, ...], vals:tuple[sint|None, ...]=(), prefix:list[int]|None=None):
    super().__init__(buf, prg, bufs, vals=vals)

    if prefix is not None: self.buf.cpu_view().view(size=len(prefix) * 4, fmt='I')[:] = array.array('I', prefix)

    self.bind_sints_to_buf(*[b.va_addr for b in bufs], buf=self.buf, fmt='Q', offset=len(prefix or []) * 4)
    for v,(val_offset,dt) in zip(vals, TinyELF.iter_sig(prg.signature[-len(vals):], len(bufs) * 8)):
      assert v is not None
      self.bind_sints_to_buf(v, buf=self.buf, fmt=dt.fmt, offset=len(prefix or []) * 4 + val_offset)

class HCQProgram(Program[HCQDeviceType]):
  def __init__(self, args_state_t:Type[HCQArgsState], dev:HCQDeviceType, obj:TinyELF, kernargs_alloc_size:int, base:int|None=None):
    self.args_state_t, self.dev, self.name, self.signature, self.kernargs_alloc_size = args_state_t, dev, obj.name, obj.signature, kernargs_alloc_size
    self.profile_key = obj.profile_key
    self.prof_prg_counter = next(self.dev.prof_prg_counter)
    if PROFILE: Compiled.profile_events += [ProfileProgramEvent(dev.device, obj.name, obj.lib, base, self.prof_prg_counter, self.profile_key)]

  @staticmethod
  def _fini(dev, buf, spec): dev.allocator.free(BufferStorage(buf, buf.meta, buf.view), buf.size, spec)

  def fill_kernargs(self, bufs:tuple[HCQBuffer, ...], vals:tuple[int|None, ...]=(), kernargs:HCQBuffer|None=None) -> HCQArgsState:
    """
    Fills arguments for the kernel, optionally allocating space from the device if `kernargs_ptr` is not provided.
    Args:
      bufs: Buffers to be written to kernel arguments.
      vals: Values to be written to kernel arguments.
      kernargs_ptr: Optional pointer to pre-allocated kernel arguments memory.
    Returns:
      Arguments state with the given buffers and values set for the program.
    """
    argsbuf = kernargs or self.dev.kernargs_buf.offset(offset=self.dev.kernargs_offset_allocator.alloc(self.kernargs_alloc_size, 8),
                                                       size=self.kernargs_alloc_size)
    return self.args_state_t(argsbuf, self, bufs, vals=vals)

  def __call__(self, *bufs:HCQBuffer, global_size:tuple[int,int,int]=(1,1,1), local_size:tuple[int,int,int]=(1,1,1),
               vals:tuple[int|None, ...]=(), wait:bool=False, timeout:int|None=None) -> float|None:
    """
    Enqueues the program for execution with the given arguments and dimensions.

    Args:
      bufs: Buffer arguments to execute the kernel with.
      global_size: Specifies the global work size for kernel execution (equivalent to CUDA's grid size).
      local_size: Specifies the local work size for kernel execution (equivalent to CUDA's block size).
      vals: Value arguments to execute the kernel with.
      wait: If True, waits for the kernel to complete execution.

    Returns:
      Execution time of the kernel if 'wait' is True, otherwise None.
    """

    kernargs = self.fill_kernargs(bufs, vals)
    q = unwrap(self.dev.hw_compute_queue_t)().wait(self.dev.timeline_signal, self.dev.timeline_value - 1).memory_barrier()

    self.dev.prof_exec_counter += 1
    with hcq_profile(self.dev, queue=q, desc=self.name, enabled=wait or PROFILE, profile_key=self.profile_key) as (sig_st, sig_en):
      q.exec(self, kernargs, global_size, local_size)

    q.signal(self.dev.timeline_signal, self.dev.next_timeline()).submit(self.dev)

    if wait: self.dev.synchronize(timeout=timeout)
    return (float(sig_en.timestamp - sig_st.timestamp) / 1e6) if wait else None

class HCQCompiled(Compiled, Generic[SignalType]):
  """
  A base class for devices compatible with the HCQ (Hardware Command Queue) API.
  """
  peer_groups: dict[str, list[HCQCompiled]] = collections.defaultdict(list)
  signal_pages: dict[str, list[HCQBuffer]] = collections.defaultdict(list) # per peer group
  signal_pool: dict[str, list[HCQBuffer]] = collections.defaultdict(list) # per peer group
  cpu_devices: list[HCQCompiled] = []

  def __init__(self, device:str, allocator:HCQAllocatorBase, compilers:list[type[Renderer]], runtime:type[Program]|None,
               signal_t:Type[SignalType]|None=None, comp_queue_t:Callable[..., HWQueue]|None=None, copy_queue_t:Callable[..., HWQueue]|None=None,
               kernargs_size=(16 << 20), sigalloc_size=0x1000, can_recover:bool=False, arch=None):
    from extra.hcq1.graph import HCQGraph
    super().__init__(device, allocator, compilers, runtime, arch=arch)
    self.graph = HCQGraph

    self.peer_group = getattr(getattr(self, 'iface', None), 'peer_group', device.split(":")[0])
    HCQCompiled.peer_groups[self.peer_group].append(self)

    self.signal_t, self.hw_compute_queue_t, self.hw_copy_queue_t = signal_t, comp_queue_t, copy_queue_t

    self.timeline_value:int = 1
    self.sig_prof_records:list[tuple[HCQSignal, HCQSignal, str|TracingKey, str, bytes|None]] = []
    self.prof_exec_counter:int = 0
    self.prof_prg_counter = itertools.count(0)

    if signal_t is not None:
      # Map signals if any
      for sig_page in HCQCompiled.signal_pages[self.peer_group]: cast(HCQAllocator, self.allocator)._map(sig_page)

      self.sigalloc_size = sigalloc_size
      self.timeline_signal, self._shadow_timeline_signal = self.new_signal(value=0, is_timeline=True), self.new_signal(value=0, is_timeline=True)

    if comp_queue_t is not None:
      self.kernargs_buf:HCQBuffer = self.allocator.alloc(kernargs_size, BufferSpec(cpu_access=True)).buf
      self.kernargs_offset_allocator:BumpAllocator = BumpAllocator(self.kernargs_buf.size, wrap=True)

    self.can_recover = can_recover # Whether the device can recover from faults or timeouts
    self.error_state:Exception|None = None # Exception if error is unrecoverable and sync will always fail

    if self._is_cpu(): HCQCompiled.cpu_devices.append(self)

  def synchronize(self, timeout:int|None=None):
    if self.error_state is not None: raise self.error_state
    if not hasattr(self, 'timeline_signal'): return

    # If we have any work on CPU devices, need to synchronize them. This is just an optimization to release GIL allowing to finish faster.
    if not self._is_cpu():
      for dev in HCQCompiled.cpu_devices: dev.synchronize()

    try: self.timeline_signal.wait(self.timeline_value - 1, timeout=timeout if timeout is not None and self.can_recover else None)
    except RuntimeError as e:
      self.error_state = e
      if hasattr(self, 'on_device_hang'): self.on_device_hang()
      raise e

    if self.timeline_value > (1 << 31): self._wrap_timeline_signal()
    if PROFILE:
      Compiled.profile_events += [ProfileRangeEvent(dev, name, st.timestamp, en.timestamp, pk) for st,en,name,dev,pk in self.sig_prof_records]
      self.sig_prof_records = []

  def next_timeline(self):
    self.timeline_value += 1
    return self.timeline_value - 1

  def new_signal(self, **kwargs) -> SignalType:
    assert self.signal_t is not None, "Device does not support signals"
    if not HCQCompiled.signal_pool[pg:=self.peer_group]:
      HCQCompiled.signal_pages[pg].append(alc:=self.allocator.alloc(self.sigalloc_size, BufferSpec(host=True, uncached=True, cpu_access=True)).buf)
      HCQCompiled.signal_pool[pg] += [alc.offset(offset=off, size=16) for off in range(0, alc.size, 16)]
      for dev in HCQCompiled.peer_groups[pg]: cast(HCQAllocator, dev.allocator)._map(alc)
    return self.signal_t(base_buf=HCQCompiled.signal_pool[pg].pop(), owner=self, **kwargs)

  def device_props(self) -> dict[str,Any]: return {} # to be overridden if needed. dict keys are backend dependent.

  def hw_compute_queues(self) -> list[tuple[str|None, Callable[[], HWQueue]]]:
    return [(None, self.hw_compute_queue_t)] if self.hw_compute_queue_t is not None else []
  def hw_copy_queues(self) -> list[tuple[str, Callable[[], HWQueue]]]:
    return [("SDMA:0", self.hw_copy_queue_t)] if self.hw_copy_queue_t is not None else []

  def _at_profile_finalize(self):
    self.synchronize() # Expect device to be synchronizes

    def _sync(d:HCQCompiled, q_t:Callable[[], HWQueue]):
      q_t().timestamp(d.timeline_signal).signal(d.timeline_signal, d.next_timeline()).submit(d)
      st = time.perf_counter_ns()
      d.timeline_signal.wait(d.timeline_value - 1)  # average of the two
      et = time.perf_counter_ns()
      return (decimal.Decimal(et+st) / 2000) - d.timeline_signal.timestamp

    for prefix, q_t in self.hw_compute_queues() + self.hw_copy_queues():
      devname = f"{self.device}:{prefix}" if prefix else self.device
      Compiled.profile_events += [ProfileDeviceEvent(devname, statistics.median([_sync(self, q_t) for _ in range(40)]), props=self.device_props())]

  def _wrap_timeline_signal(self):
    self.timeline_signal, self._shadow_timeline_signal, self.timeline_value = self._shadow_timeline_signal, self.timeline_signal, 1
    self.timeline_signal.value = 0
    cast(HCQAllocatorBase, self.allocator).b_timeline = [0] * len(cast(HCQAllocatorBase, self.allocator).b)

  def _realloc(self, oldbuf:HCQBuffer|None, new_size:int, options:BufferSpec|None=None, force=False) -> tuple[HCQBuffer, bool]:
    if oldbuf is not None: self.allocator.free(BufferStorage(oldbuf, oldbuf.meta, oldbuf.view), oldbuf.size, options=options)
    try: buf, realloced = self.allocator.alloc(new_size, options=options).buf, True
    except MemoryError:
      if force: raise
      buf, realloced = self.allocator.alloc(oldbuf.size if oldbuf is not None else new_size, options=options).buf, False
    return buf, realloced

  def _is_cpu(self) -> bool: return hasattr(self, 'device') and self.device.split(":")[0] == "CPU"

  def rdma_dev(self):
    from extra.hcq1.ops_rdma import get_rdma_device
    for i in itertools.count():
      if (dev:=next((d for d in HCQCompiled.peer_groups[self.peer_group] if type(d).__name__ == 'RDMADevice'), None)): return dev
      try: get_rdma_device(i)
      except IndexError: raise RuntimeError(f"No RDMA found for peer group '{self.peer_group}'")

  def finalize(self):
    try: self.synchronize() # Try to finalize device in any case.
    except RuntimeError as e: print(f"{self.device} synchronization failed before finalizing: {e}")
    super().finalize()

class HCQAllocatorBase(Allocator[HCQDeviceType], Generic[HCQDeviceType]):
  """
  A base allocator class compatible with the HCQ (Hardware Command Queue) API.

  This class implements basic copy operations following the HCQ API, utilizing both types of `HWQueue`.
  """

  def __init__(self, dev:HCQDeviceType, batch_size:int=(2 << 20), batch_cnt:int=32, copy_bufs=None, **kwargs):
    super().__init__(dev, **kwargs)
    self.b = copy_bufs or [self._alloc(batch_size, BufferSpec(host=True)).buf for _ in range(batch_cnt)]
    self.b_timeline, self.b_next = [0] * len(self.b), 0

  def map(self, buf:Buffer) -> BufferStorage: return BufferStorage(*self._map(buf.ensure_allocated()._buf))

  def _map(self, buf:HCQBuffer) -> tuple:
    if self.dev not in buf.mapped_devs:
      if buf.owner is None: raise RuntimeError(f"map failed: buffer {buf.va_addr} has no owner, it's a virtual buffer")
      if not hasattr(self, '_do_map'): raise NotImplementedError("map failed: no method implemented")
      if (mb:=self._do_map(buf)) is not None: buf.mappings[self.dev] = mb
      buf.mapped_devs.append(self.dev)
    mapped = buf.mappings.get(self.dev, buf)
    return mapped, mapped.meta

  @suppress_finalizing
  def _free(self, storage:BufferStorage, options:BufferSpec|None=None):
    for dev in storage.buf.mapped_devs: dev.synchronize()
    for d, mb in storage.buf.mappings.items(): d.allocator._do_unmap(mb)
    if hasattr(self, '_do_free'): self._do_free(storage.buf, options)

  def _do_unmap(self, mb): self.dev.iface.free(mb)

  def _offset(self, buf, size:int, offset:int) -> HCQBuffer: return buf.offset(offset=offset, size=size)

class HCQAllocator(HCQAllocatorBase, Generic[HCQDeviceType]):
  def _copyin(self, dest:HCQBuffer, src:memoryview):
    if self.dev.hw_copy_queue_t is None:
      self.dev.synchronize()
      with cpu_profile(f'TINY -> {self.dev.device}', f"{self.dev.device}:COPY"): ctypes.memmove(int(dest.va_addr), from_mv(src), len(src))
      return

    with hcq_profile(self.dev, queue_type=self.dev.hw_copy_queue_t, desc=TracingKey(f"TINY -> {self.dev.device}", ret=src.nbytes), enabled=PROFILE,
                     dev_suff="SDMA:0"):
      for i in range(0, src.nbytes, self.b[0].size):
        self.b_next = (self.b_next + 1) % len(self.b)
        self.dev.timeline_signal.wait(self.b_timeline[self.b_next])

        lsize = min(self.b[self.b_next].size, src.nbytes - i)
        self.b[self.b_next].cpu_view().view(size=lsize, fmt='B')[:] = src.cast('B')[i:i+lsize]
        self.dev.hw_copy_queue_t().wait(self.dev.timeline_signal, self.dev.timeline_value - 1) \
                                  .copy(dest.offset(i), self.b[self.b_next], lsize) \
                                  .signal(self.dev.timeline_signal, self.dev.next_timeline()).submit(self.dev)
        self.b_timeline[self.b_next] = self.dev.timeline_value - 1

  def copy_from_disk(self, dest:HCQBuffer, src, size):
    def _get_temp_buf():
      # Check if the next buffer is safe to be used (its signal has passed) and reserve it.
      if self.b_timeline[(self.b_next + 1) % len(self.b)] <= self.dev.timeline_signal.value:
        self.b_timeline[(self.b_next + 1) % len(self.b)], self.b_next = (1 << 64), (self.b_next + 1) % len(self.b)
        return (self.b[self.b_next].cpu_view(), self.b_next)
      return None

    assert self.dev.hw_copy_queue_t is not None
    with hcq_profile(self.dev, queue_type=self.dev.hw_copy_queue_t, desc=TracingKey(f"DISK -> {self.dev.device}", ret=size), enabled=PROFILE,
                     dev_suff="SDMA:0"):
      for (batch_info, dst_off, src_off, copy_size) in src.device.allocator._copyout_sharded(src, size, _get_temp_buf, seg_len=self.b[0].size,
                                                                                             use_ioring=type(self.b[0].cpu_view()) is MMIOInterface):
        self.dev.hw_copy_queue_t().wait(self.dev.timeline_signal, self.dev.timeline_value - 1) \
                                  .copy(dest.offset(dst_off), self.b[batch_info[1]].offset(src_off), copy_size) \
                                  .signal(self.dev.timeline_signal, self.dev.next_timeline()).submit(self.dev)
        self.b_timeline[batch_info[1]] = self.dev.timeline_value - 1

  def _copyout(self, dest:memoryview, src:HCQBuffer):
    self.dev.synchronize()
    if self.dev.hw_copy_queue_t is None:
      with cpu_profile(f'{self.dev.device} -> TINY', f"{self.dev.device}:COPY"): ctypes.memmove(from_mv(dest), int(src.va_addr), len(dest))
      return

    with hcq_profile(self.dev, queue_type=self.dev.hw_copy_queue_t, desc=TracingKey(f"{self.dev.device} -> TINY", ret=dest.nbytes), enabled=PROFILE,
                     dev_suff="SDMA:0"):
      for i in range(0, dest.nbytes, cp_size:=self.b[0].size):
        self.dev.hw_copy_queue_t().wait(self.dev.timeline_signal, self.dev.timeline_value - 1) \
                                  .copy(self.b[0], src.offset(i), lsize:=min(cp_size, dest.nbytes-i)) \
                                  .signal(self.dev.timeline_signal, self.dev.next_timeline()).submit(self.dev)
        self.dev.timeline_signal.wait(self.dev.timeline_value - 1)
        dest.cast('B')[i:i+lsize] = self.b[0].cpu_view().view(size=lsize, fmt='B')[:]

  def _transfer(self, dest:HCQBuffer, src:HCQBuffer, sz:int, src_dev:HCQDeviceType, dest_dev:HCQDeviceType):
    if src_dev.peer_group != dest_dev.peer_group: return src_dev.rdma_dev().allocator._transfer(dest, src, sz, src_dev, dest_dev)

    cast(HCQAllocator, src_dev.allocator)._map(dest)

    assert src_dev.hw_copy_queue_t is not None
    with hcq_profile(src_dev, queue_type=src_dev.hw_copy_queue_t, desc=TracingKey(f"{src_dev.device} -> {dest_dev.device}", ret=sz), enabled=PROFILE,
                     dev_suff="SDMA:0"):
      src_dev.hw_copy_queue_t().wait(src_dev.timeline_signal, src_dev.timeline_value - 1) \
                               .wait(dest_dev.timeline_signal, dest_dev.timeline_value - 1) \
                               .copy(dest, src, sz) \
                               .signal(src_dev.timeline_signal, src_dev.next_timeline()).submit(src_dev)

    if src_dev != dest_dev:
      unwrap(dest_dev.hw_compute_queue_t)().wait(src_dev.timeline_signal, src_dev.timeline_value - 1) \
                                           .wait(dest_dev.timeline_signal, dest_dev.timeline_value - 1) \
                                           .signal(dest_dev.timeline_signal, dest_dev.next_timeline()).submit(dest_dev)
